Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions .pre-commit-config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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|
Expand Down
4 changes: 4 additions & 0 deletions CHANGELOG.rst
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,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.
Expand Down
43 changes: 43 additions & 0 deletions modelopt/torch/export/plugins/hf_spec_export.py
Original file line number Diff line number Diff line change
Expand Up @@ -605,3 +605,46 @@ def _export_config(self):
}
)
return config


class DFlash2Exporter(DFlashExporter):
"""Draft model exporter for DFlash2 (DFlash backbone + convolutions + selector).

Same z-lab-compatible format as DFlash, plus the DFlash2 weights
(``layers.*.attention_conv.*`` / ``layers.*.mlp_conv.*`` /
``candidate_selector.*``, already captured by the inherited ``dflash_module.``
stripping) and the config fields the SGLang/vLLM ``DFlash2DraftModel`` loader
needs to rebuild them (``conv_kernel_size``, ``conv_group_size``,
``selector_rank``, ``selector_top_k``).

The architecture name is what selects the DFlash2 serving path: a checkpoint
declaring ``DFlashDraftModel`` loads as a plain DFlash draft and would silently
ignore the convolutions and the selector.
"""

def _export_config(self):
"""Extend the DFlash config with the DFlash2 architecture fields."""
config = super()._export_config()
draft_config = self.model.dflash_config

config["architectures"] = ["DFlash2DraftModel"]
# Present because HFDFlash2Model.modify validates them at convert time.
config["dflash_config"].update(
{
"projector_type": getattr(draft_config, "projector_type", "dflash2"),
"conv_kernel_size": draft_config.conv_kernel_size,
"conv_group_size": draft_config.conv_group_size,
"selector_rank": draft_config.selector_rank,
"selector_top_k": draft_config.selector_top_k,
# The published DFlash2 checkpoints carry block_size inside
# dflash_config; the DFlash loader reads it from the top level.
# Emit both so either contract resolves to the same value.
"block_size": config["block_size"],
}
)
# vLLM reads is_causal from the TOP level, and the published DFlash2 checkpoints
# state it explicitly rather than leaving it to be inferred from layer_types.
# Mirror dflash_config.causal, which DFlashExporter writes unconditionally from
# dflash_draft_attention; no parent writes a top-level is_causal.
config["is_causal"] = config["dflash_config"]["causal"]
return config
46 changes: 46 additions & 0 deletions modelopt/torch/speculative/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -366,6 +366,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
Expand Down
10 changes: 9 additions & 1 deletion modelopt/torch/speculative/dflash/conversion.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,12 @@
# ``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
Expand All @@ -59,14 +65,16 @@ 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"):
registry = DFlashDMRegistry
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)
Expand Down
1 change: 1 addition & 0 deletions modelopt/torch/speculative/plugins/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@

with import_plugin("transformers"):
from .hf_dflash import *
from .hf_dflash2 import *
from .hf_domino import *
from .hf_dspark import *
from .hf_eagle import *
Expand Down
31 changes: 30 additions & 1 deletion modelopt/torch/speculative/plugins/hf_dflash.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,7 +73,7 @@

import logging
from pathlib import Path
from typing import Any
from typing import Any, NamedTuple

import torch
import torch.nn.functional as F
Expand Down Expand Up @@ -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."""
Expand Down Expand Up @@ -915,6 +930,7 @@ def _compute_loss(
base_logits=None,
draft_hidden=None,
base_outputs=None,
return_terms=False,
):
"""Compute weighted cross-entropy (or KD) loss and accuracy.

Expand All @@ -927,6 +943,11 @@ def _compute_loss(
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.
Expand Down Expand Up @@ -1021,6 +1042,14 @@ def _compute_loss(
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(
Expand Down
Loading
Loading