Skip to content
Open
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
21 changes: 21 additions & 0 deletions docs/source/features/speculative-decoding.md
Original file line number Diff line number Diff line change
Expand Up @@ -168,6 +168,27 @@ llm = LLM("/path/to/target_model", speculative_config=speculative_config)

[DFlash 2](https://inco.ai/blog/dflash2/) is also supported. The same `DFlashDecodingConfig` can be used for DFlash 2; no extra arguments are required.

[LiLiCorr](https://arxiv.org/abs/2608.20530) checkpoints use the same
`DFlashDecodingConfig`. The draft checkpoint selects LiLiCorr through
`architectures: ["LiLiCorrDraftModel"]`, `dflash_config.projector_type: "lilicorr"`,
or `dflash_config.lilicorr_enabled: true`.
Its candidate scorer correlates the per-position top-k proposals before target
verification. Load the matching target separately through `LLM(model=...)` and
set `speculative_model` to the drafter. Both models must share the token vocabulary
and embedding width; the captured target layers must match the drafter's training.
The drafter currently uses the generic GQA DFlash backbone; MLA and Laguna-specific
draft layers require separate adapters.

Set `max_draft_len` to `block_size - 1`; for an eight-position checkpoint, use seven
draft tokens. Projection quantization follows the checkpoint's global or per-module
metadata, including excluded modules. The loader accepts dense, FP8, NVFP4 and
W4A16 NVFP4 projections. It preserves a
checkpoint-owned output head when `has_own_lm_head` is set. Grouped
convolutions are loaded when declared by the checkpoint. Draft attention is
non-causal, including symmetric windows on sliding-attention layers; causal
LiLiCorr checkpoints are rejected. The draft attention backend retains the
hardware restrictions listed above.

### User-provided drafting
A completely user-defined drafting method can be supplied with a `UserProvidedDecodingConfig` that includes
* `max_draft_len`: Maximum draft candidate length.
Expand Down
77 changes: 49 additions & 28 deletions tensorrt_llm/_torch/models/modeling_dflash.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,6 +74,16 @@ def _is_dflash2_architecture(config: PretrainedConfig) -> bool:
)


def declares_lilicorr(config: PretrainedConfig) -> bool:
"""Whether the checkpoint requires the LiLiCorr candidate scorer."""
settings = getattr(config, "dflash_config", {})
return bool(
settings.get("projector_type") == "lilicorr"
or settings.get("lilicorr_enabled", False)
or "LiLiCorrDraftModel" in (getattr(config, "architectures", None) or [])
)


def dflash2_grouped_conv(
hidden_states: torch.Tensor,
delta: torch.Tensor,
Expand Down Expand Up @@ -336,6 +346,7 @@ class DFlashForCausalLM(nn.Module):
# sets were built for, while an MLA drafter runs its own block decode and
# has a third implementation they cannot express.
_default_attention_backend = "VANILLA"
_uses_lilicorr = False
_supported_attention_backends = ("VANILLA", "TRTLLM", "FA4")
# Where AUTO lands when the preferred backend cannot run. None means AUTO
# propagates the reason instead: a family whose only deployable target can
Expand Down Expand Up @@ -537,7 +548,8 @@ def __init__(self, draft_config, *, dflash_attention_backend: str = "AUTO"):
self._dflash2_conv_group_size = int(dflash_config.get("conv_group_size", 0) or 0)
self._dflash2_selector_rank = int(dflash_config.get("selector_rank", 0) or 0)
self._dflash2_selector_top_k = int(dflash_config.get("selector_top_k", 0) or 0)
self._is_dflash2 = (
# LiLiCorr supplies its own scorer and permits convolution-free checkpoints.
self._is_dflash2 = not self._uses_lilicorr and (

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Forcing _is_dflash2 = False for LiLiCorr has two side effects beyond skipping the selector:

  • _validate_dflash2_config() is skipped, so a LiLiCorr checkpoint with conv_kernel_size > 0 only gets the conv_group_size check added in LiLiCorrForCausalLM.__init__.
  • explicit_causal at line 1294 is only consulted when _is_dflash2 is true, so a LiLiCorr checkpoint derived from a DFlash 2 drafter that sets top-level is_causal silently falls back to the layer-type default.

If LiLiCorr checkpoints are always non-causal and never ship is_causal, a one-line comment here saying so would prevent the next reader from re-deriving it. Otherwise consider keeping _is_dflash2 as computed and gating only the selector load/use on _uses_lilicorr.

self._dflash2_conv_taps > 0
or self._dflash2_selector_rank > 0
or _is_dflash2_architecture(pretrained_config)
Expand Down Expand Up @@ -832,7 +844,7 @@ def load_weights(self, weights: Dict, weight_mapper=None, **kwargs):
# Remap: add 'model.' prefix where needed, and extract DFlash-specific weights
remapped = {}
for key, value in weights.items():
if key in ("fc.weight", "hidden_norm.weight"):
if key.startswith("fc.") or key == "hidden_norm.weight":
# DFlash-specific projection weights - store directly
remapped[key] = value
elif key == "norm.weight":
Expand All @@ -854,16 +866,7 @@ def load_weights(self, weights: Dict, weight_mapper=None, **kwargs):
)

# Load DFlash-specific weights directly
if "fc.weight" in remapped:
self.fc = nn.Linear(
remapped["fc.weight"].shape[1],
remapped["fc.weight"].shape[0],
bias=False,
device="cuda",
dtype=remapped["fc.weight"].dtype,
)
self.fc.weight.data.copy_(remapped["fc.weight"])
del remapped["fc.weight"]
self._load_target_projection(remapped)

if "hidden_norm.weight" in remapped:
rms_norm_eps = getattr(self.config, "rms_norm_eps", 1e-6)
Expand All @@ -885,7 +888,16 @@ def load_weights(self, weights: Dict, weight_mapper=None, **kwargs):
weights=remapped, weight_mapper=weight_mapper, allow_partial_loading=True
)

def _load_dflash2_weights(self, weights: Dict) -> Dict:
def _load_target_projection(self, weights: Dict) -> None:
if "fc.weight" not in weights:
return
weight = weights.pop("fc.weight")
self.fc = nn.Linear(
weight.shape[1], weight.shape[0], bias=False, device="cuda", dtype=weight.dtype
)
self.fc.weight.data.copy_(weight)

def _load_dflash2_weights(self, weights: Dict, *, load_selector: bool = True) -> Dict:
"""Build the DFlash 2 convolutions and candidate selector.

Returns ``weights`` minus the keys consumed here.
Expand Down Expand Up @@ -918,18 +930,17 @@ def take(key: str) -> torch.Tensor:
)
)

selector = DFlash2CandidateSelector(
predecessor_codebook=take("candidate_selector.predecessor_codebook"),
successor_codebook=take("candidate_selector.successor_codebook"),
hidden_projection_weight=take("candidate_selector.hidden_projection.weight"),
top_k=self._dflash2_selector_top_k,
vocab_size=self.config.vocab_size,
rank=self._dflash2_selector_rank,
)

self.attention_convs = attention_convs
self.mlp_convs = mlp_convs
self.candidate_selector = selector
if load_selector:
self.candidate_selector = DFlash2CandidateSelector(
predecessor_codebook=take("candidate_selector.predecessor_codebook"),
successor_codebook=take("candidate_selector.successor_codebook"),
hidden_projection_weight=take("candidate_selector.hidden_projection.weight"),
top_k=self._dflash2_selector_top_k,
vocab_size=self.config.vocab_size,
rank=self._dflash2_selector_rank,
)
return {k: v for k, v in weights.items() if k not in consumed}

#: Tensors the wrapper itself owns: not in draft_model_full, built from the
Expand Down Expand Up @@ -1277,12 +1288,13 @@ def _get_attention_mask_args(self, layer_idx):
sliding_window = get_layer_attention_window(self.config, layer_idx)
is_sliding_layer = is_sliding_layer or sliding_window is not None

# A DFlash 2 checkpoint's top-level is_causal wins over the layer-type
# default (which would read an all-sliding drafter as causal). Only
# consulted for DFlash 2: is_causal is common enough elsewhere that
# honoring it everywhere could silently retarget an older drafter.
# LiLiCorr uses non-causal block attention with symmetric sliding windows.
# DFlash 2 honors explicit is_causal metadata when present; otherwise,
# causality follows the layer-type default.
explicit_causal = getattr(self.config, "is_causal", None) if self._is_dflash2 else None
if explicit_causal is not None:
if self._uses_lilicorr:
causal = False
elif explicit_causal is not None:
causal = bool(explicit_causal)
elif not is_sliding_layer:
return False, (-1, -1)
Expand Down Expand Up @@ -2060,6 +2072,15 @@ def _build_dflash_draft(model_config, draft_config, lm_head, model):
)
draft_arches = getattr(draft_config.pretrained_config, "architectures", None) or []
dflash_attention_backend = model_config.spec_config.attention_backend
if declares_lilicorr(draft_config.pretrained_config):
if any("Laguna" in arch for arch in draft_arches):
raise NotImplementedError(
"LiLiCorr currently requires the generic GQA DFlash backbone; "
"Laguna's draft-layer specialization is not supported."
)
from .modeling_lilicorr import LiLiCorrForCausalLM

return LiLiCorrForCausalLM(draft_config, dflash_attention_backend=dflash_attention_backend)
if any("Laguna" in arch for arch in draft_arches):
return DFlashLagunaForCausalLM(
draft_config,
Expand Down
Loading
Loading