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>
244 lines
9.8 KiB
Python
244 lines
9.8 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.
|
|
|
|
"""Unit tests for FakeBaseModel and the fake-base / offline paths in load_vlm_or_llm."""
|
|
|
|
import json
|
|
|
|
import pytest
|
|
import safetensors.torch
|
|
import torch
|
|
|
|
pytest.importorskip("transformers")
|
|
import transformers
|
|
|
|
from modelopt.torch.speculative.plugins.modeling_fakebase import FakeBaseConfig, FakeBaseModel
|
|
from modelopt.torch.speculative.utils import load_vlm_or_llm
|
|
|
|
_HIDDEN_SIZE = 16
|
|
_VOCAB_SIZE = 32
|
|
|
|
|
|
@pytest.fixture
|
|
def fake_config(monkeypatch):
|
|
"""Monkeypatch AutoConfig.from_pretrained to return a minimal fake config."""
|
|
cfg = transformers.PretrainedConfig()
|
|
cfg.model_type = "llama"
|
|
cfg.hidden_size = _HIDDEN_SIZE
|
|
cfg.vocab_size = _VOCAB_SIZE
|
|
cfg.num_hidden_layers = 2
|
|
cfg.max_position_embeddings = 128
|
|
cfg.tie_word_embeddings = False
|
|
monkeypatch.setattr(transformers.AutoConfig, "from_pretrained", lambda *a, **kw: cfg)
|
|
return cfg
|
|
|
|
|
|
@pytest.fixture
|
|
def fake_checkpoint(tmp_path, fake_config):
|
|
"""Minimal local safetensors checkpoint loadable by FakeBaseModel."""
|
|
tensors = {
|
|
"lm_head.weight": torch.zeros(_VOCAB_SIZE, _HIDDEN_SIZE),
|
|
"embed_tokens.weight": torch.zeros(_VOCAB_SIZE, _HIDDEN_SIZE),
|
|
# model_type "llama" is in the final-norm whitelist, so FakeBaseModel requires the norm.
|
|
"norm.weight": torch.ones(_HIDDEN_SIZE),
|
|
}
|
|
shard = tmp_path / "model-00001-of-00001.safetensors"
|
|
safetensors.torch.save_file(tensors, shard)
|
|
index = {"weight_map": dict.fromkeys(tensors, shard.name)}
|
|
(tmp_path / "model.safetensors.index.json").write_text(json.dumps(index))
|
|
return tmp_path
|
|
|
|
|
|
def test_fakebase_local_happy_path(fake_checkpoint):
|
|
model = FakeBaseModel.from_source(str(fake_checkpoint))
|
|
assert model.lm_head.weight.shape == torch.Size([_VOCAB_SIZE, _HIDDEN_SIZE])
|
|
assert model.embed_tokens.weight.shape == torch.Size([_VOCAB_SIZE, _HIDDEN_SIZE])
|
|
|
|
|
|
def test_fakebase_missing_index_raises(tmp_path, fake_config):
|
|
with pytest.raises(FileNotFoundError, match="safetensors"):
|
|
FakeBaseModel.from_source(str(tmp_path))
|
|
|
|
|
|
def test_fakebase_single_file_no_index(tmp_path, fake_config):
|
|
"""Small models often ship a single ``model.safetensors`` without an index.json."""
|
|
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),
|
|
}
|
|
safetensors.torch.save_file(tensors, tmp_path / "model.safetensors")
|
|
model = FakeBaseModel.from_source(str(tmp_path))
|
|
assert model.lm_head.weight.shape == torch.Size([_VOCAB_SIZE, _HIDDEN_SIZE])
|
|
assert model.embed_tokens.weight.shape == torch.Size([_VOCAB_SIZE, _HIDDEN_SIZE])
|
|
|
|
|
|
def test_fakebase_tied_embeddings_falls_back_to_embed(tmp_path, fake_config):
|
|
"""Tied-embeddings models (e.g. Llama-3.2-1B) omit ``lm_head`` from safetensors;
|
|
FakeBaseModel must reuse ``embed_tokens`` for both."""
|
|
fake_config.tie_word_embeddings = True
|
|
weight = torch.randn(_VOCAB_SIZE, _HIDDEN_SIZE)
|
|
safetensors.torch.save_file(
|
|
{"embed_tokens.weight": weight, "norm.weight": torch.ones(_HIDDEN_SIZE)},
|
|
tmp_path / "model.safetensors",
|
|
)
|
|
model = FakeBaseModel.from_source(str(tmp_path))
|
|
torch.testing.assert_close(model.lm_head.weight, weight)
|
|
torch.testing.assert_close(model.embed_tokens.weight, weight)
|
|
|
|
|
|
def test_fakebase_missing_lm_head_without_tying_raises(tmp_path, fake_config):
|
|
"""Without ``tie_word_embeddings`` a missing ``lm_head`` is still an error."""
|
|
safetensors.torch.save_file(
|
|
{"embed_tokens.weight": torch.zeros(_VOCAB_SIZE, _HIDDEN_SIZE)},
|
|
tmp_path / "model.safetensors",
|
|
)
|
|
with pytest.raises(RuntimeError, match="lm_head"):
|
|
FakeBaseModel.from_source(str(tmp_path))
|
|
|
|
|
|
def test_load_vlm_or_llm_returns_fakebase(fake_checkpoint):
|
|
model = load_vlm_or_llm(str(fake_checkpoint), use_offline_training=True, use_fake_base=True)
|
|
assert isinstance(model, FakeBaseModel)
|
|
|
|
|
|
def test_load_vlm_or_llm_offline_zero_layers(monkeypatch):
|
|
cfg = transformers.PretrainedConfig()
|
|
cfg.model_type = "llama"
|
|
cfg.num_hidden_layers = 4
|
|
monkeypatch.setattr(transformers.AutoConfig, "from_pretrained", lambda *a, **kw: cfg)
|
|
|
|
captured_kwargs = {}
|
|
|
|
class _FakeModel:
|
|
config = cfg
|
|
|
|
def _fake_from_pretrained(*args, **kwargs):
|
|
captured_kwargs.update(kwargs)
|
|
return _FakeModel()
|
|
|
|
monkeypatch.setattr(transformers.AutoModelForCausalLM, "from_pretrained", _fake_from_pretrained)
|
|
|
|
model = load_vlm_or_llm("fake-model", use_offline_training=True, use_fake_base=False)
|
|
assert captured_kwargs.get("num_hidden_layers") == 0
|
|
assert model.config.num_orig_hidden_layers == 4
|
|
|
|
|
|
def test_load_vlm_or_llm_uses_transformers5_vlm_auto_class(monkeypatch):
|
|
"""Transformers 5 loads VLMs through AutoModelForImageTextToText."""
|
|
cfg = transformers.PretrainedConfig()
|
|
cfg.model_type = "qwen3_vl"
|
|
cfg.text_config = object()
|
|
monkeypatch.setattr(transformers.AutoConfig, "from_pretrained", lambda *a, **kw: cfg)
|
|
|
|
captured = {}
|
|
|
|
class _FakeVLM:
|
|
@staticmethod
|
|
def from_pretrained(*args, **kwargs):
|
|
captured["args"] = args
|
|
captured["kwargs"] = kwargs
|
|
return object()
|
|
|
|
# ``transformers`` exposes auto classes lazily, so deleting this attribute
|
|
# lets its module-level ``__getattr__`` recreate the legacy class. An
|
|
# explicit ``None`` models its absence and reliably exercises the v5
|
|
# fallback.
|
|
monkeypatch.setattr(transformers, "AutoModelForVision2Seq", None, raising=False)
|
|
monkeypatch.setattr(transformers, "AutoModelForImageTextToText", _FakeVLM)
|
|
|
|
assert load_vlm_or_llm("qwen3-vl", dtype="auto") is not None
|
|
assert captured["args"] == ("qwen3-vl",)
|
|
assert captured["kwargs"]["torch_dtype"] == "auto"
|
|
|
|
|
|
class TestFakeBaseRopeTheta:
|
|
"""The RoPE base has to survive the fake base, whichever shape the target stores it in.
|
|
|
|
The draft injects the target's KV, so a mismatched base trains and exports without
|
|
complaint and only misbehaves at serve time.
|
|
"""
|
|
|
|
def test_from_source_carries_a_transformers_5_base_theta(self, tmp_path, monkeypatch):
|
|
"""The seam, not the reader.
|
|
|
|
Reading rope_theta correctly is the exporter's ``_get_rope_theta`` and is tested
|
|
there. What is pinned here is that this call site uses it: a plain
|
|
``getattr(base_cfg, "rope_theta")`` reads None from a transformers 5 config and
|
|
silently builds a draft with no RoPE base at all.
|
|
"""
|
|
base_cfg = transformers.PretrainedConfig(
|
|
model_type="llama",
|
|
hidden_size=_HIDDEN_SIZE,
|
|
vocab_size=_VOCAB_SIZE,
|
|
num_hidden_layers=2,
|
|
max_position_embeddings=128,
|
|
tie_word_embeddings=False,
|
|
rope_parameters={"rope_theta": 1000000.0},
|
|
)
|
|
monkeypatch.setattr(transformers.AutoConfig, "from_pretrained", lambda *a, **kw: base_cfg)
|
|
tensors = {
|
|
"lm_head.weight": torch.zeros(_VOCAB_SIZE, _HIDDEN_SIZE),
|
|
"embed_tokens.weight": torch.zeros(_VOCAB_SIZE, _HIDDEN_SIZE),
|
|
"norm.weight": torch.ones(_HIDDEN_SIZE),
|
|
}
|
|
shard = tmp_path / "model-00001-of-00001.safetensors"
|
|
safetensors.torch.save_file(tensors, shard)
|
|
(tmp_path / "model.safetensors.index.json").write_text(
|
|
json.dumps({"weight_map": dict.fromkeys(tensors, shard.name)})
|
|
)
|
|
|
|
model = FakeBaseModel.from_source(str(tmp_path))
|
|
|
|
assert model.config.rope_theta == 1000000.0
|
|
assert model.config.rope_parameters == {"rope_theta": 1000000.0, "rope_type": "default"}
|
|
|
|
def test_config_publishes_both_shapes(self):
|
|
"""Consumers that prefer the dict must find it on a fake base too."""
|
|
config = FakeBaseConfig(num_hidden_layers=2, hidden_size=32, rope_theta=1000000.0)
|
|
assert config.rope_theta == 1000000.0
|
|
assert config.rope_parameters == {"rope_theta": 1000000.0, "rope_type": "default"}
|
|
|
|
def test_config_drives_a_transformers_rotary_embedding(self):
|
|
"""The published dict has to satisfy transformers, not just carry the number.
|
|
|
|
This class is also the class the EAGLE draft config is built from, so the dict
|
|
reaches `LlamaRotaryEmbedding`, which indexes ``rope_parameters["rope_type"]``
|
|
unconditionally.
|
|
"""
|
|
from transformers.models.llama.modeling_llama import LlamaRotaryEmbedding
|
|
|
|
config = FakeBaseConfig(
|
|
num_hidden_layers=2,
|
|
hidden_size=64,
|
|
num_attention_heads=4,
|
|
num_key_value_heads=2,
|
|
max_position_embeddings=128,
|
|
rope_theta=1000000.0,
|
|
)
|
|
config.head_dim = 16
|
|
|
|
rotary = LlamaRotaryEmbedding(config=config)
|
|
cos, sin = rotary(torch.zeros(1, 4, 64), torch.arange(4).unsqueeze(0))
|
|
|
|
assert rotary.rope_type == "default"
|
|
assert cos.shape == (1, 4, 16)
|
|
|
|
def test_unknown_theta_publishes_no_dict(self):
|
|
"""An absent base must stay absent rather than become a wrong default."""
|
|
config = FakeBaseConfig(num_hidden_layers=2, hidden_size=32, rope_theta=None)
|
|
assert config.rope_theta is None
|
|
assert not getattr(config, "rope_parameters", None)
|