Files
Model-Optimizer/tests/unit/torch/speculative/plugins/test_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

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)