From 9b82dd2f5f8811dd8d174af265e244acb5267972 Mon Sep 17 00:00:00 2001 From: h-guo18 <67671475+h-guo18@users.noreply.github.com> Date: Tue, 22 Sep 2026 13:42:01 +0000 Subject: [PATCH 01/11] [Speculative Decoding] Add the DFlash2 draft variant Adds DFlash2 (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: a grouped dynamic depthwise convolution around every attention and MLP sublayer, giving each block position a view of its predecessors inside the block without the taps crossing the block boundary; and a 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) and the way modeling_lilicorr installs the same convolution class. This also unblocks a recipe already shipped on main: modeling_lilicorr._install_sublayer_convs imports DFlashGroupedConv from modeling_dflash2, so lilicorr_conv.yaml raises at model build today. LiLiCorr's own initialization is unchanged and is now covered by tests -- it assigns kernel_projection explicitly, so it holds whichever way DFlashGroupedConv initializes itself -- and the two texts on main that described DFlash2's older random init are corrected. 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, has merged with no change to the checkpoint contract. modeling_dflash2.py is adapted from sgl-project/SpecForge#772 and carries its MIT notice ahead of the NVIDIA dual-SPDX header, matching modeling_dflash.py. No new dependencies. Signed-off-by: h-guo18 <67671475+h-guo18@users.noreply.github.com> --- .pre-commit-config.yaml | 2 + CHANGELOG.rst | 4 + .../torch/export/plugins/hf_spec_export.py | 43 ++ modelopt/torch/speculative/config.py | 11 + .../torch/speculative/dflash/conversion.py | 10 +- .../torch/speculative/plugins/__init__.py | 1 + .../torch/speculative/plugins/hf_dflash2.py | 229 +++++++++ .../speculative/plugins/modeling_dflash2.py | 317 ++++++++++++ .../speculative/plugins/modeling_lilicorr.py | 18 +- .../general/speculative_decoding/dflash2.yaml | 99 ++++ .../speculative_decoding/lilicorr_conv.yaml | 13 +- .../speculative/plugins/test_hf_dflash2.py | 479 ++++++++++++++++++ .../speculative/plugins/test_hf_lilicorr.py | 93 ++++ .../Qwen/Qwen3-8B/hf_online_dflash2.yaml | 107 ++++ 14 files changed, 1411 insertions(+), 15 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 diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 7654acab550..0078076d0bf 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 0a6afcb0017..b506bb565ae 100755 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -21,6 +21,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 06be11b8a57..38ee13c7aaa 100644 --- a/modelopt/torch/export/plugins/hf_spec_export.py +++ b/modelopt/torch/export/plugins/hf_spec_export.py @@ -605,3 +605,46 @@ def _export_config(self): } ) 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 ca9213aea86..ce524358af1 100644 --- a/modelopt/torch/speculative/config.py +++ b/modelopt/torch/speculative/config.py @@ -366,6 +366,17 @@ class DFlashConfig(ModeloptBaseConfig): ), ) + 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 3a164538297..2f7152ab576 100644 --- a/modelopt/torch/speculative/dflash/conversion.py +++ b/modelopt/torch/speculative/dflash/conversion.py @@ -35,6 +35,12 @@ # ``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 36cde677d25..bbdcac7553d 100644 --- a/modelopt/torch/speculative/plugins/__init__.py +++ b/modelopt/torch/speculative/plugins/__init__.py @@ -31,6 +31,7 @@ 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_dflash2.py b/modelopt/torch/speculative/plugins/hf_dflash2.py new file mode 100644 index 00000000000..8c4aafabea7 --- /dev/null +++ b/modelopt/torch/speculative/plugins/hf_dflash2.py @@ -0,0 +1,229 @@ +# 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._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) + + # Where the gold token is absent from the top-k, overwrite the last (lowest + # scoring) slot with it, so every supervised position has a correct class. + gold_in_topk = (candidate_ids == target_ids.unsqueeze(-1)).any(dim=-1) + gold_slot = torch.where( + gold_in_topk, + (candidate_ids == target_ids.unsqueeze(-1)).float().argmax(dim=-1), + torch.full_like(target_ids, top_k - 1), + ) + gold_unary = logits.gather(-1, target_ids.unsqueeze(-1)) + candidate_ids = candidate_ids.scatter(-1, gold_slot.unsqueeze(-1), target_ids.unsqueeze(-1)) + unary_logits = unary_logits.scatter(-1, gold_slot.unsqueeze(-1), gold_unary) + + selector_logits = selector.score_candidates( + candidate_ids, unary_logits, hidden, predecessor_ids + ) + + flat_weights = weight_mask.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 = (gold_in_topk.reshape(-1).float() * flat_weights).sum() / denominator + # 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 _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 = super()._compute_loss( + logits, + input_ids, + anchor_positions, + block_keep_mask, + loss_mask, + base_logits, + draft_hidden=draft_hidden, + base_outputs=base_outputs, + ) + 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 00000000000..e54a9b0f956 --- /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_lilicorr.py b/modelopt/torch/speculative/plugins/modeling_lilicorr.py index 6644f97e522..7d8924d5412 100644 --- a/modelopt/torch/speculative/plugins/modeling_lilicorr.py +++ b/modelopt/torch/speculative/plugins/modeling_lilicorr.py @@ -527,14 +527,16 @@ def _install_sublayer_convs(self, config, taps: int, group_size: int) -> None: ``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 00000000000..aa48b15ffdb --- /dev/null +++ b/modelopt_recipes/general/speculative_decoding/dflash2.yaml @@ -0,0 +1,99 @@ +# 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 + 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 0fdea64b25c..7a7a06b9cfd 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/speculative/plugins/test_hf_dflash2.py b/tests/unit/torch/speculative/plugins/test_hf_dflash2.py new file mode 100644 index 00000000000..b3060ad811c --- /dev/null +++ b/tests/unit/torch/speculative/plugins/test_hf_dflash2.py @@ -0,0 +1,479 @@ +# 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 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 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 f9bdbbca355..a21a8bdecf4 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 @@ def test_pseudo_speculative_generate_still_runs(self): 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 00000000000..e300e0665b4 --- /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 From 3967717f2cfb3691d833f65ec2c09e76bff68d9b Mon Sep 17 00:00:00 2001 From: h-guo18 <67671475+h-guo18@users.noreply.github.com> Date: Wed, 23 Sep 2026 04:50:38 +0000 Subject: [PATCH 02/11] feat(speculative): DFlash2 acceptance-annealed objective, strict top-k selector Two changes that bring DFlash2's objective in line with the reference implementation's semantics. dflash_lk_loss_type selects what the block objective minimizes against the hard target. 'ce' is -log q(gold), today's behaviour and the default. '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) for a the mean q(gold) over supervised positions, so the objective moves from fitting the distribution to maximizing acceptance as acceptance improves. The share is detached, so it reshapes the objective without adding a gradient path. Both new terms read q(gold) off the per-position cross-entropy the backbone loss already forms, which the KD path never produces -- it optimizes a soft target instead. That combination is rejected at convert time rather than silently falling back to CE. The terms come from the base loss through a new default-off return_terms, so the target alignment and position weighting are shared rather than re-derived; it is marked TODO for promotion to a shared divergence seam when the family's loss code is refactored. The candidate selector now trains on the strict unary top-k. A gold token the backbone did not propose was previously substituted into the lowest-scoring slot, which supervised the selector on a candidate set serving never builds and taught it to override the unary ranking there. Those positions are a backbone recall failure, not a selector classification example: they now carry no selector gradient and leave the denominator. selector_coverage reports how often the problem was solvable at all, and reads exactly 1.0 on a fully covered batch. The degenerate cases are asserted bit-exact -- a unit CE share reproduces 'ce', a zero share reproduces 'tv' -- because a blend wired up backwards still produces a finite decreasing loss. The selector mask is tested against a covered contrast, since masking everything would pass a one-sided assertion. Co-Authored-By: Claude Opus 5 (1M context) Signed-off-by: h-guo18 <67671475+h-guo18@users.noreply.github.com> --- modelopt/torch/speculative/config.py | 35 ++++++ .../torch/speculative/plugins/hf_dflash.py | 31 ++++- .../torch/speculative/plugins/hf_dflash2.py | 78 ++++++++++--- .../general/speculative_decoding/dflash2.yaml | 5 + .../speculative/plugins/test_hf_dflash2.py | 109 ++++++++++++++++++ 5 files changed, 243 insertions(+), 15 deletions(-) diff --git a/modelopt/torch/speculative/config.py b/modelopt/torch/speculative/config.py index ce524358af1..913c82010a6 100644 --- a/modelopt/torch/speculative/config.py +++ b/modelopt/torch/speculative/config.py @@ -366,6 +366,41 @@ 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, diff --git a/modelopt/torch/speculative/plugins/hf_dflash.py b/modelopt/torch/speculative/plugins/hf_dflash.py index d2750ad2aef..f77568ecc63 100644 --- a/modelopt/torch/speculative/plugins/hf_dflash.py +++ b/modelopt/torch/speculative/plugins/hf_dflash.py @@ -73,7 +73,7 @@ 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.""" @@ -915,6 +930,7 @@ def _compute_loss( base_logits=None, draft_hidden=None, base_outputs=None, + return_terms=False, ): """Compute weighted cross-entropy (or KD) loss and accuracy. @@ -927,6 +943,11 @@ def _compute_loss( 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. @@ -1021,6 +1042,14 @@ def _compute_loss( 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 index 8c4aafabea7..625bc622688 100644 --- a/modelopt/torch/speculative/plugins/hf_dflash2.py +++ b/modelopt/torch/speculative/plugins/hf_dflash2.py @@ -73,6 +73,16 @@ def modify(self, config): ) 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): @@ -116,23 +126,21 @@ def _selector_loss(self, logits, target_ids, hidden, predecessor_ids, weight_mas unary_logits, candidate_ids = logits.topk(top_k, dim=-1) - # Where the gold token is absent from the top-k, overwrite the last (lowest - # scoring) slot with it, so every supervised position has a correct class. - gold_in_topk = (candidate_ids == target_ids.unsqueeze(-1)).any(dim=-1) - gold_slot = torch.where( - gold_in_topk, - (candidate_ids == target_ids.unsqueeze(-1)).float().argmax(dim=-1), - torch.full_like(target_ids, top_k - 1), - ) - gold_unary = logits.gather(-1, target_ids.unsqueeze(-1)) - candidate_ids = candidate_ids.scatter(-1, gold_slot.unsqueeze(-1), target_ids.unsqueeze(-1)) - unary_logits = unary_logits.scatter(-1, gold_slot.unsqueeze(-1), gold_unary) + # 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 ) - flat_weights = weight_mask.reshape(-1) + 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), @@ -146,11 +154,50 @@ def _selector_loss(self, logits, target_ids, hidden, predecessor_ids, weight_mas accuracy = ( (chosen == gold_slot.reshape(-1)).float() * flat_weights ).sum() / denominator - coverage = (gold_in_topk.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, @@ -169,7 +216,7 @@ def _compute_loss( stays the backbone's top-1, so DFlash and DFlash2 runs remain comparable; the selector's own accuracy is logged separately. """ - loss, accuracy = super()._compute_loss( + loss, accuracy, terms = super()._compute_loss( logits, input_ids, anchor_positions, @@ -178,7 +225,10 @@ def _compute_loss( 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 diff --git a/modelopt_recipes/general/speculative_decoding/dflash2.yaml b/modelopt_recipes/general/speculative_decoding/dflash2.yaml index aa48b15ffdb..1f02e44b1aa 100644 --- a/modelopt_recipes/general/speculative_decoding/dflash2.yaml +++ b/modelopt_recipes/general/speculative_decoding/dflash2.yaml @@ -79,6 +79,11 @@ dflash: # 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 diff --git a/tests/unit/torch/speculative/plugins/test_hf_dflash2.py b/tests/unit/torch/speculative/plugins/test_hf_dflash2.py index b3060ad811c..931a4a3ee46 100644 --- a/tests/unit/torch/speculative/plugins/test_hf_dflash2.py +++ b/tests/unit/torch/speculative/plugins/test_hf_dflash2.py @@ -383,6 +383,62 @@ def _selector(self, vocab=16, rank=4, top_k=3, hidden=8): 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. @@ -428,6 +484,59 @@ def test_greedy_path_feeds_each_choice_forward(self): 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).""" From f20a94e44e44664f28698f5276ad64da41b19ff1 Mon Sep 17 00:00:00 2001 From: h-guo18 <67671475+h-guo18@users.noreply.github.com> Date: Wed, 23 Sep 2026 05:01:24 +0000 Subject: [PATCH 03/11] feat(speculative): DFlash2 streaming example, and carry the base RoPE through the fake base The fake base dropped the target's RoPE base on Transformers 5. It read the flat rope_theta attribute, which v5 no longer keeps: loading Qwen3-8B's own config.json, whose rope_theta is 1e6, leaves the value only inside rope_parameters and removes the flat field. FakeBaseConfig therefore stored None, published no rope_parameters of its own, and HFDFlashModel.modify -- which prefers the dict -- fell through to the flat field and set the draft's rope_theta to None. That is the streaming and offline half of the fix #2342 landed for the online path, and it was never made: the draft injects the target's KV, so a draft built this way trains and exports without complaint against a RoPE base the target never used. test_fakebase.py had no RoPE coverage at all, which is why it survived. It now asserts both shapes resolve, that an unknown base stays None rather than becoming a wrong default, and that the config publishes the dict form consumers prefer. The streaming example mirrors hf_streaming_dflash.yaml at one serve node plus one trainer. Two deltas are DFlash2's: the capture list omits the final layer, since the recipe trains against the hard target and forms no teacher distribution, so capturing it would only move bytes the trainer never reads; and the smoke test keeps method "dflash", because vLLM has no dflash2 method and selects the path from the checkpoint's DFlash2DraftModel architecture instead. Co-Authored-By: Claude Opus 5 (1M context) Signed-off-by: h-guo18 <67671475+h-guo18@users.noreply.github.com> --- .../speculative/plugins/modeling_fakebase.py | 21 +++- .../speculative/plugins/test_fakebase.py | 47 ++++++++- .../Qwen/Qwen3-8B/hf_streaming_dflash2.yaml | 99 +++++++++++++++++++ 3 files changed, 165 insertions(+), 2 deletions(-) create mode 100644 tools/launcher/examples/Qwen/Qwen3-8B/hf_streaming_dflash2.yaml diff --git a/modelopt/torch/speculative/plugins/modeling_fakebase.py b/modelopt/torch/speculative/plugins/modeling_fakebase.py index 2b5fe989c03..e24542e15ca 100644 --- a/modelopt/torch/speculative/plugins/modeling_fakebase.py +++ b/modelopt/torch/speculative/plugins/modeling_fakebase.py @@ -72,6 +72,20 @@ _SAFETENSORS_SINGLE_FILENAMES = ["model.safetensors", "consolidated.safetensors"] +def _base_rope_theta(base_cfg): + """Read the base model's RoPE base, whichever shape its config uses. + + Transformers 5 moves rope_theta into a rope_parameters dict and drops the flat + attribute, so reading the flat field alone returns None for a Qwen3 target whose + real base is 1e6. The draft injects the target's KV, so a draft built that way + trains and exports without complaint against a RoPE base the target never used. + """ + rope_parameters = getattr(base_cfg, "rope_parameters", None) + if isinstance(rope_parameters, dict) and rope_parameters.get("rope_theta") is not None: + return rope_parameters["rope_theta"] + return getattr(base_cfg, "rope_theta", None) + + class FakeBaseConfig(PretrainedConfig): """Minimal config for FakeBaseModel that supports offline speculative decoding training.""" @@ -124,7 +138,12 @@ def __init__( ) self.intermediate_size = intermediate_size # For some drafter algo (e.g. DFlash) rope theta must match target model. Extract here. + # Published in both shapes: consumers built for Transformers 5 read the + # rope_parameters dict first, and a fake base that only carried the flat field + # would hand them nothing. self.rope_theta = rope_theta + if rope_theta is not None: + self.rope_parameters = {"rope_theta": rope_theta} if isinstance(dtype, str): dtype = getattr(torch, dtype) self.dtype = dtype @@ -203,7 +222,7 @@ def from_source(cls, source: str, trust_remote_code: bool = False) -> "FakeBaseM 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), + rope_theta=_base_rope_theta(base_cfg), final_norm_type=_select_final_norm_type( getattr(base_cfg, "model_type", None), base_cfg ), diff --git a/tests/unit/torch/speculative/plugins/test_fakebase.py b/tests/unit/torch/speculative/plugins/test_fakebase.py index cf6dfe1a6bc..8127e303810 100644 --- a/tests/unit/torch/speculative/plugins/test_fakebase.py +++ b/tests/unit/torch/speculative/plugins/test_fakebase.py @@ -24,7 +24,11 @@ pytest.importorskip("transformers") import transformers -from modelopt.torch.speculative.plugins.modeling_fakebase import FakeBaseModel +from modelopt.torch.speculative.plugins.modeling_fakebase import ( + FakeBaseConfig, + FakeBaseModel, + _base_rope_theta, +) from modelopt.torch.speculative.utils import load_vlm_or_llm _HIDDEN_SIZE = 16 @@ -162,3 +166,44 @@ def from_pretrained(*args, **kwargs): 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_reads_the_transformers_5_rope_parameters_dict(self): + """Transformers 5 moves rope_theta into rope_parameters and drops the flat field.""" + 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 not hasattr(config, "rope_theta"), "fixture no longer reproduces the v5 layout" + assert _base_rope_theta(config) == 1000000.0 + + def test_falls_back_to_a_flat_attribute(self): + """A config that only carries the flat field still resolves.""" + assert _base_rope_theta(transformers.PretrainedConfig(rope_theta=12345.0)) == 12345.0 + + def test_missing_everywhere_is_none(self): + assert _base_rope_theta(transformers.PretrainedConfig()) is None + + 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} + + 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/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 00000000000..2f22469ebe8 --- /dev/null +++ b/tools/launcher/examples/Qwen/Qwen3-8B/hf_streaming_dflash2.yaml @@ -0,0 +1,99 @@ +# 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. +# +# 3-step pipeline: +# task_0: Build input conversations (jsonl) +# task_1: Streaming train — vllm serve + DFlash2 trainer; hidden states over NIXL RDMA +# task_2: vLLM smoke test with the exported drafter +# +# task_1 uses nodes=2: node 0 runs vllm serve, node 1 the trainer. Tasks share +# /scratchspace to pass artifacts. +# +# 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, 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]. + # + # Unlike the DFlash example there is no final layer (36) in the capture list: the + # DFlash2 recipe trains against the hard target (dflash_self_logit_distillation: false, + # dflash_lk_loss_type: lambda), so no teacher distribution is formed and capturing the + # final layer would only move bytes the trainer never reads. Set + # dflash.dflash_lk_loss_type=ce and add 36 back to train with distillation instead. + 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=<> + - data.mode=streaming + - data.data_path=/scratchspace/data/train.jsonl + - training.output_dir=/scratchspace/dflash2 + - training.training_seq_len=4096 + - training.disable_tqdm=true + # Streaming corpus is prompt-only (the serve generates the response and we + # capture its hidden states), so there is no assistant span to mask -> train + # over the full sequence, same as the DFlash streaming example. + - training.answer_only_loss=false + - 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_TP: "1" + # 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: 1 + 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 From 129619664cba3a3a66b74be6dcaf1601c85af538 Mon Sep 17 00:00:00 2001 From: h-guo18 <67671475+h-guo18@users.noreply.github.com> Date: Wed, 23 Sep 2026 07:01:35 +0000 Subject: [PATCH 04/11] fix(speculative): make the DFlash2 streaming example runnable, and add a multi-GPU one The DFlash2 streaming example could not run as written. Its capture list deliberately omits the base's final layer, on the correct reasoning that the recipe trains against the hard target and so never forms a teacher distribution -- DFlashBaseModelOutput.from_offline_dict only reads base_model_hidden_states under need_logits. But the streaming dataset splits the captured planes unconditionally: _format always peels the last one off as the base hidden. With five capture ids the draft's fc was therefore handed four planes and died with "mat1 and mat2 shapes cannot be multiplied (16384x16384 and 20480x4096)". Pair the five ids with data.final_aux_is_base_hidden=true, which keeps all five as aux features; the alternative is to capture a sixth layer and move 20% more bytes per sample for a plane nothing reads. Two comments in that file were also wrong, in a way that matters because they justify a setting. The serve does not generate: the trainer POSTs the whole conversation as a prompt with max_tokens=1 and the connector captures the per-token hidden states of that prefill. The corpus therefore does carry the assistant turn -- the reason answer_only_loss stays false is that Qwen3-8B's stock chat template has no {% generation %} tags to locate it, and the file now points at the shipped template that does, as the gpt-oss streaming example already does. Streaming also never runs the base's transformer layers on the trainer, so use_fake_base_for_offline is on. The new multi-node example is the shape that was actually validated: one serve node at TP=4 and one 4-rank DDP trainer node, rather than a single GPU each, which leaves the trainer waiting on a one-GPU prefill. 600 steps in 403 s (0.67 s/step, global batch 16 x 4096 tokens), loss 17.0 -> 3.4, drafter exported. Site-specific settings that run needed -- aarch64 image, explicit walltime, IB pinning for both UCX and NCCL, node-local Triton cache -- are documented in the header rather than hardcoded, since the right value differs per cluster and a wrong one here fails silently. Co-Authored-By: Claude Opus 5 (1M context) Signed-off-by: h-guo18 <67671475+h-guo18@users.noreply.github.com> --- .../Qwen/Qwen3-8B/hf_streaming_dflash2.yaml | 28 +++- .../hf_streaming_dflash2_multi_node.yaml | 144 ++++++++++++++++++ 2 files changed, 166 insertions(+), 6 deletions(-) create mode 100644 tools/launcher/examples/Qwen/Qwen3-8B/hf_streaming_dflash2_multi_node.yaml diff --git a/tools/launcher/examples/Qwen/Qwen3-8B/hf_streaming_dflash2.yaml b/tools/launcher/examples/Qwen/Qwen3-8B/hf_streaming_dflash2.yaml index 2f22469ebe8..ea1413e54a5 100644 --- a/tools/launcher/examples/Qwen/Qwen3-8B/hf_streaming_dflash2.yaml +++ b/tools/launcher/examples/Qwen/Qwen3-8B/hf_streaming_dflash2.yaml @@ -44,22 +44,38 @@ pipeline: # # Unlike the DFlash example there is no final layer (36) in the capture list: the # DFlash2 recipe trains against the hard target (dflash_self_logit_distillation: false, - # dflash_lk_loss_type: lambda), so no teacher distribution is formed and capturing the - # final layer would only move bytes the trainer never reads. Set - # dflash.dflash_lk_loss_type=ce and add 36 back to train with distillation instead. + # dflash_lk_loss_type: lambda), so no teacher distribution is formed and the base + # hidden is never read -- DFlashBaseModelOutput.from_offline_dict only touches + # base_model_hidden_states under need_logits. Dropping that plane therefore has to be + # paired with data.final_aux_is_base_hidden=true below, because the streaming dataset + # splits the captured planes unconditionally. To train with distillation instead, set + # dflash.dflash_lk_loss_type=ce, add 36 back, and drop final_aux_is_base_hidden. 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=<> + # The trainer never runs the base's transformer layers in streaming mode -- the + # serve produces the hidden states -- so load only the embeddings, final norm and + # lm_head instead of all 36 layers on every rank. + - model.use_fake_base_for_offline=true - data.mode=streaming - data.data_path=/scratchspace/data/train.jsonl + # Keep all five captured planes as aux features. The streaming dataset otherwise + # peels the LAST plane off as base_model_hidden_states, which would hand the + # draft's fc 4x4096 where it wants 5x4096 ("mat1 and mat2 shapes cannot be + # multiplied"). See the capture-id note above for when to drop this instead. + - data.final_aux_is_base_hidden=true - training.output_dir=/scratchspace/dflash2 - training.training_seq_len=4096 - training.disable_tqdm=true - # Streaming corpus is prompt-only (the serve generates the response and we - # capture its hidden states), so there is no assistant span to mask -> train - # over the full sequence, same as the DFlash streaming example. + # The serve does NOT generate: the trainer POSTs the whole conversation as a + # prompt with max_tokens=1 and the connector captures the per-token hidden states + # of that prefill. The corpus therefore carries the assistant turn, but Qwen3-8B's + # stock chat template has no {% generation %} tags to locate it, so train over the + # full sequence. To mask to the assistant span instead, pass + # data.chat_template=examples/Qwen/Qwen3-8B/chat_template_train.jinja with + # training.answer_only_loss=true, as the gpt-oss streaming example does. - training.answer_only_loss=false - training.num_train_epochs=1 - training.max_steps=5000 diff --git a/tools/launcher/examples/Qwen/Qwen3-8B/hf_streaming_dflash2_multi_node.yaml b/tools/launcher/examples/Qwen/Qwen3-8B/hf_streaming_dflash2_multi_node.yaml new file mode 100644 index 00000000000..22b07aa9f63 --- /dev/null +++ b/tools/launcher/examples/Qwen/Qwen3-8B/hf_streaming_dflash2_multi_node.yaml @@ -0,0 +1,144 @@ +# DFlash2 streaming speculative decoding pipeline for Qwen3-8B — MULTI-GPU NODES. +# +# Same three steps and same transport as hf_streaming_dflash2.yaml; the difference is the +# node shape. That example gives serve and trainer one GPU each, which leaves the trainer +# waiting on a single-GPU prefill. Here each node contributes all four of its GPUs: +# node 0 is one vLLM replica at TP=4, node 1 is a 4-rank DDP trainer. See +# common/eagle3/train_eagle_streaming.sh for the dispatch and rendezvous. +# +# Scale by raising `nodes` and SERVE_NODES together; keep per-serve in-flight +# (trainer_ranks x STREAMING_NUM_WORKERS / serve_replicas) roughly constant, and +# HS_POOL_SLOTS a few times above it. +# +# 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 +# +# Validated on a GB200 cluster (4 GPU/node, aarch64): 600 steps in 403 s (0.67 s/step, +# global batch 16 x 4096 tokens), loss 17.0 -> 3.4, exported 81 tensors. That run fed a +# pre-staged corpus instead of task_0 and used the overrides listed under "Cluster +# overrides" below. +# +# Cluster overrides this file deliberately does NOT hardcode, because the right value is +# site-specific: +# * container — vllm/vllm-openai:latest is x86; an aarch64 cluster needs an aarch64 +# image. The image must ship nixl (recent vllm-openai images do); if it does not, +# the serve dies at connector init with ModuleNotFoundError. +# * slurm_config.time — some clusters have a DefaultTime of minutes, and a job that +# outlives it dies as a bare Slurm TIMEOUT with no traceback. +# * fabric — on an InfiniBand cluster, pin BOTH stacks to the working HCAs or each +# picks a dead one and wedges silently (UCX as NIXL_ERR_BACKEND on the first fetch, +# NCCL as a watchdog SIGABRT half an hour later): +# NIXL_BACKENDS: "UCX" +# UCX_TLS: "rc,ud,sm,self" # host transports only +# UCX_NET_DEVICES: ":1,:1,..." +# NCCL_IB_DISABLE: "0" +# NCCL_IB_HCA: ",,..." +# On EFA use NIXL_BACKENDS=LIBFABRIC with FI_PROVIDER=efa instead. +# * caches — srun runs with --no-container-mount-home, so a $HOME cache path silently +# resolves inside the container overlay and is lost with the job. TRITON_CACHE_DIR +# must additionally be node-local: concurrent ranks race on a shared filesystem. +# +# Usage: +# uv run launch.py --yaml examples/Qwen/Qwen3-8B/hf_streaming_dflash2_multi_node.yaml --yes + +job_name: Qwen3-8B_DFlash2_streaming_multi_node +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. + # Capture ids and the aux/base split are explained in hf_streaming_dflash2.yaml. + 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 does not generate — it prefills the corpus conversation — so the + # assistant turn is present, but Qwen3-8B's stock chat template has no + # {% generation %} tags to locate it. Pass + # data.chat_template=examples/Qwen/Qwen3-8B/chat_template_train.jinja with + # answer_only_loss=true to mask to the assistant span instead. + - 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, the producer silently skips capture, and the fetch hangs + # (the trainer has a fail-loud guard for it). 32 slots x 4608 tok x 5 planes x + # 4096 x 2 B = 5.4 GiB of pinned host memory; peak in-flight here is + # 4 ranks x 4 workers = 16. + - 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 From b35bf76e5238f8b157119f3ee584eed074ce23be Mon Sep 17 00:00:00 2001 From: h-guo18 <67671475+h-guo18@users.noreply.github.com> Date: Wed, 23 Sep 2026 08:06:34 +0000 Subject: [PATCH 05/11] test(speculative): stop pinning a transformers-5-only layout in the RoPE tests The tf_min CI job (transformers 4.57) failed on `assert not hasattr(config, "rope_theta"), "fixture no longer reproduces the v5 layout"`. The guard was doing its job -- it fired the moment the fixture stopped reproducing the condition -- but the condition it asserted is not universal: 4.57 has only the flat field and no `rope_parameters` dict at all, while 5.12 has only the dict. The assertion encoded the newer layout as if it were the only one. The test that uses a real config now asserts the outcome rather than the layout, since the reader has to work on both. The layouts themselves are pinned by two tests that build them explicitly instead of depending on what the installed version happens to produce, so neither can drift out from under the suite again. One of those is new coverage: a config carrying BOTH a dict and a disagreeing flat field must resolve to the dict. That is the case where getting it wrong is silent -- the draft trains and exports with a RoPE base the target does not use, and only misbehaves at serve time. Verified against transformers 4.57.1 and 5.12.1: all seven pass on both. Co-Authored-By: Claude Opus 5 (1M context) Signed-off-by: h-guo18 <67671475+h-guo18@users.noreply.github.com> --- .../speculative/plugins/test_fakebase.py | 29 +++++++++++++++++-- 1 file changed, 26 insertions(+), 3 deletions(-) diff --git a/tests/unit/torch/speculative/plugins/test_fakebase.py b/tests/unit/torch/speculative/plugins/test_fakebase.py index 8127e303810..9c7edf5980b 100644 --- a/tests/unit/torch/speculative/plugins/test_fakebase.py +++ b/tests/unit/torch/speculative/plugins/test_fakebase.py @@ -175,8 +175,15 @@ class TestFakeBaseRopeTheta: complaint and only misbehaves at serve time. """ - def test_reads_the_transformers_5_rope_parameters_dict(self): - """Transformers 5 moves rope_theta into rope_parameters and drops the flat field.""" + def test_reads_a_real_config_whichever_layout_it_uses(self): + """A real config resolves on every supported transformers version. + + Where the value lives moved underneath us: transformers 5.12 keeps it only in + ``rope_parameters``, while the minimum supported version (4.57) has only the flat + field and no dict at all. So this asserts the outcome and not the layout; the two + layouts are pinned individually by the tests below, which build them explicitly + rather than depending on what the installed version happens to produce. + """ config = transformers.Qwen3Config( hidden_size=32, num_hidden_layers=2, @@ -186,7 +193,23 @@ def test_reads_the_transformers_5_rope_parameters_dict(self): vocab_size=64, rope_theta=1000000.0, ) - assert not hasattr(config, "rope_theta"), "fixture no longer reproduces the v5 layout" + assert _base_rope_theta(config) == 1000000.0 + + def test_reads_the_rope_parameters_dict(self): + """The transformers 5.12+ layout: the value lives only in the dict.""" + config = transformers.PretrainedConfig(rope_parameters={"rope_theta": 1000000.0}) + assert _base_rope_theta(config) == 1000000.0 + + def test_prefers_the_dict_over_a_disagreeing_flat_field(self): + """Both present and disagreeing: the dict wins. + + A config can carry both, and they can disagree. Reading the flat field first would + give the draft a RoPE base the target does not use -- which trains and exports + without complaint, and only misbehaves at serve time. + """ + config = transformers.PretrainedConfig( + rope_theta=10000.0, rope_parameters={"rope_theta": 1000000.0} + ) assert _base_rope_theta(config) == 1000000.0 def test_falls_back_to_a_flat_attribute(self): From b82195b5800c594d6c21670f7aadff5b23beface Mon Sep 17 00:00:00 2001 From: h-guo18 <67671475+h-guo18@users.noreply.github.com> Date: Thu, 24 Sep 2026 07:05:04 +0000 Subject: [PATCH 06/11] refactor(speculative): read the base RoPE through the exporter's reader, and guard it This PR added a third implementation of "where does a config keep rope_theta", `_base_rope_theta`, next to the two already on main. That question has been answered independently in several places for months and re-fixed one call site at a time -- the exporter's `_get_rope_theta` carried the wrong precedence from 2026-07-30 to 2026-09-09, and this file read the flat attribute only from 2026-07-06 until this PR. Adding a fourth answer is how that continues, so the fake base now calls the exporter's reader, which is a strict superset (it also handles the legacy `rope_scaling` spelling). The other duplicates are left for a follow-up; they span export, utils and speculative and do not belong here. The reason it kept being re-fixed is that nothing failed when it was wrong. Both halves now have a guard, and both were confirmed by reverting the code they cover: * `TestGetRopeTheta::test_prefers_the_dict_over_a_disagreeing_flat_field` pins the precedence. A config can hold 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 yields a drafter whose RoPE base is 100x off. Flipping the order back to the 2026-07-30 form previously passed all 302 tests; it now fails. * `TestFakeBaseRopeTheta::test_from_source_carries_a_transformers_5_base_theta` pins the seam rather than the reader -- it goes through `from_source` and asserts a transformers-5-shaped base config reaches the FakeBaseConfig. Reverting the call site to a plain `getattr` now fails it. The reader tests move to the exporter's test file, where the reader lives, so there is one place to add to next time rather than one per caller. Verified on transformers 4.57.1 and 5.12.1: 304 passed on both. Co-Authored-By: Claude Opus 5 (1M context) Signed-off-by: h-guo18 <67671475+h-guo18@users.noreply.github.com> --- .../speculative/plugins/modeling_fakebase.py | 23 +++---- .../torch/export/test_hf_spec_rope_export.py | 67 +++++++++++++++++- .../speculative/plugins/test_fakebase.py | 68 ++++++++----------- 3 files changed, 101 insertions(+), 57 deletions(-) diff --git a/modelopt/torch/speculative/plugins/modeling_fakebase.py b/modelopt/torch/speculative/plugins/modeling_fakebase.py index e24542e15ca..7e2e0bf00c0 100644 --- a/modelopt/torch/speculative/plugins/modeling_fakebase.py +++ b/modelopt/torch/speculative/plugins/modeling_fakebase.py @@ -72,20 +72,6 @@ _SAFETENSORS_SINGLE_FILENAMES = ["model.safetensors", "consolidated.safetensors"] -def _base_rope_theta(base_cfg): - """Read the base model's RoPE base, whichever shape its config uses. - - Transformers 5 moves rope_theta into a rope_parameters dict and drops the flat - attribute, so reading the flat field alone returns None for a Qwen3 target whose - real base is 1e6. The draft injects the target's KV, so a draft built that way - trains and exports without complaint against a RoPE base the target never used. - """ - rope_parameters = getattr(base_cfg, "rope_parameters", None) - if isinstance(rope_parameters, dict) and rope_parameters.get("rope_theta") is not None: - return rope_parameters["rope_theta"] - return getattr(base_cfg, "rope_theta", None) - - class FakeBaseConfig(PretrainedConfig): """Minimal config for FakeBaseModel that supports offline speculative decoding training.""" @@ -198,6 +184,8 @@ def from_source(cls, source: str, trust_remote_code: bool = False) -> "FakeBaseM 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 ) @@ -222,7 +210,12 @@ def from_source(cls, source: str, trust_remote_code: bool = False) -> "FakeBaseM 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=_base_rope_theta(base_cfg), + # Shared with the exporter deliberately: where a config keeps rope_theta + # depends on the transformers version, and a local getattr got that wrong + # here for two months while the exporter had it right. The draft injects + # the target's KV, so a wrong base trains and exports without complaint + # and only misbehaves at 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/tests/unit/torch/export/test_hf_spec_rope_export.py b/tests/unit/torch/export/test_hf_spec_rope_export.py index fbeb218793e..6d23d58ca6e 100644 --- a/tests/unit/torch/export/test_hf_spec_rope_export.py +++ b/tests/unit/torch/export/test_hf_spec_rope_export.py @@ -19,8 +19,13 @@ from unittest.mock import MagicMock import torch +import transformers -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,63 @@ 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.""" + config = transformers.PretrainedConfig(rope_parameters={"rope_theta": 1000000.0}) + assert _get_rope_theta(config) == 1000000.0 + + def test_reads_the_legacy_rope_scaling_dict(self): + """Older transformers spell the same dict rope_scaling.""" + config = transformers.PretrainedConfig(rope_scaling={"rope_theta": 1000000.0}) + assert _get_rope_theta(config) == 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 = transformers.PretrainedConfig( + 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(transformers.PretrainedConfig(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(transformers.PretrainedConfig()) is None + assert _get_rope_theta(transformers.PretrainedConfig(), 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. + """ + 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 9c7edf5980b..a27729e1c74 100644 --- a/tests/unit/torch/speculative/plugins/test_fakebase.py +++ b/tests/unit/torch/speculative/plugins/test_fakebase.py @@ -24,11 +24,7 @@ pytest.importorskip("transformers") import transformers -from modelopt.torch.speculative.plugins.modeling_fakebase import ( - FakeBaseConfig, - FakeBaseModel, - _base_rope_theta, -) +from modelopt.torch.speculative.plugins.modeling_fakebase import FakeBaseConfig, FakeBaseModel from modelopt.torch.speculative.utils import load_vlm_or_llm _HIDDEN_SIZE = 16 @@ -175,49 +171,39 @@ class TestFakeBaseRopeTheta: complaint and only misbehaves at serve time. """ - def test_reads_a_real_config_whichever_layout_it_uses(self): - """A real config resolves on every supported transformers version. + def test_from_source_carries_a_transformers_5_base_theta(self, tmp_path, monkeypatch): + """The seam, not the reader. - Where the value lives moved underneath us: transformers 5.12 keeps it only in - ``rope_parameters``, while the minimum supported version (4.57) has only the flat - field and no dict at all. So this asserts the outcome and not the layout; the two - layouts are pinned individually by the tests below, which build them explicitly - rather than depending on what the installed version happens to produce. + 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. """ - config = transformers.Qwen3Config( - hidden_size=32, + base_cfg = transformers.PretrainedConfig( + model_type="llama", + hidden_size=_HIDDEN_SIZE, + vocab_size=_VOCAB_SIZE, num_hidden_layers=2, - num_attention_heads=4, - num_key_value_heads=2, - intermediate_size=64, - vocab_size=64, - rope_theta=1000000.0, + max_position_embeddings=128, + tie_word_embeddings=False, + rope_parameters={"rope_theta": 1000000.0}, ) - assert _base_rope_theta(config) == 1000000.0 - - def test_reads_the_rope_parameters_dict(self): - """The transformers 5.12+ layout: the value lives only in the dict.""" - config = transformers.PretrainedConfig(rope_parameters={"rope_theta": 1000000.0}) - assert _base_rope_theta(config) == 1000000.0 - - def test_prefers_the_dict_over_a_disagreeing_flat_field(self): - """Both present and disagreeing: the dict wins. - - A config can carry both, and they can disagree. Reading the flat field first would - give the draft a RoPE base the target does not use -- which trains and exports - without complaint, and only misbehaves at serve time. - """ - config = transformers.PretrainedConfig( - rope_theta=10000.0, 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)}) ) - assert _base_rope_theta(config) == 1000000.0 - def test_falls_back_to_a_flat_attribute(self): - """A config that only carries the flat field still resolves.""" - assert _base_rope_theta(transformers.PretrainedConfig(rope_theta=12345.0)) == 12345.0 + model = FakeBaseModel.from_source(str(tmp_path)) - def test_missing_everywhere_is_none(self): - assert _base_rope_theta(transformers.PretrainedConfig()) is None + assert model.config.rope_theta == 1000000.0 + assert model.config.rope_parameters == {"rope_theta": 1000000.0} def test_config_publishes_both_shapes(self): """Consumers that prefer the dict must find it on a fake base too.""" From 9027c42e1c5a6dfe693c0a557b053e01815842bb Mon Sep 17 00:00:00 2001 From: h-guo18 <67671475+h-guo18@users.noreply.github.com> Date: Thu, 24 Sep 2026 07:26:30 +0000 Subject: [PATCH 07/11] docs(launcher): one streaming example for DFlash2, under the plain name The two DFlash2 streaming examples differed by eleven lines and neither was single-node: both ran two nodes, one serve and one trainer, and the split was one GPU each versus four. Keeping both meant maintaining the same configuration twice, which is how the two drift. The four-GPU shape is the one that has actually been run end to end (600 steps, 0.67 s/step, loss 17.0 -> 3.4, drafter exported), so that is what survives, and it takes the plain name. Only two of the repo's nineteen streaming examples ship a base/_multi_node pair; thirteen are _multi_node alone, so one file per variant is the common shape here. Dropping the suffix rather than the file is deliberate. Every other _multi_node example sets SERVE_NODES >= 2 -- the suffix marks the serve fan-out path, not the GPU count -- and this one has a single serve replica, so with no base file left beside it the name would have been the only wrong one in the tree and nothing would have made that visible. The header now says how to go both ways: down to one GPU per node, and up to several serve replicas. Co-Authored-By: Claude Opus 5 (1M context) Signed-off-by: h-guo18 <67671475+h-guo18@users.noreply.github.com> --- .../Qwen/Qwen3-8B/hf_streaming_dflash2.yaml | 86 ++++++++--- .../hf_streaming_dflash2_multi_node.yaml | 144 ------------------ 2 files changed, 66 insertions(+), 164 deletions(-) delete mode 100644 tools/launcher/examples/Qwen/Qwen3-8B/hf_streaming_dflash2_multi_node.yaml diff --git a/tools/launcher/examples/Qwen/Qwen3-8B/hf_streaming_dflash2.yaml b/tools/launcher/examples/Qwen/Qwen3-8B/hf_streaming_dflash2.yaml index ea1413e54a5..ae334b188cf 100644 --- a/tools/launcher/examples/Qwen/Qwen3-8B/hf_streaming_dflash2.yaml +++ b/tools/launcher/examples/Qwen/Qwen3-8B/hf_streaming_dflash2.yaml @@ -3,15 +3,49 @@ # 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. +# 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. +# +# On a cluster with fewer GPUs per node, drop gpus_per_node and SERVE_TP together (1 and 1 +# still works, it just leaves the trainer waiting on a single-GPU prefill) and lower +# per_device_train_batch_size to keep the global batch you want. To scale up instead, +# raise `nodes` and SERVE_NODES together so the trainer fans out over several serve +# replicas, as hf_streaming_dflash_multi_node.yaml does; keep per-serve in-flight +# (trainer_ranks x STREAMING_NUM_WORKERS / serve_replicas) roughly constant and +# HS_POOL_SLOTS a few times above it. # # 3-step pipeline: # task_0: Build input conversations (jsonl) -# task_1: Streaming train — vllm serve + DFlash2 trainer; hidden states over NIXL RDMA +# task_1: Streaming train — 1 serve node (TP=4) + 1 trainer node (4-rank DDP) # task_2: vLLM smoke test with the exported drafter # -# task_1 uses nodes=2: node 0 runs vllm serve, node 1 the trainer. Tasks share -# /scratchspace to pass artifacts. +# Validated on a GB200 cluster (4 GPU/node, aarch64): 600 steps in 403 s (0.67 s/step, +# global batch 16 x 4096 tokens), loss 17.0 -> 3.4, exported 81 tensors. That run fed a +# pre-staged corpus instead of task_0 and used the overrides listed under "Cluster +# overrides" below. +# +# Cluster overrides this file deliberately does NOT hardcode, because the right value is +# site-specific: +# * container — vllm/vllm-openai:latest is x86; an aarch64 cluster needs an aarch64 +# image. The image must ship nixl (recent vllm-openai images do); if it does not, +# the serve dies at connector init with ModuleNotFoundError. +# * slurm_config.time — some clusters have a DefaultTime of minutes, and a job that +# outlives it dies as a bare Slurm TIMEOUT with no traceback. +# * fabric — on an InfiniBand cluster, pin BOTH stacks to the working HCAs or each +# picks a dead one and wedges silently (UCX as NIXL_ERR_BACKEND on the first fetch, +# NCCL as a watchdog SIGABRT half an hour later): +# NIXL_BACKENDS: "UCX" +# UCX_TLS: "rc,ud,sm,self" # host transports only +# UCX_NET_DEVICES: ":1,:1,..." +# NCCL_IB_DISABLE: "0" +# NCCL_IB_HCA: ",,..." +# On EFA use NIXL_BACKENDS=LIBFABRIC with FI_PROVIDER=efa instead. +# * caches — srun runs with --no-container-mount-home, so a $HOME cache path silently +# resolves inside the container overlay and is lost with the job. TRITON_CACHE_DIR +# must additionally be node-local: concurrent ranks race on a shared filesystem. # # Usage: # uv run launch.py --yaml examples/Qwen/Qwen3-8B/hf_streaming_dflash2.yaml --yes @@ -38,7 +72,7 @@ pipeline: gpus_per_node: 1 container: nvcr.io/nvidia/tensorrt-llm/release:1.3.0rc20 - # Step 2: Streaming DFlash2 training — node 0 vllm serve, node 1 trainer. + # 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]. # @@ -55,28 +89,26 @@ pipeline: args: - --config modules/Model-Optimizer/modelopt_recipes/general/speculative_decoding/dflash2.yaml - model.model_name_or_path=<> - # The trainer never runs the base's transformer layers in streaming mode -- the - # serve produces the hidden states -- so load only the embeddings, final norm and - # lm_head instead of all 36 layers on every rank. + # 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 - # Keep all five captured planes as aux features. The streaming dataset otherwise - # peels the LAST plane off as base_model_hidden_states, which would hand the - # draft's fc 4x4096 where it wants 5x4096 ("mat1 and mat2 shapes cannot be - # multiplied"). See the capture-id note above for when to drop this instead. + # 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 does NOT generate: the trainer POSTs the whole conversation as a - # prompt with max_tokens=1 and the connector captures the per-token hidden states - # of that prefill. The corpus therefore carries the assistant turn, but Qwen3-8B's - # stock chat template has no {% generation %} tags to locate it, so train over the - # full sequence. To mask to the assistant span instead, pass + # The serve does not generate — it prefills the corpus conversation — so the + # assistant turn is present, but Qwen3-8B's stock chat template has no + # {% generation %} tags to locate it. Pass # data.chat_template=examples/Qwen/Qwen3-8B/chat_template_train.jinja with - # training.answer_only_loss=true, as the gpt-oss streaming example does. + # answer_only_loss=true to mask to the assistant span instead. - 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 @@ -86,14 +118,28 @@ pipeline: - 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_TP: "1" + # 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, the producer silently skips capture, and the fetch hangs + # (the trainer has a fail-loud guard for it). 32 slots x 4608 tok x 5 planes x + # 4096 x 2 B = 5.4 GiB of pinned host memory; peak in-flight here is + # 4 ranks x 4 workers = 16. + - 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: 1 + gpus_per_node: 4 container: vllm/vllm-openai:latest # Step 3: vLLM smoke test (uses the exported checkpoint from training). diff --git a/tools/launcher/examples/Qwen/Qwen3-8B/hf_streaming_dflash2_multi_node.yaml b/tools/launcher/examples/Qwen/Qwen3-8B/hf_streaming_dflash2_multi_node.yaml deleted file mode 100644 index 22b07aa9f63..00000000000 --- a/tools/launcher/examples/Qwen/Qwen3-8B/hf_streaming_dflash2_multi_node.yaml +++ /dev/null @@ -1,144 +0,0 @@ -# DFlash2 streaming speculative decoding pipeline for Qwen3-8B — MULTI-GPU NODES. -# -# Same three steps and same transport as hf_streaming_dflash2.yaml; the difference is the -# node shape. That example gives serve and trainer one GPU each, which leaves the trainer -# waiting on a single-GPU prefill. Here each node contributes all four of its GPUs: -# node 0 is one vLLM replica at TP=4, node 1 is a 4-rank DDP trainer. See -# common/eagle3/train_eagle_streaming.sh for the dispatch and rendezvous. -# -# Scale by raising `nodes` and SERVE_NODES together; keep per-serve in-flight -# (trainer_ranks x STREAMING_NUM_WORKERS / serve_replicas) roughly constant, and -# HS_POOL_SLOTS a few times above it. -# -# 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 -# -# Validated on a GB200 cluster (4 GPU/node, aarch64): 600 steps in 403 s (0.67 s/step, -# global batch 16 x 4096 tokens), loss 17.0 -> 3.4, exported 81 tensors. That run fed a -# pre-staged corpus instead of task_0 and used the overrides listed under "Cluster -# overrides" below. -# -# Cluster overrides this file deliberately does NOT hardcode, because the right value is -# site-specific: -# * container — vllm/vllm-openai:latest is x86; an aarch64 cluster needs an aarch64 -# image. The image must ship nixl (recent vllm-openai images do); if it does not, -# the serve dies at connector init with ModuleNotFoundError. -# * slurm_config.time — some clusters have a DefaultTime of minutes, and a job that -# outlives it dies as a bare Slurm TIMEOUT with no traceback. -# * fabric — on an InfiniBand cluster, pin BOTH stacks to the working HCAs or each -# picks a dead one and wedges silently (UCX as NIXL_ERR_BACKEND on the first fetch, -# NCCL as a watchdog SIGABRT half an hour later): -# NIXL_BACKENDS: "UCX" -# UCX_TLS: "rc,ud,sm,self" # host transports only -# UCX_NET_DEVICES: ":1,:1,..." -# NCCL_IB_DISABLE: "0" -# NCCL_IB_HCA: ",,..." -# On EFA use NIXL_BACKENDS=LIBFABRIC with FI_PROVIDER=efa instead. -# * caches — srun runs with --no-container-mount-home, so a $HOME cache path silently -# resolves inside the container overlay and is lost with the job. TRITON_CACHE_DIR -# must additionally be node-local: concurrent ranks race on a shared filesystem. -# -# Usage: -# uv run launch.py --yaml examples/Qwen/Qwen3-8B/hf_streaming_dflash2_multi_node.yaml --yes - -job_name: Qwen3-8B_DFlash2_streaming_multi_node -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. - # Capture ids and the aux/base split are explained in hf_streaming_dflash2.yaml. - 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 does not generate — it prefills the corpus conversation — so the - # assistant turn is present, but Qwen3-8B's stock chat template has no - # {% generation %} tags to locate it. Pass - # data.chat_template=examples/Qwen/Qwen3-8B/chat_template_train.jinja with - # answer_only_loss=true to mask to the assistant span instead. - - 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, the producer silently skips capture, and the fetch hangs - # (the trainer has a fail-loud guard for it). 32 slots x 4608 tok x 5 planes x - # 4096 x 2 B = 5.4 GiB of pinned host memory; peak in-flight here is - # 4 ranks x 4 workers = 16. - - 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 From e10b003a371a9a0f0c33540ea3cbac73755cfc1a Mon Sep 17 00:00:00 2001 From: h-guo18 <67671475+h-guo18@users.noreply.github.com> Date: Fri, 25 Sep 2026 11:57:45 +0000 Subject: [PATCH 08/11] test(export): keep the RoPE tests collectable without transformers The `partial-install (torch)` CI job installs modelopt without transformers, and the module-level `import transformers` this file gained alongside the new `TestGetRopeTheta` class made collection fail there -- taking the pre-existing exporter tests in the same file down with it, which is worse than the new tests simply not running. The five layout tests never needed it: `_get_rope_theta` reads its input with `getattr`, so `SimpleNamespace` exercises the same path and is what the rest of this file already uses for fake configs. That keeps the precedence guard alive in the torch-only job rather than skipping it. Only the test that asserts a real config resolves genuinely needs transformers, and it now takes `pytest.importorskip` inside the test, as test_quant_aware_conversion.py does, so it skips alone. Verified three ways: with transformers blocked at import (13 passed, 1 skipped, no collection error), and with transformers 4.57.1 and 5.12.1 (14 passed each). The precedence mutation -- flipping `_get_rope_theta` back to reading the flat field first -- still fails the guard after the switch to SimpleNamespace. Co-Authored-By: Claude Opus 5 (1M context) Signed-off-by: h-guo18 <67671475+h-guo18@users.noreply.github.com> --- .../torch/export/test_hf_spec_rope_export.py | 26 ++++++++++--------- 1 file changed, 14 insertions(+), 12 deletions(-) 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 6d23d58ca6e..7cb0a317d06 100644 --- a/tests/unit/torch/export/test_hf_spec_rope_export.py +++ b/tests/unit/torch/export/test_hf_spec_rope_export.py @@ -18,8 +18,8 @@ from types import SimpleNamespace from unittest.mock import MagicMock +import pytest import torch -import transformers from modelopt.torch.export.plugins.hf_spec_export import ( DFlashExporter, @@ -170,13 +170,15 @@ class TestGetRopeTheta: def test_reads_the_rope_parameters_dict(self): """The transformers 5.12+ layout: the value lives only in the dict.""" - config = transformers.PretrainedConfig(rope_parameters={"rope_theta": 1000000.0}) - assert _get_rope_theta(config) == 1000000.0 + 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.""" - config = transformers.PretrainedConfig(rope_scaling={"rope_theta": 1000000.0}) - assert _get_rope_theta(config) == 1000000.0 + 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. @@ -186,27 +188,27 @@ def test_prefers_the_dict_over_a_disagreeing_flat_field(self): 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 = transformers.PretrainedConfig( - rope_theta=10000.0, rope_parameters={"rope_theta": 1000000.0} - ) + 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(transformers.PretrainedConfig(rope_theta=12345.0)) == 12345.0 + 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(transformers.PretrainedConfig()) is None - assert _get_rope_theta(transformers.PretrainedConfig(), 7.0) == 7.0 + 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. + 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, From 3d77d6cbaa6f2b9a895a5dc7bc370eed6fd87c07 Mon Sep 17 00:00:00 2001 From: h-guo18 <67671475+h-guo18@users.noreply.github.com> Date: Sun, 27 Sep 2026 08:10:40 +0000 Subject: [PATCH 09/11] fix(speculative): give the fake base's rope_parameters a rope_type The offline EAGLE3 example tests fail on this branch with modeling_eagle.py:84 LlamaRotaryEmbedding(config=self.config, ...) KeyError: 'rope_type' and take test_offline_resume_training_kimi down after them, since it reads the checkpoint the failed run never wrote. Both pass on main. Earlier in this PR `FakeBaseConfig` started publishing a `rope_parameters` dict so that consumers written for transformers 5 -- which read the dict before the flat field -- would find the target's RoPE base on a fake base too. That dict carried only `rope_theta`. The shape matters more than it looks, because this class is not only the base config: `hf_eagle.modify` builds the draft config as `type(base_config).from_dict(arch_config)`, so a FakeBaseConfig is also what the EAGLE draft is constructed from, and the dict lands where transformers' rotary embeddings index `rope_parameters["rope_type"]` unconditionally. A dict without it raises rather than falling back, and the manual publish pre-empts the well-formed dict transformers would otherwise build. `"default"` is both what every real config produces -- Qwen3Config and LlamaConfig each yield `{"rope_theta": ..., "rope_type": "default"}` -- and what the draft wants: the long-context scaling families are deliberately not inherited, they are re-applied at export. The new test builds a real `LlamaRotaryEmbedding` from a FakeBaseConfig and checks it produces usable tables, rather than asserting the dict's contents. Asserting contents is what the existing test did, and it stayed green through the regression: it pinned the value and not the requirement. Removing `rope_type` again fails three tests, including that one. Co-Authored-By: Claude Opus 5 (1M context) Signed-off-by: h-guo18 <67671475+h-guo18@users.noreply.github.com> --- .../speculative/plugins/modeling_fakebase.py | 9 +++++- .../speculative/plugins/test_fakebase.py | 30 +++++++++++++++++-- 2 files changed, 36 insertions(+), 3 deletions(-) diff --git a/modelopt/torch/speculative/plugins/modeling_fakebase.py b/modelopt/torch/speculative/plugins/modeling_fakebase.py index 7e2e0bf00c0..162033bc605 100644 --- a/modelopt/torch/speculative/plugins/modeling_fakebase.py +++ b/modelopt/torch/speculative/plugins/modeling_fakebase.py @@ -129,7 +129,14 @@ def __init__( # would hand them nothing. self.rope_theta = rope_theta if rope_theta is not None: - self.rope_parameters = {"rope_theta": rope_theta} + # rope_type is not optional padding: this class is also the class the EAGLE + # draft config is built from (`type(base_config).from_dict(arch_config)` in + # hf_eagle.modify), and transformers' rotary embeddings index + # config.rope_parameters["rope_type"] unconditionally. A dict carrying only + # rope_theta raises KeyError there. "default" matches what every real config + # produces and what the draft wants -- the long-context scaling families are + # deliberately not inherited, they are re-applied at export. + self.rope_parameters = {"rope_theta": rope_theta, "rope_type": "default"} if isinstance(dtype, str): dtype = getattr(torch, dtype) self.dtype = dtype diff --git a/tests/unit/torch/speculative/plugins/test_fakebase.py b/tests/unit/torch/speculative/plugins/test_fakebase.py index a27729e1c74..642af7c1a5b 100644 --- a/tests/unit/torch/speculative/plugins/test_fakebase.py +++ b/tests/unit/torch/speculative/plugins/test_fakebase.py @@ -203,13 +203,39 @@ def test_from_source_carries_a_transformers_5_base_theta(self, tmp_path, monkeyp model = FakeBaseModel.from_source(str(tmp_path)) assert model.config.rope_theta == 1000000.0 - assert model.config.rope_parameters == {"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} + 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. Publishing only rope_theta raised KeyError there and took the + offline EAGLE3 example tests down while every unit test stayed green. + """ + 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.""" From 9cf116cce565b26f2eb53893c86816f584051d3e Mon Sep 17 00:00:00 2001 From: h-guo18 <67671475+h-guo18@users.noreply.github.com> Date: Sun, 27 Sep 2026 08:26:18 +0000 Subject: [PATCH 10/11] fix(speculative): give the fake base's rope_parameters a rope_type The offline EAGLE3 example tests fail on this branch with modeling_eagle.py:84 LlamaRotaryEmbedding(config=self.config, ...) KeyError: 'rope_type' and take test_offline_resume_training_kimi down after them, since it reads the checkpoint the failed run never wrote. Both pass on main. Earlier in this PR `FakeBaseConfig` started publishing a `rope_parameters` dict so that consumers written for transformers 5 -- which read the dict before the flat field -- would find the target's RoPE base on a fake base too. That dict carried only `rope_theta`. The shape matters more than it looks, because this class is not only the base config: `hf_eagle.modify` builds the draft config as `type(base_config).from_dict(arch_config)`, so a FakeBaseConfig is also what the EAGLE draft is constructed from, and the dict lands where transformers' rotary embeddings index `rope_parameters["rope_type"]` unconditionally. A dict without it raises rather than falling back, and the manual publish pre-empts the well-formed dict transformers would otherwise build. `"default"` is both what every real config produces -- Qwen3Config and LlamaConfig each yield `{"rope_theta": ..., "rope_type": "default"}` -- and what the draft wants: the long-context scaling families are deliberately not inherited, they are re-applied at export. The new test builds a real `LlamaRotaryEmbedding` from a FakeBaseConfig and checks it produces usable tables, rather than asserting the dict's contents. Asserting contents is what the existing test did, and it stayed green through the regression: it pinned the value and not the requirement. Removing `rope_type` again fails three tests, including that one. Co-Authored-By: Claude Opus 5 (1M context) Signed-off-by: h-guo18 <67671475+h-guo18@users.noreply.github.com> --- .../speculative/plugins/modeling_fakebase.py | 20 +++++-------------- .../speculative/plugins/test_fakebase.py | 3 +-- 2 files changed, 6 insertions(+), 17 deletions(-) diff --git a/modelopt/torch/speculative/plugins/modeling_fakebase.py b/modelopt/torch/speculative/plugins/modeling_fakebase.py index 162033bc605..9ef6a1d2671 100644 --- a/modelopt/torch/speculative/plugins/modeling_fakebase.py +++ b/modelopt/torch/speculative/plugins/modeling_fakebase.py @@ -124,18 +124,11 @@ def __init__( ) self.intermediate_size = intermediate_size # For some drafter algo (e.g. DFlash) rope theta must match target model. Extract here. - # Published in both shapes: consumers built for Transformers 5 read the - # rope_parameters dict first, and a fake base that only carried the flat field - # would hand them nothing. + # 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: - # rope_type is not optional padding: this class is also the class the EAGLE - # draft config is built from (`type(base_config).from_dict(arch_config)` in - # hf_eagle.modify), and transformers' rotary embeddings index - # config.rope_parameters["rope_type"] unconditionally. A dict carrying only - # rope_theta raises KeyError there. "default" matches what every real config - # produces and what the draft wants -- the long-context scaling families are - # deliberately not inherited, they are re-applied at export. self.rope_parameters = {"rope_theta": rope_theta, "rope_type": "default"} if isinstance(dtype, str): dtype = getattr(torch, dtype) @@ -217,11 +210,8 @@ def from_source(cls, source: str, trust_remote_code: bool = False) -> "FakeBaseM 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), - # Shared with the exporter deliberately: where a config keeps rope_theta - # depends on the transformers version, and a local getattr got that wrong - # here for two months while the exporter had it right. The draft injects - # the target's KV, so a wrong base trains and exports without complaint - # and only misbehaves at serve time. + # 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/tests/unit/torch/speculative/plugins/test_fakebase.py b/tests/unit/torch/speculative/plugins/test_fakebase.py index 642af7c1a5b..be959b22e63 100644 --- a/tests/unit/torch/speculative/plugins/test_fakebase.py +++ b/tests/unit/torch/speculative/plugins/test_fakebase.py @@ -216,8 +216,7 @@ def test_config_drives_a_transformers_rotary_embedding(self): 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. Publishing only rope_theta raised KeyError there and took the - offline EAGLE3 example tests down while every unit test stayed green. + unconditionally. """ from transformers.models.llama.modeling_llama import LlamaRotaryEmbedding From aaccc4ba69accff9490223cf9f0710496020aa31 Mon Sep 17 00:00:00 2001 From: h-guo18 <67671475+h-guo18@users.noreply.github.com> Date: Sun, 27 Sep 2026 08:26:20 +0000 Subject: [PATCH 11/11] docs(speculative): cut the commentary back to what a reader needs The streaming example had drifted to 55% comments against 18-41% for its siblings, with a 52-line header. Most of the excess was a record of the run that validated it rather than documentation of the file: its measured step time and loss curve, and a catalogue of cluster-specific failure modes written while debugging them. Those belong in the PR and in the commit that made each fix, not in an example someone copies. The header keeps what it is, the node shape, how to scale it, and which settings are deliberately left unset; the capture-id and RDMA-sizing notes keep the constraint and drop the derivation. Same treatment for the fake base's RoPE comments, where two lines of code had carried fifteen of commentary, including how long the old bug had been there. No behaviour change: 305 unit tests pass, and removing `rope_type` still fails the three tests that guard it. Co-Authored-By: Claude Opus 5 (1M context) Signed-off-by: h-guo18 <67671475+h-guo18@users.noreply.github.com> --- .../Qwen/Qwen3-8B/hf_streaming_dflash2.yaml | 67 +++++-------------- 1 file changed, 15 insertions(+), 52 deletions(-) diff --git a/tools/launcher/examples/Qwen/Qwen3-8B/hf_streaming_dflash2.yaml b/tools/launcher/examples/Qwen/Qwen3-8B/hf_streaming_dflash2.yaml index ae334b188cf..89c06eeceed 100644 --- a/tools/launcher/examples/Qwen/Qwen3-8B/hf_streaming_dflash2.yaml +++ b/tools/launcher/examples/Qwen/Qwen3-8B/hf_streaming_dflash2.yaml @@ -7,45 +7,18 @@ # 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. -# -# On a cluster with fewer GPUs per node, drop gpus_per_node and SERVE_TP together (1 and 1 -# still works, it just leaves the trainer waiting on a single-GPU prefill) and lower -# per_device_train_batch_size to keep the global batch you want. To scale up instead, -# raise `nodes` and SERVE_NODES together so the trainer fans out over several serve -# replicas, as hf_streaming_dflash_multi_node.yaml does; keep per-serve in-flight -# (trainer_ranks x STREAMING_NUM_WORKERS / serve_replicas) roughly constant and -# HS_POOL_SLOTS a few times above it. +# 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 # -# Validated on a GB200 cluster (4 GPU/node, aarch64): 600 steps in 403 s (0.67 s/step, -# global batch 16 x 4096 tokens), loss 17.0 -> 3.4, exported 81 tensors. That run fed a -# pre-staged corpus instead of task_0 and used the overrides listed under "Cluster -# overrides" below. -# -# Cluster overrides this file deliberately does NOT hardcode, because the right value is -# site-specific: -# * container — vllm/vllm-openai:latest is x86; an aarch64 cluster needs an aarch64 -# image. The image must ship nixl (recent vllm-openai images do); if it does not, -# the serve dies at connector init with ModuleNotFoundError. -# * slurm_config.time — some clusters have a DefaultTime of minutes, and a job that -# outlives it dies as a bare Slurm TIMEOUT with no traceback. -# * fabric — on an InfiniBand cluster, pin BOTH stacks to the working HCAs or each -# picks a dead one and wedges silently (UCX as NIXL_ERR_BACKEND on the first fetch, -# NCCL as a watchdog SIGABRT half an hour later): -# NIXL_BACKENDS: "UCX" -# UCX_TLS: "rc,ud,sm,self" # host transports only -# UCX_NET_DEVICES: ":1,:1,..." -# NCCL_IB_DISABLE: "0" -# NCCL_IB_HCA: ",,..." -# On EFA use NIXL_BACKENDS=LIBFABRIC with FI_PROVIDER=efa instead. -# * caches — srun runs with --no-container-mount-home, so a $HOME cache path silently -# resolves inside the container overlay and is lost with the job. TRITON_CACHE_DIR -# must additionally be node-local: concurrent ranks race on a shared filesystem. +# 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 @@ -74,16 +47,10 @@ pipeline: # 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]. - # - # Unlike the DFlash example there is no final layer (36) in the capture list: the - # DFlash2 recipe trains against the hard target (dflash_self_logit_distillation: false, - # dflash_lk_loss_type: lambda), so no teacher distribution is formed and the base - # hidden is never read -- DFlashBaseModelOutput.from_offline_dict only touches - # base_model_hidden_states under need_logits. Dropping that plane therefore has to be - # paired with data.final_aux_is_base_hidden=true below, because the streaming dataset - # splits the captured planes unconditionally. To train with distillation instead, set - # dflash.dflash_lk_loss_type=ce, add 36 back, and drop final_aux_is_base_hidden. + # 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: @@ -101,11 +68,9 @@ pipeline: - training.output_dir=/scratchspace/dflash2 - training.training_seq_len=4096 - training.disable_tqdm=true - # The serve does not generate — it prefills the corpus conversation — so the - # assistant turn is present, but Qwen3-8B's stock chat template has no - # {% generation %} tags to locate it. Pass - # data.chat_template=examples/Qwen/Qwen3-8B/chat_template_train.jinja with - # answer_only_loss=true to mask to the assistant span instead. + # 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 @@ -126,10 +91,8 @@ pipeline: - 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, the producer silently skips capture, and the fetch hangs - # (the trainer has a fail-loud guard for it). 32 slots x 4608 tok x 5 planes x - # 4096 x 2 B = 5.4 GiB of pinned host memory; peak in-flight here is - # 4 ranks x 4 workers = 16. + # 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"