Files
Model-Optimizer/modelopt/torch/speculative/plugins/modeling_fakebase.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

334 lines
15 KiB
Python

# 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.
"""Lightweight fake base model for offline speculative decoding training."""
import json
import os
import torch
import torch.nn as nn
import transformers
from huggingface_hub import hf_hub_download
from huggingface_hub.errors import EntryNotFoundError
from safetensors import safe_open
from transformers import (
AutoConfig,
AutoModel,
AutoModelForCausalLM,
PretrainedConfig,
PreTrainedModel,
)
from .modeling_final_norm import _FINAL_NORM_CLASSES, _select_final_norm_type
# Candidate module paths searched in order — shared with HFEagleModel._find_base_model_parts
_EMBED_TOKENS_PATHS = [
"embed_tokens",
"language_model.model.embed_tokens",
"model.embed_tokens",
"backbone.embeddings",
"language_model.backbone.embeddings",
"model.language_model.embed_tokens",
"tok_embeddings", # Mistral native checkpoints (consolidated.safetensors)
]
_LM_HEAD_PATHS = [
"lm_head",
"language_model.lm_head",
"output", # Mistral native checkpoints (consolidated.safetensors)
]
_FINAL_NORM_PATHS = [
"model.norm",
"language_model.model.norm",
"norm",
"backbone.norm_f",
"backbone.norm",
"language_model.backbone.norm",
"model.language_model.norm",
]
_BASE_MODEL_PATHS = [
"language_model.model",
"model.language_model",
"model",
"backbone",
"language_model.backbone",
]
_VLM_CONFIG_ATTRS = ["text_config", "llm_config"]
_SAFETENSORS_INDEX_FILENAME = "model.safetensors.index.json"
# Single-file safetensors names to try, in order. Mistral native checkpoints
# use ``consolidated.safetensors`` instead of the HF-standard ``model.safetensors``.
_SAFETENSORS_SINGLE_FILENAMES = ["model.safetensors", "consolidated.safetensors"]
class FakeBaseConfig(PretrainedConfig):
"""Minimal config for FakeBaseModel that supports offline speculative decoding training."""
model_type = "fake_base_model"
def __init__(
self,
num_hidden_layers=None,
hidden_size=None,
vocab_size=None,
max_position_embeddings=None,
dtype=torch.bfloat16,
tie_word_embeddings=False,
num_orig_hidden_layers=None,
num_attention_heads=None,
num_key_value_heads=None,
intermediate_size=None,
rms_norm_eps=1e-6,
rope_theta=None,
final_norm_type=None,
**kwargs,
):
"""Initialize FakeBaseConfig with minimal model configuration parameters."""
super().__init__(tie_word_embeddings=tie_word_embeddings, **kwargs)
self.rms_norm_eps = rms_norm_eps
# Which self-implemented final-norm class FakeBaseModel builds, or None to build no norm
# (model whose final-norm type we don't know). See _FINAL_NORM_CLASSES /
# _FINAL_NORM_TYPE_BY_MODEL_TYPE. Persisted so a reloaded config rebuilds the same norm.
self.final_norm_type = final_norm_type
self.num_hidden_layers = num_hidden_layers
# Mirror the original base layer count. The non-fake offline path loads with
# num_hidden_layers=0 and stashes the real count here (see utils.load_vlm_or_llm);
# the fake base keeps num_hidden_layers as the real count, so default to it. DFlash's
# offline modify() reads num_orig_hidden_layers directly (hf_dflash.py), so it must
# always be present on the base config.
# TODO: Deprecate the old offline path.
self.num_orig_hidden_layers = (
num_orig_hidden_layers if num_orig_hidden_layers is not None else num_hidden_layers
)
self.hidden_size = hidden_size
self.vocab_size = vocab_size
self.max_position_embeddings = max_position_embeddings
# Attention/MLP dims needed when exporting a draft head built on a fake base: the
# DFlash exporter (hf_spec_export._export_config) references base_config.{num_attention_heads,
# num_key_value_heads, intermediate_size} as getattr fallbacks, which Python evaluates
# eagerly, so they must exist even though the fake base has no real layers.
self.num_attention_heads = num_attention_heads
self.num_key_value_heads = (
num_key_value_heads if num_key_value_heads is not None else num_attention_heads
)
self.intermediate_size = intermediate_size
# For some drafter algo (e.g. DFlash) rope theta must match target model. Extract here.
# Published in both shapes: Transformers 5 consumers read the dict first, and it has
# to be well-formed -- the EAGLE draft config is built from this class too, and
# rotary embeddings index rope_parameters["rope_type"] unconditionally.
self.rope_theta = rope_theta
if rope_theta is not None:
self.rope_parameters = {"rope_theta": rope_theta, "rope_type": "default"}
if isinstance(dtype, str):
dtype = getattr(torch, dtype)
self.dtype = dtype
class FakeBaseModel(PreTrainedModel):
"""Minimal base model for offline speculative decoding.
Contains only ``lm_head``, ``embed_tokens``, and the minimal config needed by the EAGLE
training loop. The full model weights are never loaded, keeping memory usage low.
Weights are loaded from a local HuggingFace checkpoint directory. Weight key names and
VLM config nesting are auto-detected from the shared path constants.
"""
config_class = FakeBaseConfig
def __init__(self, config: FakeBaseConfig, **kwargs):
"""Initialize FakeBaseModel structure from a FakeBaseConfig.
To construct a FakeBaseModel from an original HuggingFace checkpoint (e.g. a Llama
repo), use the :meth:`from_source` classmethod instead.
"""
super().__init__(config, **kwargs)
# Initialize dummy module and attributes for compatibility with HFEagleModel
self.model = nn.Module()
self.model.layers = nn.ModuleList()
self.model.dtype = config.dtype
self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, dtype=config.dtype)
self.lm_head = nn.Linear(
config.hidden_size, config.vocab_size, bias=False, dtype=config.dtype
)
# Final pre-lm_head norm, applied before lm_head when reconstructing base logits for
# self-logit-distillation (vLLM-captured final hidden states are un-normed). Built ONLY
# when the base model's final-norm type is known (config.final_norm_type set); otherwise
# no ``norm`` attribute exists and downstream skips re-norming. The concrete class comes
# from config.final_norm_type (see _FINAL_NORM_CLASSES); weight loaded in from_source.
if config.final_norm_type is not None:
norm_cls = _FINAL_NORM_CLASSES[config.final_norm_type]
self.norm = norm_cls(config.hidden_size, eps=config.rms_norm_eps, dtype=config.dtype)
# Initialize weights and apply final processing
self.post_init()
@classmethod
def from_source(cls, source: str, trust_remote_code: bool = False) -> "FakeBaseModel":
"""Load lm_head and embed_tokens from a local directory or HuggingFace Hub repo.
Args:
source: Path to a local HuggingFace checkpoint directory, or a HuggingFace Hub
repo ID (e.g. ``"meta-llama/Llama-3.1-8B"``). The source type is detected
automatically: if ``source`` is an existing local directory it is treated as a
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
)
# For vlms, detect language model config based on _VLM_CONFIG_ATTRS
base_cfg = next(
(
getattr(orig_config, attr)
for attr in _VLM_CONFIG_ATTRS
if getattr(orig_config, attr, None) is not None
),
orig_config,
)
# Extract necessary info for spec training from base config
config = FakeBaseConfig(
num_hidden_layers=getattr(base_cfg, "num_hidden_layers", None),
hidden_size=getattr(base_cfg, "hidden_size", None),
vocab_size=getattr(base_cfg, "vocab_size", None),
max_position_embeddings=getattr(base_cfg, "max_position_embeddings", None),
dtype=getattr(base_cfg, "dtype", torch.bfloat16),
tie_word_embeddings=getattr(base_cfg, "tie_word_embeddings", False),
num_attention_heads=getattr(base_cfg, "num_attention_heads", None),
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: 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
),
)
model = cls(config)
# Load lm_head, embed_tokens, and (for known models) the final norm into the model.
model._load_weights(source)
return model
@staticmethod
def _find_weight_key(weight_map: dict, paths: list[str], label: str) -> str:
"""Return the first ``path + '.weight'`` found in ``weight_map``."""
for path in paths:
key = path + ".weight"
if key in weight_map:
return key
tried = [p + ".weight" for p in paths]
raise RuntimeError(f"Cannot find {label} in checkpoint; tried: {tried}")
@staticmethod
def _load_index(source: str) -> dict:
"""Load weight_map from a sharded index, or synthesize one from a single safetensors file.
Sharded checkpoints ship ``model.safetensors.index.json`` mapping every key to its shard;
small checkpoints ship a single ``model.safetensors`` with no index — we read its keys
and synthesize the equivalent weight_map so downstream code stays the same.
"""
def _try_fetch(name: str) -> str | None:
if os.path.isdir(source):
path = os.path.join(source, name)
return path if os.path.isfile(path) else None
try:
return hf_hub_download(repo_id=source, filename=name)
except EntryNotFoundError:
return None
if (index_path := _try_fetch(_SAFETENSORS_INDEX_FILENAME)) is not None:
with open(index_path) as f:
return json.load(f).get("weight_map", {})
for single_name in _SAFETENSORS_SINGLE_FILENAMES:
if (single_path := _try_fetch(single_name)) is not None:
with safe_open(single_path, framework="pt") as h:
return dict.fromkeys(h.keys(), single_name)
raise FileNotFoundError(
f"No {_SAFETENSORS_INDEX_FILENAME} or {_SAFETENSORS_SINGLE_FILENAMES} found at "
f"{source!r}. FakeBaseModel only supports safetensors checkpoints; "
"pytorch_model.bin is not supported."
)
@staticmethod
def _resolve_shard_paths(source: str, shard_filenames: list[str]) -> list[str]:
"""Return local filesystem paths for each shard filename.
For a local directory the paths are joined directly; for a HuggingFace Hub repo ID the
shards are downloaded via ``hf_hub_download`` (cached on subsequent calls).
"""
if os.path.isdir(source):
return [os.path.join(source, name) for name in shard_filenames]
return [hf_hub_download(repo_id=source, filename=name) for name in shard_filenames]
def _load_weights(self, source: str) -> None:
"""Load lm_head, embed_tokens, and (for known models) the final norm into this model.
Reads only the tensors needed (never materializes a whole shard) and copies them into
the already-constructed submodules. For unknown models (``final_norm_type`` unset) the
norm is neither loaded nor present (``self.norm`` does not exist; see :meth:`__init__`).
"""
weight_map = self._load_index(source)
embed_tokens_key = self._find_weight_key(weight_map, _EMBED_TOKENS_PATHS, "embed_tokens")
try:
lm_head_key = self._find_weight_key(weight_map, _LM_HEAD_PATHS, "lm_head")
except RuntimeError:
# Tied embeddings: lm_head shares embed_tokens weight and isn't stored separately.
if not self.config.tie_word_embeddings:
raise
lm_head_key = embed_tokens_key
# Pull only the tensor we need; avoids materializing the whole file.
def _read(key: str) -> torch.Tensor:
(path,) = self._resolve_shard_paths(source, [weight_map[key]])
with safe_open(path, framework="pt", device="cpu") as h:
return h.get_tensor(key)
# Explicit shape checks: copy_ would broadcast a wrong-but-compatible shape silently.
# Use raises (not asserts) so the guard survives python -O / PYTHONOPTIMIZE.
hidden, vocab = self.config.hidden_size, self.config.vocab_size
embed_tokens_w = _read(embed_tokens_key)
lm_head_w = _read(lm_head_key)
if embed_tokens_w.shape != (vocab, hidden):
raise ValueError(
f"embed_tokens weight shape {tuple(embed_tokens_w.shape)} != ({vocab}, {hidden})"
)
if lm_head_w.shape != (vocab, hidden):
raise ValueError(
f"lm_head weight shape {tuple(lm_head_w.shape)} != ({vocab}, {hidden})"
)
self.embed_tokens.weight.data.copy_(embed_tokens_w)
self.lm_head.weight.data.copy_(lm_head_w)
# Final norm only for models whose norm type we know (final_norm_type set); when known it
# MUST be present — a missing key is a hard error, never a silent skip. Unknown: no norm.
if self.config.final_norm_type is not None:
norm_w = _read(self._find_weight_key(weight_map, _FINAL_NORM_PATHS, "final_norm"))
if norm_w.shape != (hidden,):
raise ValueError(f"final-norm weight shape {tuple(norm_w.shape)} != ({hidden},)")
self.norm.weight.data.copy_(norm_w)
def forward(self, *args, **kwargs):
"""Not implemented: FakeBaseModel omits full model weights and cannot run inference."""
raise NotImplementedError("FakeBaseModel forward is not implemented.")
# Register so that AutoConfig / AutoModel / AutoModelForCausalLM can resolve "fake_base_model".
AutoConfig.register("fake_base_model", FakeBaseConfig)
AutoModel.register(FakeBaseConfig, FakeBaseModel)
AutoModelForCausalLM.register(FakeBaseConfig, FakeBaseModel)