mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
### What does this PR do? Type of change: new feature Adds **DFlash2** ([blog](https://inco.ai/blog/dflash2/)) as a draft variant of the existing DFlash mode, selected with `dflash_architecture_config.projector_type="dflash2"` alongside `domino`, `dspark` and `lilicorr`. DFlash2 keeps DFlash's one-pass parallel backbone and adds two components that recover the acceptance a purely parallel draft loses: - **Grouped dynamic depthwise convolution** around every attention and MLP sublayer, giving each block position a view of its predecessors *inside* the block. Taps do not cross the block boundary, so the draft stays one forward pass. - **Low-rank candidate selector** scoring transitions between adjacent positions' top-k candidates, so serving walks one coherent path instead of taking an independent argmax per position. Both start as exact no-ops — the convolution's `base_kernel` is an identity and `kernel_projection` is zeroed; the selector's `successor_codebook` is zeroed — so a freshly built DFlash2 draft *is* its DFlash backbone, and enabling the variant is an extension rather than a perturbation. This matches the reference implementation ([SpecForge#772](https://github.com/sgl-project/SpecForge/pull/772), merged) and the way `modeling_lilicorr` already installs this same convolution class. **This unblocks a recipe already shipped on `main`.** `modeling_lilicorr._install_sublayer_convs` imports `DFlashGroupedConv` from `modeling_dflash2`, so `modelopt_recipes/general/speculative_decoding/lilicorr_conv.yaml` raises at model build today and its CHANGELOG entry documents a feature that cannot run. Landing this makes it runnable. Module and parameter names match the SGLang/vLLM `DFlash2DraftModel` loaders. Verified against the released `z-lab/Qwen3.8-27B-DFlash2` checkpoint: **81 tensors, 21 name patterns, zero difference in either direction**. The serving side, [vllm-project/vllm#52816](https://github.com/vllm-project/vllm/pull/52816), has since merged (`b389ac294`) with no change to the checkpoint contract. ### Usage ```bash python examples/speculative_decoding/main.py \ --config modelopt_recipes/general/speculative_decoding/dflash2.yaml \ model.model_name_or_path=Qwen/Qwen3-8B \ data.data_path=<corpus>.jsonl \ training.output_dir=<out> ``` ```yaml # modelopt_recipes/general/speculative_decoding/dflash2.yaml dflash: dflash_selector_loss_alpha: 1.0 # weight of the candidate-selector CE term dflash_architecture_config: projector_type: dflash2 conv_kernel_size: 2 # taps; must not exceed the block size conv_group_size: 16 # must divide hidden_size selector_rank: 256 selector_top_k: 16 ``` ### Testing <img width="2000" height="1320" alt="image" src="https://github.com/user-attachments/assets/869004d1-92c6-41a4-a03d-fd824a06255c" /> **Unit** — 25 CPU tests in `tests/unit/torch/speculative/plugins/test_hf_dflash2.py`; the full `tests/unit/torch/speculative/` suite passes with no regressions. The ones worth keeping pin invariants that a decreasing loss does not catch: - the convolution is an exact identity on the **default** construction, and its taps stay inside the block while a position still sees its predecessors; - the block-offset contract shared by the training objective and `CandidateSelector.greedy_path` — a misaligned objective still converges; - the export fields the vLLM loader requires, including the top-level `block_size` that `DFlash2Exporter` derives the nested copy from; - which selector factors receive gradient on the first step. `successor_codebook` starts at zero, so `predecessor_codebook` and `hidden_projection` take one step to begin moving. That is a warm start, not a dead branch, and both sides are asserted. **End-to-end** — trained on Qwen3-8B against a plain DFlash control with every other argument identical (plot above). Monotonic convergence, no NaN/divergence, no DDP unused-parameter issues. Note the losses are **not comparable** across arms: DFlash2's includes the selector CE term. **Serving (vLLM)** — the exported drafter loads and drafts under the merged DFlash2 path (`RESOLVED draft architectures: ['DFlash2DraftModel']`). Two notes for anyone reproducing: vLLM sizes the convolution from `1 + num_speculative_tokens` at runtime rather than from the checkpoint, so a `block_size=16` drafter is only correct at `num_speculative_tokens=15`; and at that value the upstream path currently hits an illegal memory access in `_cache_draft_logits` ([vllm#55279](https://github.com/vllm-project/vllm/issues/55279)), independent of which checkpoint is used. ### Before your PR is "*Ready for review*" - Is this change backward compatible?: ✅ — additive. New `projector_type`, its own registry and exporter, one new config field; DFlash / Domino / DSpark / LiLiCorr numerics and `state_dict` contents are untouched. - If you copied code from any other sources or added a new PIP dependency, did you follow guidance in `CONTRIBUTING.md`: ✅ — `modeling_dflash2.py` is adapted from [SpecForge#772](https://github.com/sgl-project/SpecForge/pull/772) and carries its MIT notice. No new dependencies. - Did you write any new necessary tests?: ✅ - Did you update [Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?: ✅ — under `0.48.0`. - Did you get Claude approval on this PR?: ✅ — run on 2026-08-20; all review threads addressed and resolved. ### Additional Information Rebased onto current `main`. Two commits from the original branch were dropped because [#2342](https://github.com/NVIDIA/Model-Optimizer/pull/2342) landed them first, with authorship preserved: the no-op sublayer seam in `modeling_dflash.py`, and the `rope_theta`/`rope_parameters` fix — `main`'s version of the latter is stricter, so this PR no longer touches `hf_dflash.py` at all. <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit * **New Features** * Added DFlash2 speculative decoding with grouped dynamic convolutions and low-rank candidate selection. * Added configurable selector-loss weighting, including an option to disable it. * Added DFlash2 model conversion and export support. * Added checkpoints compatible with SGLang and vLLM DFlash2 serving. * **Documentation** * Added training recipes and a Qwen3-8B online DFlash2 training configuration. * **Tests** * Added coverage for conversion, training, metrics, gradients, and export compatibility. <!-- end of auto-generated comment: release notes by coderabbit.ai --> --------- Signed-off-by: h-guo18 <67671475+h-guo18@users.noreply.github.com> Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
318 lines
15 KiB
Python
318 lines
15 KiB
Python
# 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) = <predecessor_codebook[p] * hidden_projection(h),
|
|
successor_codebook[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)
|