[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>
This commit is contained in:
h-guo18
2026-09-29 11:43:02 +08:00
committed by GitHub
co-authored by Claude Opus 5
parent be74001256
commit c2aaa44f60
19 changed files with 1922 additions and 19 deletions
+2
View File
@@ -117,6 +117,8 @@ repos:
modelopt/torch/speculative/plugins/modeling_domino.py|
modelopt/torch/speculative/plugins/hf_dflash.py|
modelopt/torch/speculative/plugins/modeling_dflash.py|
modelopt/torch/speculative/plugins/hf_dflash2.py|
modelopt/torch/speculative/plugins/modeling_dflash2.py|
modelopt/torch/speculative/plugins/hf_dspark.py|
modelopt/torch/speculative/plugins/modeling_dspark.py|
modelopt/torch/speculative/plugins/hf_lilicorr.py|
+4
View File
@@ -26,6 +26,10 @@ Changelog
- Add an end-to-end BEVFormer ONNX PTQ example with temporal calibration data generation, INT8 and FP8 quantization, TensorRT engine building, and nuScenes accuracy evaluation. See `examples/onnx_ptq/bevformer/README.md <https://github.com/NVIDIA/Model-Optimizer/tree/main/examples/onnx_ptq/bevformer>`_ for details.
- Add a reusable local-Hessian NVFP4 PTQ recipe and the quantization recipe used for ``nvidia/Qwen3.8-27B-NVFP4``.
*Speculative Decoding*
- Add the DFlash2 draft variant, selected with ``dflash_architecture_config.projector_type="dflash2"``: DFlash's one-pass parallel backbone plus a grouped dynamic convolution around every attention/MLP sublayer (``conv_kernel_size`` / ``conv_group_size``) and a low-rank candidate selector (``selector_rank`` / ``selector_top_k``, weighted by ``dflash_selector_loss_alpha``). Exported checkpoints declare ``DFlash2DraftModel`` and load in the SGLang/vLLM DFlash2 serving path.
*Megatron Framework (M-LM / M-Bridge)*
- Add an end-to-end W4A4 NVFP4 PTQ and QAD tutorial for Qwen3.6-35B-A3B also covering evaluation and vLLM throughput benchmarking. See `examples/megatron_bridge/tutorials/Qwen3.6-35B-A3B/README.md <https://github.com/NVIDIA/Model-Optimizer/tree/main/examples/megatron_bridge/tutorials/Qwen3.6-35B-A3B/>`_ for details.
@@ -605,3 +605,46 @@ class DSparkExporter(DFlashExporter):
}
)
return config
class DFlash2Exporter(DFlashExporter):
"""Draft model exporter for DFlash2 (DFlash backbone + convolutions + selector).
Same z-lab-compatible format as DFlash, plus the DFlash2 weights
(``layers.*.attention_conv.*`` / ``layers.*.mlp_conv.*`` /
``candidate_selector.*``, already captured by the inherited ``dflash_module.``
stripping) and the config fields the SGLang/vLLM ``DFlash2DraftModel`` loader
needs to rebuild them (``conv_kernel_size``, ``conv_group_size``,
``selector_rank``, ``selector_top_k``).
The architecture name is what selects the DFlash2 serving path: a checkpoint
declaring ``DFlashDraftModel`` loads as a plain DFlash draft and would silently
ignore the convolutions and the selector.
"""
def _export_config(self):
"""Extend the DFlash config with the DFlash2 architecture fields."""
config = super()._export_config()
draft_config = self.model.dflash_config
config["architectures"] = ["DFlash2DraftModel"]
# Present because HFDFlash2Model.modify validates them at convert time.
config["dflash_config"].update(
{
"projector_type": getattr(draft_config, "projector_type", "dflash2"),
"conv_kernel_size": draft_config.conv_kernel_size,
"conv_group_size": draft_config.conv_group_size,
"selector_rank": draft_config.selector_rank,
"selector_top_k": draft_config.selector_top_k,
# The published DFlash2 checkpoints carry block_size inside
# dflash_config; the DFlash loader reads it from the top level.
# Emit both so either contract resolves to the same value.
"block_size": config["block_size"],
}
)
# vLLM reads is_causal from the TOP level, and the published DFlash2 checkpoints
# state it explicitly rather than leaving it to be inferred from layer_types.
# Mirror dflash_config.causal, which DFlashExporter writes unconditionally from
# dflash_draft_attention; no parent writes a top-level is_causal.
config["is_causal"] = config["dflash_config"]["causal"]
return config
+46
View File
@@ -375,6 +375,52 @@ class DFlashConfig(ModeloptBaseConfig):
),
)
dflash_lk_loss_type: Literal["ce", "tv", "lambda"] = ModeloptField(
default="ce",
description=(
"DFlash2 only: which divergence the block objective minimizes against the "
"hard target. 'ce' is -log q(gold), today's behavior. 'tv' is 1 - q(gold), the "
"total variation to the one-hot target, which is also the per-position expected "
"acceptance loss. 'lambda' anneals between them: the CE share is "
"dflash_lk_ce_scale * exp(-dflash_lk_ce_decay * a), where a is the mean q(gold) "
"over supervised positions, so the objective moves from fitting the "
"distribution to maximizing acceptance as acceptance improves. "
"'lambda' and 'tv' require dflash_self_logit_distillation=false: both read "
"q(gold) from the per-position cross-entropy, which the KD path does not "
"produce. Ignored unless dflash_architecture_config.projector_type == 'dflash2'."
),
)
dflash_lk_ce_scale: float = ModeloptField(
default=1.0,
ge=0.0,
description=(
"DFlash2 only: scale of the CE share in the dflash_lk_loss_type='lambda' "
"blend. 1.0 starts the run as pure CE. Ignored for other loss types."
),
)
dflash_lk_ce_decay: float = ModeloptField(
default=1.0,
ge=0.0,
description=(
"DFlash2 only: how fast the CE share decays as acceptance rises in the "
"dflash_lk_loss_type='lambda' blend. 0 pins the blend at dflash_lk_ce_scale. "
"Ignored for other loss types."
),
)
dflash_selector_loss_alpha: float = ModeloptField(
default=1.0,
ge=0.0,
description=(
"DFlash2 only: weight of the candidate-selector cross-entropy term, added to "
"the backbone loss. The selector re-ranks the backbone's top-k candidates per "
"block position; 0 trains the backbone and convolutions only. "
"Ignored unless dflash_architecture_config.projector_type == 'dflash2'."
),
)
@model_validator(mode="after")
def _check_dpace_alpha(self) -> "DFlashConfig":
# Validate at construction regardless of the active objective, so a bad alpha
@@ -35,6 +35,12 @@ DominoDMRegistry = _DMRegistryCls(prefix="Domino")
# ``dflash_architecture_config.projector_type == "dspark"`` and kept in its own
# registry so its wrapper (HFDSparkModel) does not overwrite HFDFlashModel.
DSparkDMRegistry = _DMRegistryCls(prefix="DSpark")
# DFlash2 also reuses the dflash mode/config/recipe, converting the base model to a
# DFlash backbone whose sublayers are wrapped in grouped dynamic convolutions, plus a
# low-rank candidate selector. Selected via
# ``dflash_architecture_config.projector_type == "dflash2"`` and kept in its own
# registry so its wrapper (HFDFlash2Model) does not overwrite HFDFlashModel.
DFlash2DMRegistry = _DMRegistryCls(prefix="DFlash2")
# LiLiCorr also reuses the dflash mode/config/recipe, converting the base model to a
# DFlash backbone augmented with a reranker over the candidate lattice the backbone
# already produces. Selected via
@@ -59,6 +65,8 @@ def convert_to_dflash_model(model: nn.Module, config: DFlashConfig) -> ConvertRe
registry = DominoDMRegistry
elif projector_type == "dspark":
registry = DSparkDMRegistry
elif projector_type == "dflash2":
registry = DFlash2DMRegistry
elif projector_type == "lilicorr":
registry = LiLiCorrDMRegistry
elif projector_type in (None, "dflash"):
@@ -66,7 +74,7 @@ def convert_to_dflash_model(model: nn.Module, config: DFlashConfig) -> ConvertRe
else:
raise ValueError(
f"Unsupported dflash_architecture_config.projector_type: {projector_type!r}. "
"Expected 'dflash' (default), 'domino', 'dspark' or 'lilicorr'."
"Expected 'dflash' (default), 'domino', 'dspark', 'dflash2' or 'lilicorr'."
)
original_cls = type(model)
@@ -31,6 +31,7 @@ with import_plugin("megatron_medusa"):
with import_plugin("transformers"):
from .hf_dflash import *
from .hf_dflash2 import *
from .hf_domino import *
from .hf_dspark import *
from .hf_eagle import *
@@ -73,7 +73,7 @@ Draft model components:
import logging
from pathlib import Path
from typing import Any
from typing import Any, NamedTuple
import torch
import torch.nn.functional as F
@@ -200,6 +200,21 @@ def _dpace_position_weights(
return weights.to(dtype=confidences.dtype)
class _DFlashLossTerms(NamedTuple):
"""The unreduced pieces behind the block loss, for variants that re-derive it.
``ce_per_token`` is ``None`` on the KD path, which never forms a per-position
cross-entropy. ``weights`` carries the position weighting (decay or D-PACE) and
normalizes by ``weight_sum``; ``supervised_mask`` is the unweighted mask the
reported accuracy uses.
"""
ce_per_token: torch.Tensor | None
weights: torch.Tensor
weight_sum: torch.Tensor
supervised_mask: torch.Tensor
@DFlashDMRegistry.register({PreTrainedModel: "hf.PreTrainedModel"})
class HFDFlashModel(DFlashModel):
"""DFlash Model for HuggingFace transformers."""
@@ -823,6 +838,7 @@ class HFDFlashModel(DFlashModel):
base_logits=None,
draft_hidden=None,
base_outputs=None,
return_terms=False,
):
"""Compute weighted cross-entropy (or KD) loss and accuracy.
@@ -835,6 +851,11 @@ class HFDFlashModel(DFlashModel):
base_logits: Base model logits for KD loss [B, seq_len, vocab], or None for CE.
draft_hidden: Draft hidden states [B, N*block_size, H] behind ``logits``.
Unused here; passed for variants whose head consumes them.
return_terms: Also return the unreduced pieces behind the loss, so a variant
can recompose the block objective from a different divergence without
rebuilding the target alignment and position weighting.
TODO: promote this into a shared divergence seam when the DFlash-family
loss code is refactored; DFlash2 is the only consumer today.
Returns:
(loss, accuracy) tuple.
@@ -929,6 +950,14 @@ class HFDFlashModel(DFlashModel):
loss = flat_logits.sum() * 0.0
accuracy = 0.0
if return_terms:
terms = _DFlashLossTerms(
ce_per_token=loss_per_token,
weights=flat_weights,
weight_sum=valid_count,
supervised_mask=binary_eval_mask,
)
return loss, accuracy, terms
return loss, accuracy
def forward(
@@ -0,0 +1,279 @@
# 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.
"""HF DFlash2 model wrapper — DFlash training plus the candidate-selector objective.
DFlash2 differs from DFlash only in the draft module (grouped dynamic convolutions
around every sublayer, plus a candidate selector) and in one extra loss term, so
this wrapper reuses ``HFDFlashModel``'s forward wholesale and overrides just
:meth:`_compute_loss`.
The convolutions need no supervision of their own: they sit inside the backbone
and are trained by the backbone loss. The selector does, because at serving time
it — not an independent argmax — picks the drafted token at each block position.
Selector supervision (following the SGLang/SpecForge reference):
- Take the backbone's top-k candidates per block position.
- Score each candidate against its *teacher-forced* predecessor token, so the
positions train in parallel exactly as the backbone does.
- When the gold token is missing from the top-k, substitute it into the last
candidate slot. Without this the selector sees no positive class on the hard
positions and never learns those edges.
"""
import torch
import torch.nn.functional as F
from transformers import PreTrainedModel
from ..dflash.conversion import DFlash2DMRegistry
from .hf_dflash import HFDFlashModel
from .modeling_dflash2 import DFlash2Module
__all__ = ["HFDFlash2Model"]
@DFlash2DMRegistry.register({PreTrainedModel: "hf.PreTrainedModel"})
class HFDFlash2Model(HFDFlashModel):
"""DFlash model with DFlash2's sublayer convolutions and candidate selector.
Registered in ``DFlash2DMRegistry`` so that ``convert_to_dflash_model`` routes
to it when ``dflash_architecture_config.projector_type == "dflash2"``.
"""
def _build_draft_module(self, dflash_config):
"""Build the DFlash2 draft module (DFlash backbone + convolutions + selector)."""
return DFlash2Module(dflash_config)
def modify(self, config):
"""Initialize the DFlash2 draft module and read the selector loss weight."""
arch_config = config.dflash_architecture_config
missing = [
name
for name in ("conv_kernel_size", "conv_group_size", "selector_rank", "selector_top_k")
if arch_config.get(name) is None
]
if missing:
raise ValueError(
f"DFlash2 (projector_type='dflash2') requires {missing} in "
"dflash_architecture_config (convolution taps/group size and the "
"candidate selector's rank/top-k)."
)
super().modify(config)
self.dflash_selector_loss_alpha = getattr(config, "dflash_selector_loss_alpha", 1.0)
self.dflash_lk_loss_type = getattr(config, "dflash_lk_loss_type", "ce")
self.dflash_lk_ce_scale = getattr(config, "dflash_lk_ce_scale", 1.0)
self.dflash_lk_ce_decay = getattr(config, "dflash_lk_ce_decay", 1.0)
if self.dflash_lk_loss_type != "ce" and self.dflash_self_logit_distillation:
raise ValueError(
f"dflash_lk_loss_type={self.dflash_lk_loss_type!r} needs the draft's "
"probability of the gold token, which the KD path never forms -- it "
"optimizes a soft target instead. Set dflash_self_logit_distillation=false, "
"or dflash_lk_loss_type='ce' to keep distillation."
)
self._selector_metrics = None
def forward(self, *args, **kwargs):
"""Run the DFlash training forward and attach the candidate-selector metrics.
The variant is the objective, so the pipeline above it is inherited verbatim;
this only carries out what ``_compute_loss`` produced. Without it
``selector_coverage`` -- the only signal distinguishing a selector that is
choosing from one being handed the gold token -- never reaches the logs.
"""
self._selector_metrics = None
outputs = super().forward(*args, **kwargs)
if self._selector_metrics is not None:
outputs["selector_metrics"] = self._selector_metrics
self._selector_metrics = None
return outputs
def get_exporter(self):
"""Get the exporter for the DFlash2 draft model."""
from modelopt.torch.export.plugins.hf_spec_export import DFlash2Exporter
return DFlash2Exporter(self)
def _selector_loss(self, logits, target_ids, hidden, predecessor_ids, weight_mask):
"""Cross-entropy over the selector's candidate set, and its top-1 accuracy.
Args:
logits: Backbone logits per block position ``[B, N, block_size, V]``.
target_ids: Gold token ids ``[B, N, block_size]``.
hidden: Backbone hidden states ``[B, N, block_size, H]``.
predecessor_ids: Teacher-forced predecessor ids ``[B, N, block_size]``.
weight_mask: Per-position loss weights ``[B, N, block_size]``.
Returns:
``(loss, accuracy, coverage)`` — coverage is the fraction of supervised
positions whose gold token was already in the backbone's top-k, i.e. how
often the selector is choosing rather than being handed the answer.
"""
selector = self.dflash_module.candidate_selector
top_k = selector.top_k
unary_logits, candidate_ids = logits.topk(top_k, dim=-1)
# Train on the strict top-k, the candidate set serving actually builds. A gold
# token the backbone did not propose is a backbone recall failure, not a
# selector classification example: substituting it in would teach the selector
# to override the unary ranking on a set it will never be shown. Those
# positions carry no selector gradient and leave the denominator instead.
gold_matches = candidate_ids == target_ids.unsqueeze(-1)
gold_in_topk = gold_matches.any(dim=-1)
gold_slot = gold_matches.long().argmax(dim=-1)
selector_logits = selector.score_candidates(
candidate_ids, unary_logits, hidden, predecessor_ids
)
covered = weight_mask * gold_in_topk.to(weight_mask.dtype)
flat_weights = covered.reshape(-1)
denominator = flat_weights.sum() + 1e-6
per_token = F.cross_entropy(
selector_logits.float().reshape(-1, top_k),
gold_slot.reshape(-1),
reduction="none",
)
loss = (per_token * flat_weights).sum() / denominator
with torch.no_grad():
chosen = selector_logits.argmax(dim=-1).reshape(-1)
accuracy = (
(chosen == gold_slot.reshape(-1)).float() * flat_weights
).sum() / denominator
# Coverage keeps the full supervised mask as its denominator: it measures
# how often the selector was given a solvable problem at all.
# clamp, not an epsilon: a fully covered batch must read exactly 1.0.
supervised = weight_mask.reshape(-1).sum()
coverage = flat_weights.sum() / supervised.clamp(min=1.0)
# Detached tensors, not Python scalars: .item() would force a CPU-GPU sync on
# every training step. The trainer converts them at the logging boundary.
return loss, accuracy.detach(), coverage.detach()
def _lk_loss(self, terms):
"""Re-weight the block objective between cross-entropy and acceptance.
Both terms are read off the same per-position cross-entropy the backbone loss
already produced, so the target alignment and position weighting are shared::
q = exp(-ce) draft probability of the gold token
L_ce = <ce>_w today's objective
L_tv = <1 - q>_w total variation to the one-hot target,
i.e. the per-position acceptance loss
a = <q>_mask mean acceptance over supervised positions
L = s*exp(-d*a) * L_ce + (1 - s*exp(-d*a)) * L_tv
``<.>_w`` averages under the position weighting, ``<.>_mask`` under the
unweighted supervised mask, matching how the reported accuracy is normalized.
The blend weight is detached, so it reshapes the objective without adding a
gradient path of its own.
"""
ce, weights, weight_sum = terms.ce_per_token, terms.weights, terms.weight_sum
assert ce is not None, (
"the KD path produced no per-position cross-entropy, so q(gold) is "
"unavailable and the blend would silently fall back to it; modify() is "
"supposed to have rejected this combination at convert time"
)
gold_probability = torch.exp(-ce)
tv_loss = ((1.0 - gold_probability) * weights).sum() / weight_sum
if self.dflash_lk_loss_type == "tv":
return tv_loss
ce_loss = (ce * weights).sum() / weight_sum
mask = terms.supervised_mask
acceptance = (gold_probability.detach() * mask).sum() / (mask.sum() + 1e-6)
ce_share = self.dflash_lk_ce_scale * torch.exp(-self.dflash_lk_ce_decay * acceptance)
return ce_share * ce_loss + (1.0 - ce_share) * tv_loss
def _compute_loss(
self,
logits,
input_ids,
anchor_positions,
block_keep_mask,
loss_mask,
base_logits=None,
draft_hidden=None,
base_outputs=None,
):
"""Backbone DFlash loss plus the candidate-selector cross-entropy.
Reuses ``HFDFlashModel._compute_loss`` for the backbone term, then rebuilds
the same target/weight alignment for the selector term. Reported accuracy
stays the backbone's top-1, so DFlash and DFlash2 runs remain comparable;
the selector's own accuracy is logged separately.
"""
loss, accuracy, terms = super()._compute_loss(
logits,
input_ids,
anchor_positions,
block_keep_mask,
loss_mask,
base_logits,
draft_hidden=draft_hidden,
base_outputs=base_outputs,
return_terms=True,
)
if self.dflash_lk_loss_type != "ce":
loss = self._lk_loss(terms)
if self.dflash_selector_loss_alpha <= 0 or draft_hidden is None:
return loss, accuracy
bsz, seq_len = input_ids.shape
block_size = self.dflash_block_size
n_blocks = anchor_positions.shape[1]
device = input_ids.device
offsets = torch.arange(block_size, device=device).view(1, 1, -1)
label_indices = anchor_positions.unsqueeze(-1) + offsets
valid_label = label_indices < seq_len
safe_label_indices = label_indices.clamp(max=seq_len - 1)
expanded_ids = input_ids.unsqueeze(1).expand(-1, n_blocks, -1)
target_ids = torch.gather(expanded_ids, 2, safe_label_indices)
# Same supervision mask as the backbone loss: valid block, in bounds, not the
# anchor slot, and inside the answer span. Position weighting (decay/D-PACE) is
# deliberately not applied — it shapes *where* the backbone spends capacity,
# while the selector should learn every position's transition equally.
weight_mask = block_keep_mask.unsqueeze(-1).expand(-1, -1, block_size).float()
weight_mask = weight_mask * valid_label.float()
weight_mask = weight_mask * (offsets > 0).float()
weight_mask = weight_mask * torch.gather(
loss_mask.unsqueeze(1).expand(-1, n_blocks, -1), 2, safe_label_indices
)
# Teacher-forced predecessor of block position k is the real token at anchor+k-1.
# Only offsets 1..block_size-1 are supervised (weight_mask zeroes slot 0), so the
# first supervised position, k=1, has the anchor's own token as its predecessor.
# Slot 0's entry here resolves to anchor-1 and never contributes.
#
# This is the train/serve contract: CandidateSelector.greedy_path seeds its walk
# from the anchor token, so ITS position 0 corresponds to block offset 1. A caller
# that hands greedy_path the full 0..block_size-1 candidate set is off by one.
predecessor_ids = torch.gather(expanded_ids, 2, (safe_label_indices - 1).clamp(min=0))
selector_loss, selector_accuracy, selector_coverage = self._selector_loss(
logits.reshape(bsz, n_blocks, block_size, -1),
target_ids,
draft_hidden.reshape(bsz, n_blocks, block_size, -1),
predecessor_ids,
weight_mask,
)
self._selector_metrics = {
"selector_accuracy": selector_accuracy,
"selector_coverage": selector_coverage,
}
return loss + self.dflash_selector_loss_alpha * selector_loss, accuracy
@@ -0,0 +1,317 @@
# 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)
@@ -124,7 +124,12 @@ class FakeBaseConfig(PretrainedConfig):
)
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
@@ -179,6 +184,8 @@ class FakeBaseModel(PreTrainedModel):
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
)
@@ -203,7 +210,9 @@ class FakeBaseModel(PreTrainedModel):
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),
rope_theta=getattr(base_cfg, "rope_theta", None),
# 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
),
@@ -527,14 +527,16 @@ class LiLiCorrModule(DFlashModule):
``DFlashDecoderLayer`` already exposes, so the convolution itself is shared code
rather than a second implementation of the same arithmetic.
The initialization is the one deliberate difference. DFlash2 draws
``kernel_projection`` from ``normal_(0, initializer_range)``, so its convolution
is not the identity at step 0. Here it is zero by default, and since
``base_kernel`` is identity at tap 0 the whole wrapper is then an *exact*
identity at init: ``prepare`` emits a zero dynamic kernel, ``coefficients ==
base``, and the convolution returns its input unchanged. That makes the
difference between a conv and a non-conv run attributable to the convolutions
rather than to a perturbed starting point.
The initialization is written here rather than inherited.
``conv_projection_init_std`` is zero by default, and since ``base_kernel`` is
identity at tap 0 the whole wrapper is then an *exact* identity at init:
``prepare`` emits a zero dynamic kernel, ``coefficients == base``, and the
convolution returns its input unchanged. That makes the difference between a
conv and a non-conv run attributable to the convolutions rather than to a
perturbed starting point. Assigning the weight explicitly is what keeps that
property LiLiCorr's own: it holds whatever ``DFlashGroupedConv`` does for
DFlash2, which since the DFlash2 merge also zeroes the projection in its own
constructor.
``conv_projection_init_std`` is a separate key from ``initializer_range`` on
purpose: the latter also seeds the reranker, so overloading it would couple two
@@ -0,0 +1,104 @@
# DFlash2 speculative-decoding training recipe.
#
# DFlash2 (https://inco.ai/blog/dflash2/) reuses the DFlash mode/pipeline and adds
# two components, selected via dflash_architecture_config.projector_type=dflash2:
# - a grouped dynamic depthwise convolution around every attention/MLP sublayer,
# giving each block position a view of its predecessors inside the block;
# - a low-rank candidate selector that scores transitions between adjacent block
# positions' top-k candidates, so serving walks one coherent path.
# The selector is trained by an extra cross-entropy term weighted by
# dflash_selector_loss_alpha. Online training is the default path (data.mode=online).
# Override fields via an OmegaConf dotlist.
# modelopt-schema: modelopt.recipe.config.ModelOptDFlashRecipe
metadata:
description: DFlash2 training recipe (DFlash backbone + sublayer conv + candidate selector).
# maps to ModelArguments (main.py)
model:
model_name_or_path:
trust_remote_code: false
use_fake_base_for_offline: false
# maps to DataArguments (main.py)
data:
mode: online
data_path:
offline_data_path:
# Jinja chat template with {% generation %} tags for answer_only_loss.
chat_template:
# maps to TrainingArguments (main.py)
training:
# --- commonly modified ---
output_dir:
num_train_epochs: 6
per_device_train_batch_size: 1
learning_rate: 6.0e-4
warmup_ratio: 0.04
training_seq_len: 3072
logging_steps: 50
save_steps: 2000
cp_size: 1
dp_shard_size: 1
disable_tqdm: true
# Keep off: eval takes a plain per-position argmax, so the candidate selector is
# not applied. The convolutions are -- they run inside the backbone layers -- so AR
# here measures backbone + convolutions and understates the trained model. Compare
# via export + the offline acceptance-length harness instead.
estimate_ar: false
ar_validate_steps: 0
answer_only_loss: true
# --- rarely modified ---
do_eval: false
lr_scheduler_type: linear
save_strategy: steps
weight_decay: 0.0
max_grad_norm: 1.0
dataloader_drop_last: true
bf16: true
tf32: true
remove_unused_columns: false
# Safe default: the selector params are unused when
# dflash_selector_loss_alpha == 0, which would otherwise trip DDP.
ddp_find_unused_parameters: true
ddp_timeout: 1800
report_to: tensorboard
# maps to DFlashConfig (modelopt/torch/speculative/config.py).
dflash:
dflash_block_size: 16
dflash_num_anchors: 256
dflash_use_torch_compile: false
dflash_self_logit_distillation: false
# gamma for exponential loss decay (block_size=16 -> 7).
dflash_loss_decay_factor: 7.0
# Qwen3 has no native mask token; 151669 is an unused id used by the reference.
dflash_mask_token_id: 151669
# Weight of the candidate-selector cross-entropy term (0 disables it and trains
# the backbone + convolutions only).
dflash_selector_loss_alpha: 1.0
# Anneal the block objective from cross-entropy toward acceptance as acceptance
# rises; see dflash_lk_loss_type. Requires dflash_self_logit_distillation: false.
dflash_lk_loss_type: lambda
dflash_lk_ce_scale: 1.0
dflash_lk_ce_decay: 1.0
dflash_architecture_config:
num_hidden_layers: 5
# Draft attention/MLP dims — set explicitly (the draft is an independent
# Qwen3 model and does NOT inherit these from the base). GQA: 8 KV heads.
num_attention_heads: 32
num_key_value_heads: 8
head_dim: 128
intermediate_size: 12288
projector_type: dflash2
# Grouped dynamic depthwise convolution. conv_kernel_size is the number of taps
# (2 = each position also sees its predecessor); it must not exceed the block
# size. conv_group_size must divide hidden_size.
conv_kernel_size: 2
conv_group_size: 16
# Candidate selector: rank of the transition codebooks, and how many of the
# backbone top-k candidates per position it re-ranks.
selector_rank: 256
selector_top_k: 16
@@ -13,12 +13,13 @@
# either. It is a standalone recipe because recipes do not compose; keep the two files in
# sync when editing shared fields.
#
# THE INIT IS THE ONE DELIBERATE DIFFERENCE FROM DFlash2. `conv_projection_init_std: 0.0`
# zeroes `kernel_projection`, and `base_kernel` is identity at tap 0, so the wrapper is an
# EXACT identity at step 0: a run from this recipe begins as the plain reranker and any
# difference is attributable to the convolutions rather than to a perturbed start. DFlash2
# draws the same projection from `normal_(0, initializer_range)` and is therefore not the
# identity at init. Do not "align" the two -- they answer different questions.
# THE INIT IS THIS RECIPE'S OWN. `conv_projection_init_std: 0.0` zeroes
# `kernel_projection`, and `base_kernel` is identity at tap 0, so the wrapper is an EXACT
# identity at step 0: a run from this recipe begins as the plain reranker and any
# difference is attributable to the convolutions rather than to a perturbed start.
# `_install_sublayer_convs` assigns the weight explicitly, so this holds whatever
# `DFlashGroupedConv` does for DFlash2 -- which since the DFlash2 merge also zeroes it.
# Raise this key if you want a perturbed start instead.
#
# MEMORY. The convolutions add ~42M trainable parameters (20 tensors for a 5-layer draft),
# and the activations they hold are the binding constraint. At an 8B target, combined with
@@ -18,9 +18,14 @@
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
import torch
from modelopt.torch.export.plugins.hf_spec_export import DFlashExporter, EagleExporter
from modelopt.torch.export.plugins.hf_spec_export import (
DFlashExporter,
EagleExporter,
_get_rope_theta,
)
DEFAULT_ROPE_SCALING = {
"rope_type": "yarn",
@@ -152,3 +157,65 @@ def test_dflash_rope_theta_inherits_base_rope_parameters():
config = exporter._export_config()
assert config["rope_theta"] == 5000000.0
class TestGetRopeTheta:
"""Where a config keeps rope_theta depends on the transformers version.
Every consumer of a base config -- the exporter, the draft builder, the fake base --
has to agree on this, so they share this one reader. Reading it wrong is silent: the
draft trains and exports without complaint against a RoPE base the target never used,
and only misbehaves at serve time.
"""
def test_reads_the_rope_parameters_dict(self):
"""The transformers 5.12+ layout: the value lives only in the dict."""
assert _get_rope_theta(SimpleNamespace(rope_parameters={"rope_theta": 1000000.0})) == (
1000000.0
)
def test_reads_the_legacy_rope_scaling_dict(self):
"""Older transformers spell the same dict rope_scaling."""
assert _get_rope_theta(SimpleNamespace(rope_scaling={"rope_theta": 1000000.0})) == (
1000000.0
)
def test_prefers_the_dict_over_a_disagreeing_flat_field(self):
"""Both present and disagreeing: the dict wins.
This is the regression guard. The precedence was the other way round on main from
2026-07-30 to 2026-09-09, and no test noticed -- a config can carry the real base
in the dict while the class default (10000.0 for Qwen3) stays visible as a flat
rope_theta, so reading flat first exports a drafter whose RoPE base is 100x off.
"""
config = SimpleNamespace(rope_theta=10000.0, rope_parameters={"rope_theta": 1000000.0})
assert _get_rope_theta(config) == 1000000.0
def test_falls_back_to_a_flat_attribute(self):
"""The transformers 4.x layout: only the flat field exists."""
assert _get_rope_theta(SimpleNamespace(rope_theta=12345.0)) == 12345.0
def test_missing_everywhere_returns_the_default(self):
"""An absent base must stay absent rather than become a wrong number."""
assert _get_rope_theta(SimpleNamespace()) is None
assert _get_rope_theta(SimpleNamespace(), 7.0) == 7.0
def test_reads_a_real_config_whichever_layout_it_uses(self):
"""A real config resolves on every supported transformers version.
Asserts the outcome, not the layout: 5.12 keeps the value only in the dict while
the minimum supported version (4.57) has only the flat field and no dict at all.
The layouts themselves are pinned above, built explicitly, so they stay covered
where transformers is absent -- this is the only test here that needs it.
"""
transformers = pytest.importorskip("transformers")
config = transformers.Qwen3Config(
hidden_size=32,
num_hidden_layers=2,
num_attention_heads=4,
num_key_value_heads=2,
intermediate_size=64,
vocab_size=64,
rope_theta=1000000.0,
)
assert _get_rope_theta(config) == 1000000.0
@@ -24,7 +24,7 @@ import torch
pytest.importorskip("transformers")
import transformers
from modelopt.torch.speculative.plugins.modeling_fakebase import FakeBaseModel
from modelopt.torch.speculative.plugins.modeling_fakebase import FakeBaseConfig, FakeBaseModel
from modelopt.torch.speculative.utils import load_vlm_or_llm
_HIDDEN_SIZE = 16
@@ -162,3 +162,82 @@ def test_load_vlm_or_llm_uses_transformers5_vlm_auto_class(monkeypatch):
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)
@@ -0,0 +1,588 @@
# 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.
"""CPU unit tests for the DFlash2 speculative decoding plugin.
DFlash2 reuses the DFlash mode/pipeline and adds grouped dynamic convolutions
around every attention/MLP sublayer plus a low-rank candidate selector. These
tests cover conversion routing, the convolution's two structural invariants
(identity at initialization, no leakage across the block boundary), the selector
training objective, and the export format against the SGLang/vLLM
``DFlash2DraftModel`` layout (``attention_conv.*`` / ``mlp_conv.*`` /
``candidate_selector.*``).
"""
import json
from copy import deepcopy
import pytest
import torch
from _test_utils.torch.transformers_models import get_tiny_llama
from safetensors.torch import load_file
import modelopt.torch.speculative as mtsp
from modelopt.torch.speculative.config import DFLASH_DEFAULT_CFG
from modelopt.torch.speculative.plugins.hf_dflash import HFDFlashModel
from modelopt.torch.speculative.plugins.hf_dflash2 import HFDFlash2Model
from modelopt.torch.speculative.plugins.modeling_dflash import (
DFlashModule,
_IdentitySublayerWrapper,
)
from modelopt.torch.speculative.plugins.modeling_dflash2 import (
CandidateSelector,
DFlash2Module,
DFlashGroupedConv,
)
BLOCK_SIZE = 4
NUM_DRAFT_LAYERS = 2
SEQ_LEN = 16 # must be a multiple of BLOCK_SIZE
CONV_KERNEL_SIZE = 2
CONV_GROUP_SIZE = 4
SELECTOR_RANK = 8
SELECTOR_TOP_K = 5
ARCH_FIELDS = ["conv_kernel_size", "conv_group_size", "selector_rank", "selector_top_k"]
def _get_dflash2_config(selector_loss_alpha=1.0, block_size=BLOCK_SIZE, **arch_overrides):
"""Create a DFlash2 config for testing (dflash mode + projector_type=dflash2)."""
config = deepcopy(DFLASH_DEFAULT_CFG["config"])
config["dflash_block_size"] = block_size
config["dflash_use_torch_compile"] = False
config["dflash_mask_token_id"] = 0 # token 0 as mask for the tiny model
config["dflash_self_logit_distillation"] = False
config["dflash_selector_loss_alpha"] = selector_loss_alpha
config["dflash_architecture_config"] = {
"num_hidden_layers": NUM_DRAFT_LAYERS,
"projector_type": "dflash2",
"conv_kernel_size": CONV_KERNEL_SIZE,
"conv_group_size": CONV_GROUP_SIZE,
"selector_rank": SELECTOR_RANK,
"selector_top_k": SELECTOR_TOP_K,
**arch_overrides,
}
return config
def _make_batch(vocab_size):
torch.manual_seed(0)
input_ids = torch.randint(1, vocab_size, (2, SEQ_LEN))
return input_ids, torch.ones_like(input_ids), input_ids.clone()
class TestDFlash2Convert:
"""Test DFlash2 conversion routing and module construction."""
def test_convert_creates_dflash2_model(self):
"""projector_type=dflash2 routes to HFDFlash2Model (a HFDFlashModel subclass)."""
model = get_tiny_llama(num_hidden_layers=4)
mtsp.convert(model, [("dflash", _get_dflash2_config())])
assert isinstance(model, HFDFlash2Model)
assert isinstance(model, HFDFlashModel)
assert isinstance(model.dflash_module, DFlash2Module)
def test_every_sublayer_wrapped_in_a_convolution(self):
"""Both sublayer wrappers on every draft layer become real convolutions."""
model = get_tiny_llama(num_hidden_layers=4)
mtsp.convert(model, [("dflash", _get_dflash2_config())])
layers = model.dflash_module.layers
assert len(layers) == NUM_DRAFT_LAYERS
for layer in layers:
for conv in (layer.attention_conv, layer.mlp_conv):
assert isinstance(conv, DFlashGroupedConv)
assert conv.taps == CONV_KERNEL_SIZE
assert conv.group_size == CONV_GROUP_SIZE
def test_selector_shapes(self):
"""The candidate selector's codebooks and projection are sized from the config."""
model = get_tiny_llama(num_hidden_layers=4)
mtsp.convert(model, [("dflash", _get_dflash2_config())])
selector = model.dflash_module.candidate_selector
vocab = model.dflash_config.vocab_size
assert selector.top_k == SELECTOR_TOP_K
assert selector.predecessor_codebook.shape == (vocab, SELECTOR_RANK)
assert selector.successor_codebook.shape == (vocab, SELECTOR_RANK)
assert selector.hidden_projection.out_features == SELECTOR_RANK
assert selector.hidden_projection.bias is None
def test_new_params_trainable(self):
"""The convolution and selector parameters are trainable."""
model = get_tiny_llama(num_hidden_layers=4)
mtsp.convert(model, [("dflash", _get_dflash2_config())])
new = [
(n, p)
for n, p in model.named_parameters()
if "_conv." in n or "candidate_selector" in n
]
assert len(new) >= 2 * NUM_DRAFT_LAYERS * 2 + 3
assert all(p.requires_grad for _, p in new)
@pytest.mark.parametrize("field", ARCH_FIELDS)
def test_missing_architecture_field_raises(self, field):
"""projector_type=dflash2 without a required architecture field is an error."""
config = _get_dflash2_config()
del config["dflash_architecture_config"][field]
model = get_tiny_llama(num_hidden_layers=4)
with pytest.raises(ValueError, match=field):
mtsp.convert(model, [("dflash", config)])
def test_conv_kernel_larger_than_block_raises(self):
"""A convolution tap count exceeding the block size is an error."""
model = get_tiny_llama(num_hidden_layers=4)
config = _get_dflash2_config(conv_kernel_size=BLOCK_SIZE + 1)
with pytest.raises(ValueError, match="conv_kernel_size"):
mtsp.convert(model, [("dflash", config)])
def test_conv_group_size_must_divide_hidden(self):
"""A conv_group_size that does not divide hidden_size is an error."""
model = get_tiny_llama(num_hidden_layers=4)
config = _get_dflash2_config(conv_group_size=model.config.hidden_size - 1)
with pytest.raises(ValueError, match="conv_group_size"):
mtsp.convert(model, [("dflash", config)])
def test_dflash_mode_still_creates_plain_dflash(self):
"""Without projector_type=dflash2, conversion still yields a plain DFlash model."""
config = deepcopy(DFLASH_DEFAULT_CFG["config"])
config["dflash_mask_token_id"] = 0
config["dflash_architecture_config"] = {"num_hidden_layers": NUM_DRAFT_LAYERS}
model = get_tiny_llama(num_hidden_layers=4)
mtsp.convert(model, [("dflash", config)])
assert isinstance(model, HFDFlashModel)
assert not isinstance(model, HFDFlash2Model)
assert type(model.dflash_module) is DFlashModule
# The sublayer seam stays a parameterless no-op for a plain DFlash draft.
for layer in model.dflash_module.layers:
assert isinstance(layer.attention_conv, _IdentitySublayerWrapper)
assert isinstance(layer.mlp_conv, _IdentitySublayerWrapper)
assert not any("_conv." in n for n, _ in model.dflash_module.named_parameters())
class TestDFlashGroupedConv:
"""Test the convolution's structural invariants directly."""
def _conv(self, hidden_size=32, taps=2):
torch.manual_seed(0)
return DFlashGroupedConv(
hidden_size=hidden_size, block_size=BLOCK_SIZE, taps=taps, group_size=CONV_GROUP_SIZE
).double()
def _trained_conv(self, **kwargs):
"""A conv whose dynamic kernel is non-zero, i.e. what training produces.
The default construction is an exact identity, so the structural assertions
below would pass vacuously on it.
"""
conv = self._conv(**kwargs)
with torch.no_grad():
torch.nn.init.normal_(conv.kernel_projection.weight, std=0.5)
return conv
def test_identity_at_initialization(self):
"""A freshly built conv is an exact identity -- no manual zeroing needed.
Both halves start as a no-op: the base kernel is an identity and
``kernel_projection`` is zero-initialized. This is what makes enabling
DFlash2 a stable extension of a DFlash backbone rather than a perturbation
of it, and it is asserted on the DEFAULT construction so that a change to
either half fails here.
"""
conv = self._conv()
assert conv.kernel_projection.weight.abs().max() == 0.0
x = torch.randn(2, SEQ_LEN, 32, dtype=torch.double)
out = conv.finish(*conv.prepare(x))
assert torch.equal(out, x)
def test_taps_do_not_cross_the_block_boundary(self):
"""Perturbing the last position of a block leaves later blocks untouched."""
conv = self._trained_conv()
x = torch.randn(2, SEQ_LEN, 32, dtype=torch.double)
baseline = conv.finish(*conv.prepare(x))
perturbed_input = x.clone()
perturbed_input[:, BLOCK_SIZE - 1] += 5.0
perturbed = conv.finish(*conv.prepare(perturbed_input))
assert torch.allclose(baseline[:, BLOCK_SIZE:], perturbed[:, BLOCK_SIZE:], atol=1e-12)
assert not torch.allclose(baseline[:, :BLOCK_SIZE], perturbed[:, :BLOCK_SIZE])
def test_intra_block_dependency_is_backward_only(self):
"""A position influences its successors inside the block, never its predecessors.
This is the point of the convolution: it injects the sequential dependency the
parallel backbone lacks, without letting a position see the future.
"""
conv = self._trained_conv()
x = torch.randn(2, SEQ_LEN, 32, dtype=torch.double)
baseline = conv.finish(*conv.prepare(x))
perturbed_input = x.clone()
perturbed_input[:, 1] += 5.0
perturbed = conv.finish(*conv.prepare(perturbed_input))
assert torch.allclose(baseline[:, 0], perturbed[:, 0], atol=1e-12)
assert not torch.allclose(baseline[:, 2], perturbed[:, 2])
def test_sequence_length_must_be_block_aligned(self):
"""A sequence length not divisible by the block size is an error."""
conv = self._conv()
with pytest.raises(ValueError, match="block_size"):
conv.prepare(torch.randn(1, BLOCK_SIZE + 1, 32, dtype=torch.double))
class TestDFlash2Forward:
"""Test the DFlash2 training forward (online path on CPU)."""
def test_forward_grads_reach_conv_and_selector(self):
"""Backward fills gradients on the convolutions, the selector and the backbone."""
model = get_tiny_llama(num_hidden_layers=4)
mtsp.convert(model, [("dflash", _get_dflash2_config())])
model.train()
input_ids, attention_mask, labels = _make_batch(model.dflash_config.vocab_size)
out = model(input_ids=input_ids, attention_mask=attention_mask, labels=labels)
assert out.loss.requires_grad
assert out.loss.dim() == 0
out.loss.backward()
module = model.dflash_module
selector = module.candidate_selector
for grad in (
selector.successor_codebook.grad,
module.layers[0].attention_conv.base_kernel.grad,
module.layers[0].mlp_conv.kernel_projection.weight.grad,
module.fc.weight.grad,
):
assert grad is not None and torch.isfinite(grad).all()
assert grad.abs().sum() > 0
# The zero-initialized ``successor_codebook`` is one factor of a bilinear form,
# so on the FIRST step the other two factors get exactly zero gradient. This is
# a one-step delay, not a dead branch: the assertions below show they train as
# soon as ``successor_codebook`` moves off zero. The convolution has no such
# delay -- its delta is added to a non-zero base kernel, hence the grad above.
for grad in (selector.predecessor_codebook.grad, selector.hidden_projection.weight.grad):
assert grad is not None and torch.isfinite(grad).all()
assert grad.abs().sum() == 0
def test_selector_factors_train_once_the_successor_codebook_moves(self):
"""After the first step the whole bilinear selector receives gradient.
Guards the zero-init of ``successor_codebook`` against becoming a permanently
dead branch rather than the intended one-step warm start.
"""
model = get_tiny_llama(num_hidden_layers=4)
mtsp.convert(model, [("dflash", _get_dflash2_config())])
model.train()
selector = model.dflash_module.candidate_selector
with torch.no_grad(): # stand in for the first optimizer step
torch.nn.init.normal_(selector.successor_codebook, std=0.02)
input_ids, attention_mask, labels = _make_batch(model.dflash_config.vocab_size)
model(input_ids=input_ids, attention_mask=attention_mask, labels=labels).loss.backward()
for grad in (
selector.predecessor_codebook.grad,
selector.successor_codebook.grad,
selector.hidden_projection.weight.grad,
):
assert grad is not None and torch.isfinite(grad).all()
assert grad.abs().sum() > 0
def test_selector_metrics_reported(self):
"""The forward records selector accuracy and top-k coverage."""
model = get_tiny_llama(num_hidden_layers=4)
mtsp.convert(model, [("dflash", _get_dflash2_config())])
model.train()
input_ids, attention_mask, labels = _make_batch(model.dflash_config.vocab_size)
out = model(input_ids=input_ids, attention_mask=attention_mask, labels=labels)
# Carried out on the forward output, not left on a private attribute, so the
# trainer can log them; and kept as tensors to avoid a per-step CPU-GPU sync.
metrics = out["selector_metrics"]
for key in ("selector_accuracy", "selector_coverage"):
assert torch.is_tensor(metrics[key])
assert 0.0 <= metrics[key].item() <= 1.0
def test_selector_alpha_zero_disables_the_term(self):
"""alpha=0 trains the backbone and convolutions only; the selector gets no grad."""
model = get_tiny_llama(num_hidden_layers=4)
mtsp.convert(model, [("dflash", _get_dflash2_config(selector_loss_alpha=0.0))])
model.train()
input_ids, attention_mask, labels = _make_batch(model.dflash_config.vocab_size)
out = model(input_ids=input_ids, attention_mask=attention_mask, labels=labels)
out.loss.backward()
module = model.dflash_module
codebook_grad = module.candidate_selector.predecessor_codebook.grad
assert codebook_grad is None or codebook_grad.abs().sum() == 0
# The convolutions still train: they live inside the backbone.
conv_grad = module.layers[0].attention_conv.kernel_projection.weight.grad
assert conv_grad is not None and conv_grad.abs().sum() > 0
def test_selector_loss_increases_total_loss(self):
"""The selector term adds to the backbone loss rather than replacing it."""
input_ids, attention_mask, labels = _make_batch(32)
losses = {}
for alpha in (0.0, 1.0):
torch.manual_seed(0)
model = get_tiny_llama(num_hidden_layers=4)
mtsp.convert(model, [("dflash", _get_dflash2_config(selector_loss_alpha=alpha))])
model.train()
torch.manual_seed(0)
out = model(input_ids=input_ids, attention_mask=attention_mask, labels=labels)
losses[alpha] = float(out.loss.detach())
assert losses[1.0] > losses[0.0]
def test_overfits_a_single_batch(self):
"""A few steps on one batch drive backbone and selector accuracy up.
Guards the target/predecessor alignment: a misaligned selector objective still
produces a finite decreasing loss, but its accuracy does not reach 1.
"""
model = get_tiny_llama(num_hidden_layers=4)
mtsp.convert(model, [("dflash", _get_dflash2_config())])
model.train()
input_ids, attention_mask, labels = _make_batch(model.dflash_config.vocab_size)
optimizer = torch.optim.AdamW([p for p in model.parameters() if p.requires_grad], lr=5e-3)
for _ in range(60):
out = model(input_ids=input_ids, attention_mask=attention_mask, labels=labels)
optimizer.zero_grad()
out.loss.backward()
optimizer.step()
assert out.train_acc[0][0] > 0.9
assert out["selector_metrics"]["selector_accuracy"].item() > 0.9
class TestCandidateSelectorAlignment:
"""Pin the offset convention shared by the training objective and the lattice walk."""
def _selector(self, vocab=16, rank=4, top_k=3, hidden=8):
torch.manual_seed(0)
selector = CandidateSelector(
hidden_size=hidden, vocab_size=vocab, rank=rank, top_k=top_k, std=0.5
).double()
with torch.no_grad(): # a fresh selector is a no-op; give it a real transition
torch.nn.init.normal_(selector.successor_codebook, std=0.5)
return selector
def _selector_terms(self, selector, target_ids, candidate_ids):
"""Drive HFDFlash2Model._selector_loss with a hand-built candidate set."""
b, length, k = candidate_ids.shape
model = HFDFlash2Model.__new__(HFDFlash2Model)
model.dflash_module = type("_M", (), {"candidate_selector": selector})()
logits = torch.zeros(b, length, 32, dtype=torch.double)
# Put the chosen candidates on top so logits.topk reproduces candidate_ids.
for slot in range(k):
logits.scatter_(-1, candidate_ids[..., slot : slot + 1], float(k - slot))
return HFDFlash2Model._selector_loss(
model,
logits,
target_ids,
torch.randn(b, length, 8, dtype=torch.double),
torch.zeros(b, length, dtype=torch.long),
torch.ones(b, length, dtype=torch.double),
)
def test_a_miss_carries_no_selector_gradient(self):
"""A gold token outside the backbone's top-k is excluded, not substituted in.
Serving only ever shows the strict top-k, so supervising a set with the gold
forced into it would teach the selector to override the unary ranking on a
candidate set it will never see. Contrasted against a covered set below, which
must still train -- masking everything would pass a one-sided assertion.
"""
selector = self._selector(vocab=32, top_k=2)
candidates = torch.tensor([[[5, 6], [5, 6], [5, 6]]] * 2)
_, _, covered = self._selector_terms(
selector, torch.full((2, 3), 5, dtype=torch.long), candidates
)
assert float(covered) == 1.0
loss, _, coverage = self._selector_terms(
selector, torch.full((2, 3), 9, dtype=torch.long), candidates
)
assert float(coverage) == 0.0
assert float(loss.detach()) == 0.0
selector.zero_grad()
loss.backward()
assert selector.successor_codebook.grad.abs().sum() == 0.0
def test_a_covered_position_does_train_the_selector(self):
"""The mask must not be so aggressive that nothing trains."""
selector = self._selector(vocab=32, top_k=2)
candidates = torch.tensor([[[5, 6], [5, 6], [5, 6]]] * 2)
loss, _, coverage = self._selector_terms(
selector, torch.full((2, 3), 5, dtype=torch.long), candidates
)
assert float(coverage) == 1.0
assert float(loss.detach()) > 0.0
selector.zero_grad()
loss.backward()
assert selector.successor_codebook.grad.abs().sum() > 0.0
def test_greedy_path_position_zero_is_seeded_by_the_anchor(self):
"""``greedy_path`` position 0 scores against the anchor, i.e. block offset 1.
``HFDFlash2Model._compute_loss`` supervises offsets 1..block_size-1 and pairs
offset 1 with the anchor's own token. If either side changes its base offset
the objective and the walk silently disagree by one position, which a finite
decreasing loss does not catch.
"""
selector = self._selector()
b, length, k, hidden = 2, 3, 3, 8
candidate_ids = torch.randint(0, 16, (b, length, k))
unary = torch.randn(b, length, k, dtype=torch.double)
hiddens = torch.randn(b, length, hidden, dtype=torch.double)
anchor = torch.randint(0, 16, (b,))
path = selector.greedy_path(candidate_ids, unary, hiddens, anchor)
expected_first = candidate_ids[:, 0].gather(
1,
selector.score_candidates(
candidate_ids[:, 0], unary[:, 0], hiddens[:, 0], anchor
).argmax(dim=-1, keepdim=True),
)[:, 0]
assert torch.equal(path[:, 0], expected_first)
def test_greedy_path_feeds_each_choice_forward(self):
"""Position n+1 is scored against the token position n actually selected."""
selector = self._selector()
b, length, k, hidden = 2, 3, 3, 8
candidate_ids = torch.randint(0, 16, (b, length, k))
unary = torch.randn(b, length, k, dtype=torch.double)
hiddens = torch.randn(b, length, hidden, dtype=torch.double)
anchor = torch.randint(0, 16, (b,))
path = selector.greedy_path(candidate_ids, unary, hiddens, anchor)
expected_second = candidate_ids[:, 1].gather(
1,
selector.score_candidates(
candidate_ids[:, 1], unary[:, 1], hiddens[:, 1], path[:, 0]
).argmax(dim=-1, keepdim=True),
)[:, 0]
assert torch.equal(path[:, 1], expected_second)
class TestDFlash2BlockObjective:
"""The cross-entropy/acceptance blend selected by ``dflash_lk_loss_type``."""
def _loss(self, **overrides):
torch.manual_seed(0)
config = _get_dflash2_config()
config.update(overrides)
model = get_tiny_llama(num_hidden_layers=4)
mtsp.convert(model, [("dflash", config)])
model.train()
torch.manual_seed(1)
input_ids, attention_mask, labels = _make_batch(model.dflash_config.vocab_size)
out = model(input_ids=input_ids, attention_mask=attention_mask, labels=labels)
return float(out.loss.detach())
def test_zero_decay_and_unit_scale_is_exactly_cross_entropy(self):
"""A constant CE share of 1 must leave the objective bit-identical to 'ce'.
This is the degenerate case that catches a blend wired up backwards: a wrong
sign or a swapped term still produces a finite decreasing loss.
"""
assert self._loss(
dflash_lk_loss_type="lambda", dflash_lk_ce_scale=1.0, dflash_lk_ce_decay=0.0
) == self._loss(dflash_lk_loss_type="ce")
def test_zero_scale_is_exactly_the_acceptance_term(self):
"""A CE share of 0 must leave the objective bit-identical to 'tv'."""
assert self._loss(dflash_lk_loss_type="lambda", dflash_lk_ce_scale=0.0) == self._loss(
dflash_lk_loss_type="tv"
)
def test_blend_lies_between_its_two_terms(self):
ce = self._loss(dflash_lk_loss_type="ce")
tv = self._loss(dflash_lk_loss_type="tv")
blended = self._loss(dflash_lk_loss_type="lambda")
assert min(ce, tv) <= blended <= max(ce, tv)
def test_acceptance_term_is_a_probability(self):
"""1 - q(gold) is a weighted mean of probabilities, so it cannot leave [0, 1]."""
loss = self._loss(dflash_lk_loss_type="tv", dflash_selector_loss_alpha=0.0)
assert 0.0 <= loss <= 1.0
@pytest.mark.parametrize("loss_type", ["tv", "lambda"])
def test_distillation_conflict_is_rejected(self, loss_type):
"""Both terms read q(gold), which the KD path never forms."""
config = _get_dflash2_config()
config["dflash_lk_loss_type"] = loss_type
config["dflash_self_logit_distillation"] = True
model = get_tiny_llama(num_hidden_layers=4)
with pytest.raises(ValueError, match="dflash_self_logit_distillation"):
mtsp.convert(model, [("dflash", config)])
class TestDFlash2Export:
"""Test the DFlash2 export format (weights + config)."""
def _export(self, tmp_path):
model = get_tiny_llama(num_hidden_layers=4)
mtsp.convert(model, [("dflash", _get_dflash2_config())])
export_dir = tmp_path / "exported"
model.get_exporter().export(export_dir)
return export_dir
def test_export_weight_keys_match_reference(self, tmp_path):
"""Exported weights carry the DFlash2 tensors under reference names, no prefix."""
sd = load_file(str(self._export(tmp_path) / "model.safetensors"))
for key in sd:
assert "dflash_module." not in key
assert "rotary_emb" not in key
assert "candidate_selector.predecessor_codebook" in sd
assert "candidate_selector.successor_codebook" in sd
assert "candidate_selector.hidden_projection.weight" in sd
for layer_idx in range(NUM_DRAFT_LAYERS):
for wrapper in ("attention_conv", "mlp_conv"):
assert f"layers.{layer_idx}.{wrapper}.base_kernel" in sd
assert f"layers.{layer_idx}.{wrapper}.kernel_projection.weight" in sd
def test_export_config_declares_dflash2_architecture(self, tmp_path):
"""config.json selects the DFlash2 serving path and carries its fields.
The architecture name matters: a checkpoint declaring ``DFlashDraftModel``
loads as a plain DFlash draft and silently ignores these weights.
"""
with open(self._export(tmp_path) / "config.json") as f:
cfg = json.load(f)
assert cfg["architectures"] == ["DFlash2DraftModel"]
# Emitted by the published-contract commit and read by the vLLM loader. Assert the
# top-level block_size too, not just the nested one: DFlash2Exporter derives the
# nested value FROM the top-level, so a wrong top-level passes the nested check.
assert cfg["is_causal"] is False
assert cfg["block_size"] == BLOCK_SIZE
dflash_config = cfg["dflash_config"]
assert dflash_config["block_size"] == BLOCK_SIZE
assert dflash_config["projector_type"] == "dflash2"
assert dflash_config["conv_kernel_size"] == CONV_KERNEL_SIZE
assert dflash_config["conv_group_size"] == CONV_GROUP_SIZE
assert dflash_config["selector_rank"] == SELECTOR_RANK
assert dflash_config["selector_top_k"] == SELECTOR_TOP_K
assert "mask_token_id" in dflash_config
assert "target_layer_ids" in dflash_config
@@ -420,6 +420,99 @@ class TestLiLiCorrForward:
assert draft_tokens.shape == (1, 3)
class TestLiLiCorrSublayerConvs:
"""The optional grouped-convolution path, shared with DFlash2 via DFlashGroupedConv.
This is the composition `lilicorr_conv.yaml` ships and the only place
`_install_sublayer_convs` runs, so it is also what guards LiLiCorr's init from
changes made on the DFlash2 side of the shared class.
"""
CONV_KWARGS = {"conv_kernel_size": 2, "conv_group_size": 8}
def _conv_converted(self, **arch_overrides):
config = _get_lilicorr_config()
config["dflash_architecture_config"].update(self.CONV_KWARGS)
config["dflash_architecture_config"].update(arch_overrides)
model = get_tiny_llama(num_hidden_layers=4)
mtsp.convert(model, [("dflash", config)])
return model
def test_convs_replace_the_no_op_wrappers(self):
"""Both sublayer wrappers on every draft layer become a real convolution."""
from modelopt.torch.speculative.plugins.modeling_dflash2 import DFlashGroupedConv
module = self._conv_converted().dflash_module
assert len(module.layers) == NUM_DRAFT_LAYERS
for layer in module.layers:
for wrapper_name in ("attention_conv", "mlp_conv"):
conv = getattr(layer, wrapper_name)
assert isinstance(conv, DFlashGroupedConv)
assert conv.taps == self.CONV_KWARGS["conv_kernel_size"]
assert conv.group_size == self.CONV_KWARGS["conv_group_size"]
def test_plain_lilicorr_keeps_the_no_op_wrappers(self):
"""Without the two geometry keys the draft is the plain reranker."""
from modelopt.torch.speculative.plugins.modeling_dflash import _IdentitySublayerWrapper
module = _converted().dflash_module
for layer in module.layers:
assert isinstance(layer.attention_conv, _IdentitySublayerWrapper)
assert isinstance(layer.mlp_conv, _IdentitySublayerWrapper)
def test_default_init_is_an_exact_identity(self):
"""`conv_projection_init_std` defaults to 0, so a conv run starts as the plain reranker.
This is LiLiCorr's own choice, written by `_install_sublayer_convs` after the
conv is constructed. It must not depend on how `DFlashGroupedConv` happens to
initialize itself for DFlash2.
"""
module = self._conv_converted().dflash_module
for layer in module.layers:
for wrapper_name in ("attention_conv", "mlp_conv"):
conv = getattr(layer, wrapper_name)
assert conv.kernel_projection.weight.abs().max() == 0.0
conv = module.layers[0].attention_conv.double()
x = torch.randn(2, SEQ_LEN, conv.base_kernel.shape[-1], dtype=torch.double)
assert torch.equal(conv.finish(*conv.prepare(x)), x)
def test_non_zero_init_std_perturbs_the_start(self):
"""A non-zero `conv_projection_init_std` is still honoured, and only LiLiCorr sets it."""
module = self._conv_converted(conv_projection_init_std=0.5).dflash_module
for layer in module.layers:
for wrapper_name in ("attention_conv", "mlp_conv"):
conv = getattr(layer, wrapper_name)
assert conv.kernel_projection.weight.abs().max() > 0.0
conv = module.layers[0].attention_conv.double()
x = torch.randn(2, SEQ_LEN, conv.base_kernel.shape[-1], dtype=torch.double)
assert not torch.allclose(conv.finish(*conv.prepare(x)), x)
def test_forward_trains_the_convs(self):
"""The conv path produces a finite loss and gradients reach the convolutions."""
model = self._conv_converted()
model.train()
out = model(**_make_batch(model.dflash_config.vocab_size))
assert torch.isfinite(out.loss).all()
out.loss.backward()
for layer in model.dflash_module.layers:
for wrapper_name in ("attention_conv", "mlp_conv"):
conv = getattr(layer, wrapper_name)
for grad in (conv.base_kernel.grad, conv.kernel_projection.weight.grad):
assert grad is not None and torch.isfinite(grad).all()
assert grad.abs().sum() > 0
def test_one_geometry_key_alone_is_rejected(self):
"""Half the geometry would silently build a draft with no convolutions."""
config = _get_lilicorr_config()
config["dflash_architecture_config"]["conv_kernel_size"] = 2
model = get_tiny_llama(num_hidden_layers=4)
with pytest.raises(ValueError, match="conv_kernel_size"):
mtsp.convert(model, [("dflash", config)])
class TestLiLiCorrOptimization:
"""The objective is trainable: a fixed batch is driven down.
@@ -0,0 +1,107 @@
# DFlash2 online speculative decoding training for Qwen3-8B.
#
# DFlash2 = the DFlash draft backbone plus two additions:
# * a grouped dynamic depthwise convolution wrapped around every attention and
# MLP sublayer (taps clipped at block boundaries, identity-initialized so a
# fresh DFlash2 draft computes exactly what its DFlash backbone would), and
# * a low-rank candidate selector that scores transitions between adjacent
# block positions' top-k candidates, so serving walks one coherent path.
# See the dflash2.yaml recipe and
# modelopt/torch/speculative/plugins/{modeling,hf}_dflash2.py.
#
# 2-step pipeline:
# task_0: Build training conversations (Daring-Anteater multi-turn SFT, 50K)
# task_1: Online DFlash2 training + export of the drafter checkpoint
#
# As configured this is a short convergence check (max_steps=2000), matching the
# other Qwen3-8B online examples so it finishes on one node. To reproduce the
# published Qwen3-8B DFlash2 curve instead, see "Full run" below.
#
# Reference: inco.ai/blog/dflash2 | vLLM PR #52816 (serving support)
#
# Usage:
# uv run launch.py --yaml examples/Qwen/Qwen3-8B/hf_online_dflash2.yaml --yes
# uv run slurm.py --yaml modules/Model-Optimizer/tools/launcher/examples/Qwen/Qwen3-8B/hf_online_dflash2.yaml --yes
#
# Full run (the strict A/B against DFlash / DSpark / Domino, 3 epochs = ~92K
# steps on the 1.96M-conversation Spec-Decoding-Dataset-v1, 8 nodes x 8 H100):
# - point data.data_path at that corpus instead of task_0's output
# - training.num_train_epochs=3 and drop training.max_steps
# - training.save_steps=4000
# - slurm_config.nodes=8 (global batch stays 64 = nodes x gpus x bs x accum)
# Every other knob below is already the A/B setting.
job_name: Qwen3-8B_DFlash2_online
pipeline:
global_vars:
hf_model: /hf-local/Qwen/Qwen3-8B
# Step 1: Build input conversations. example_data_config.yaml enables only the
# daring-anteater source (train: 50000) — multi-turn SFT with real assistant
# completions. --full-conversations keeps those completions so answer_only_loss
# has assistant spans to mask. make_dataset.sh writes /scratchspace/data/train.jsonl.
task_0:
script: common/eagle3/make_dataset.sh
args:
- -f modules/Model-Optimizer/examples/dataset/example_data_config.yaml
- --full-conversations
slurm_config:
_factory_: "slurm_factory"
nodes: 1
ntasks_per_node: 1
gpus_per_node: 1
container: nvcr.io/nvidia/tensorrt-llm/release:1.3.0rc10
# Step 2: Online DFlash2 training (the script exports the drafter at the end).
# Consumes the conversations built in task_0 (shared via /scratchspace).
task_1:
script: common/specdec/dflash_online_training.sh
args:
- --config modules/Model-Optimizer/modelopt_recipes/general/speculative_decoding/dflash2.yaml
- model.model_name_or_path=<<global_vars.hf_model>>
- data.data_path=/scratchspace/data/train.jsonl
- data.chat_template=examples/Qwen/Qwen3-8B/chat_template_train.jinja
- training.output_dir=/scratchspace/dflash2_bs16
- training.per_device_train_batch_size=1
- training.num_train_epochs=1
- training.max_steps=2000
- training.training_seq_len=4096
- training.learning_rate=6.0e-4
- training.warmup_ratio=0.04
- training.warmup_steps=0
- training.lr_scheduler_type=linear
- training.save_steps=5000
- training.logging_steps=100
- training.disable_tqdm=true
- training.answer_only_loss=true
# Draft backbone — identical to the DFlash / DSpark / Domino arms so the
# only difference between them is the correction head.
- dflash.dflash_block_size=16
- dflash.dflash_num_anchors=512
- dflash.dflash_loss_decay_factor=7
- dflash.dflash_mask_token_id=151669
- dflash.dflash_self_logit_distillation=false
- dflash.dflash_architecture_config.num_hidden_layers=5
- dflash.dflash_architecture_config.num_attention_heads=32
- dflash.dflash_architecture_config.num_key_value_heads=8
- dflash.dflash_architecture_config.head_dim=128
- dflash.dflash_architecture_config.intermediate_size=12288
# DFlash2 knobs (also set in the recipe; repeated here for visibility).
# A draft dim NOT set explicitly falls back to the Qwen3Config default, not
# to the base model's — hence the five dims above are always spelled out.
- dflash.dflash_architecture_config.projector_type=dflash2
- dflash.dflash_architecture_config.conv_kernel_size=2
- dflash.dflash_architecture_config.conv_group_size=16
- dflash.dflash_architecture_config.selector_rank=256
- dflash.dflash_architecture_config.selector_top_k=16
- dflash.dflash_selector_loss_alpha=1.0
# Sliding-window draft attention, matching the published DFlash2 drafters.
- dflash.dflash_swa_window_size=2048
environment:
- MAX_FINAL_LOSS: "5.0"
- MIN_FINAL_ACC: "0.15"
slurm_config:
_factory_: "slurm_factory"
nodes: 1
ntasks_per_node: 1
gpus_per_node: 8
@@ -0,0 +1,124 @@
# DFlash2 streaming speculative decoding pipeline for Qwen3-8B.
#
# Same streaming transport as hf_streaming_dflash.yaml: a live `vllm serve` captures the
# target model's hidden states and moves them to the trainer over NIXL RDMA (no disk
# round-trip). DFlash2 trains the same block-diffusion backbone plus the grouped sublayer
# convolutions and the candidate selector. See common/eagle3/train_eagle_streaming.sh for
# the dispatch and rendezvous.
#
# Two nodes, all four GPUs of each: node 0 is one vLLM replica at TP=4, node 1 is a
# 4-rank DDP trainer. Scale down by lowering gpus_per_node and SERVE_TP together; scale up
# by raising nodes and SERVE_NODES together, as hf_streaming_dflash_multi_node.yaml does.
#
# 3-step pipeline:
# task_0: Build input conversations (jsonl)
# task_1: Streaming train — 1 serve node (TP=4) + 1 trainer node (4-rank DDP)
# task_2: vLLM smoke test with the exported drafter
#
# Site-specific settings left unset on purpose: an aarch64 cluster needs its own image
# (the default is x86), slurm_config.time where the queue's default is short, the NIXL and
# NCCL HCA pinning an InfiniBand fabric needs, and container-visible cache paths (srun
# runs with --no-container-mount-home, and TRITON_CACHE_DIR must be node-local).
#
# Usage:
# uv run launch.py --yaml examples/Qwen/Qwen3-8B/hf_streaming_dflash2.yaml --yes
job_name: Qwen3-8B_DFlash2_streaming
pipeline:
allow_to_fail: false
skip: false
note:
global_vars:
hf_model: /hf-local/Qwen/Qwen3-8B
# Step 1: Build input conversations
task_0:
script: common/eagle3/make_dataset.sh
args:
- -f modules/Model-Optimizer/examples/dataset/example_data_config.yaml
- --full-conversations
slurm_config:
_factory_: "slurm_factory"
nodes: 1
ntasks_per_node: 1
gpus_per_node: 1
container: nvcr.io/nvidia/tensorrt-llm/release:1.3.0rc20
# Step 2: Streaming DFlash2 training — node 0 vllm serve (TP=4), node 1 trainer.
# DFlash2 extracts 5 target layers (build_target_layer_ids(36,5)=[1,9,17,25,33], the
# draft's fc input); vLLM's capture ids are those +1 -> [2,10,18,26,34]. No final layer
# (36): this recipe trains against the hard target, so the base hidden is never read.
# Dropping it must be paired with data.final_aux_is_base_hidden=true below. To train
# with distillation instead, set dflash_lk_loss_type=ce and add 36 back without it.
task_1:
script: common/eagle3/train_eagle_streaming.sh
args:
- --config modules/Model-Optimizer/modelopt_recipes/general/speculative_decoding/dflash2.yaml
- model.model_name_or_path=<<global_vars.hf_model>>
# Streaming never runs the base's transformer layers on the trainer; load only the
# embeddings, final norm and lm_head rather than all 36 layers on every rank.
- model.use_fake_base_for_offline=true
- data.mode=streaming
- data.data_path=/scratchspace/data/train.jsonl
# All five captured planes are aux features; without this the streaming dataset
# peels the last one off as the (unused) KD target and the draft's fc gets
# 4x4096 where it wants 5x4096.
- data.final_aux_is_base_hidden=true
- training.output_dir=/scratchspace/dflash2
- training.training_seq_len=4096
- training.disable_tqdm=true
# The serve prefills the corpus conversation rather than generating, so the
# assistant turn is present; Qwen3-8B's stock template just has no
# {% generation %} tags to locate it. Pass chat_template_train.jinja to mask to it.
- training.answer_only_loss=false
# 4 per device x 4 ranks = global batch 16 sequences = 65,536 tokens/step.
- training.per_device_train_batch_size=4
- training.num_train_epochs=1
- training.max_steps=5000
# dflash2.yaml sets report_to=tensorboard, which hard-fails if tensorboard
# isn't in the serve container; the streaming trainer doesn't need it.
- training.report_to=none
environment:
- HF_MODEL_CKPT: <<global_vars.hf_model>>
# No spaces: nemo_run emits unquoted `export FOO=value`, so spaces would split.
- EAGLE_CAPTURE_IDS: "[2,10,18,26,34]"
# Serve replica nodes (Slurm nodes 0..SERVE_NODES-1); the rest are trainers.
- SERVE_NODES: "1"
- SERVE_TP: "4"
# training_seq_len plus headroom for the single decode step.
- SERVE_MAX_MODEL_LEN: "4608"
- SERVE_MAX_NUM_SEQS: "32"
- SERVE_GPU_MEM_UTIL: "0.9"
# RDMA pool slot capacity in tokens. Must be >= training_seq_len or long prompts
# overflow the slot and the producer silently skips capture. 32 slots is 5.4 GiB of
# pinned host memory against a peak in-flight of 4 ranks x 4 workers.
- HS_MAX_TOKENS: "4608"
- HS_POOL_SLOTS: "32"
- STREAMING_NUM_WORKERS: "4"
# DFlash2 uses a custom modeling file; export must trust remote code.
- EXPORT_EXTRA_ARGS: "--trust_remote_code"
slurm_config:
_factory_: "slurm_factory"
nodes: 2
ntasks_per_node: 1
gpus_per_node: 4
container: vllm/vllm-openai:latest
# Step 3: vLLM smoke test (uses the exported checkpoint from training).
# The method stays "dflash": vLLM has no separate dflash2 method and selects the
# DFlash2 path from the checkpoint's architectures: ["DFlash2DraftModel"].
task_2:
script: common/specdec/vllm_smoke_test.sh
environment:
- HF_MODEL_CKPT: <<global_vars.hf_model>>
- DRAFT_MODEL: /scratchspace/export
- SPEC_METHOD: "dflash"
- NUM_SPEC_TOKENS: "7"
- MIN_ACCEPTANCE_LENGTH: "1.2"
slurm_config:
_factory_: "slurm_factory"
container: vllm/vllm-openai:nightly
nodes: 1
ntasks_per_node: 1
gpus_per_node: 1