mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
[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:
@@ -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|
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user