Files
Model-Optimizer/modelopt/torch/speculative/plugins/modeling_dflash2.py
T
h-guo18andClaude Opus 5 c2aaa44f60 [Speculative Decoding] DFlash2 draft variant (grouped sublayer convolution + candidate selector) (#2216)
### 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>
2026-09-29 11:43:02 +08:00

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)