diff --git a/docs/source/features/speculative-decoding.md b/docs/source/features/speculative-decoding.md index 5badac3eea6b..8ffb87f68ecf 100644 --- a/docs/source/features/speculative-decoding.md +++ b/docs/source/features/speculative-decoding.md @@ -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. diff --git a/tensorrt_llm/_torch/models/modeling_dflash.py b/tensorrt_llm/_torch/models/modeling_dflash.py index b9693465b6fe..7e1b1061dad8 100644 --- a/tensorrt_llm/_torch/models/modeling_dflash.py +++ b/tensorrt_llm/_torch/models/modeling_dflash.py @@ -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, @@ -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 @@ -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 ( self._dflash2_conv_taps > 0 or self._dflash2_selector_rank > 0 or _is_dflash2_architecture(pretrained_config) @@ -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": @@ -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) @@ -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. @@ -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 @@ -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) @@ -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, diff --git a/tensorrt_llm/_torch/models/modeling_lilicorr.py b/tensorrt_llm/_torch/models/modeling_lilicorr.py new file mode 100644 index 000000000000..290c82ce2827 --- /dev/null +++ b/tensorrt_llm/_torch/models/modeling_lilicorr.py @@ -0,0 +1,488 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""DFlash drafter with LiLiCorr candidate scoring and checkpoint-owned logits.""" + +import torch +import torch.nn.functional as F +from torch import nn + +from tensorrt_llm.logger import logger +from tensorrt_llm.models.modeling_utils import QuantConfig +from tensorrt_llm.quantization.mode import QuantAlgo + +from ..model_config import ModelConfig +from ..modules.gated_mlp import GatedMLP +from ..modules.linear import Linear +from .checkpoints.base_weight_mapper import BaseWeightMapper +from .modeling_dflash import DFlashForCausalLM + + +class LiLiCorrRMSNorm(nn.Module): + def __init__(self, hidden_size: int, eps: float) -> None: + super().__init__() + self.weight = nn.Parameter(torch.ones(hidden_size)) + self.eps = eps + + def forward(self, hidden: torch.Tensor) -> torch.Tensor: + normalized = hidden.float() * torch.rsqrt( + hidden.float().square().mean(-1, keepdim=True) + self.eps + ) + return normalized.to(hidden.dtype) * self.weight + + +class LiLiCorrAttention(nn.Module): + def __init__(self, hidden_size: int, num_heads: int) -> None: + super().__init__() + self.num_heads = num_heads + self.head_dim = hidden_size // num_heads + self.in_proj_weight = nn.Parameter(torch.empty(3 * hidden_size, hidden_size)) + self.in_proj_bias = nn.Parameter(torch.empty(3 * hidden_size)) + self.out_proj = nn.Linear(hidden_size, hidden_size) + + def forward(self, hidden: torch.Tensor, bias: torch.Tensor) -> torch.Tensor: + batch, length, width = hidden.shape + qkv = F.linear(hidden, self.in_proj_weight, self.in_proj_bias) + q, k, v = ( + x.reshape(batch, length, self.num_heads, self.head_dim).transpose(1, 2) + for x in qkv.chunk(3, dim=-1) + ) + attended = F.scaled_dot_product_attention(q, k, v, attn_mask=bias) + return self.out_proj(attended.transpose(1, 2).reshape(batch, length, width)) + + +class LiLiCorrMLP(nn.Module): + def __init__(self, input_size: int, intermediate_size: int, output_size: int) -> None: + super().__init__() + self.up_proj = nn.Linear(input_size, intermediate_size) + self.down_proj = nn.Linear(intermediate_size, output_size) + + def forward(self, hidden: torch.Tensor) -> torch.Tensor: + return self.down_proj(F.silu(self.up_proj(hidden))) + + +class LiLiCorrLayer(nn.Module): + def __init__(self, hidden_size: int, num_heads: int, mlp_ratio: float, eps: float) -> None: + super().__init__() + self.attn_norm = LiLiCorrRMSNorm(hidden_size, eps) + self.attn = LiLiCorrAttention(hidden_size, num_heads) + self.mlp_norm = LiLiCorrRMSNorm(hidden_size, eps) + self.mlp = LiLiCorrMLP(hidden_size, int(hidden_size * mlp_ratio), hidden_size) + + def forward(self, hidden: torch.Tensor, bias: torch.Tensor) -> torch.Tensor: + hidden = hidden + self.attn(self.attn_norm(hidden), bias) + return hidden + self.mlp(self.mlp_norm(hidden)) + + +class LiLiCorrHead(nn.Module): + """Score adjacent candidates using all candidates and the accepted target context.""" + + def __init__( + self, + *, + model_hidden_size: int, + hidden_size: int, + num_layers: int, + num_heads: int, + mlp_ratio: float, + block_size: int, + candidate_topk: int, + factor_dim: int, + rms_norm_eps: float, + vector_eps: float, + logit_scale: float, + ) -> None: + super().__init__() + if ( + min(model_hidden_size, hidden_size, num_layers, num_heads, candidate_topk, factor_dim) + < 1 + ): + raise ValueError("LiLiCorr dimensions must be positive") + if hidden_size % num_heads or block_size < 2 or mlp_ratio <= 0: + raise ValueError("Invalid LiLiCorr attention, block or MLP dimensions") + if not 0 < vector_eps < 0.5 or logit_scale <= 0: + raise ValueError("Invalid LiLiCorr normalization epsilon or logit scale") + self.block_size = block_size + self.candidate_topk = candidate_topk + self.hidden_size = hidden_size + self.vector_eps = vector_eps + self.logit_scale = logit_scale + self.token_proj = ( + nn.Identity() + if model_hidden_size == hidden_size + else nn.Linear(model_hidden_size, hidden_size) + ) + self.pass_hidden_proj = nn.Linear(model_hidden_size, hidden_size) + self.feature_norm = nn.LayerNorm(5) + self.feature_mlp = LiLiCorrMLP(5, hidden_size, hidden_size) + self.slot_embedding = nn.Parameter(torch.zeros(1, block_size - 1, 1, hidden_size)) + self.rank_embedding = nn.Parameter(torch.zeros(1, 1, candidate_topk, hidden_size)) + self.relative_slot_bias = nn.Parameter(torch.zeros(num_heads, 2 * block_size - 1)) + self.same_slot_bias = nn.Parameter(torch.zeros(num_heads)) + self.context_proj = nn.Linear(model_hidden_size, hidden_size) + self.layers = nn.ModuleList( + LiLiCorrLayer(hidden_size, num_heads, mlp_ratio, rms_norm_eps) + for _ in range(num_layers) + ) + self.output_norm = LiLiCorrRMSNorm(hidden_size, rms_norm_eps) + self.anchor_norm = LiLiCorrRMSNorm(hidden_size, rms_norm_eps) + self.factor_input_proj = nn.Linear(3 * hidden_size, hidden_size) + self.out_head = nn.Linear(hidden_size, factor_dim) + self.in_head = nn.Linear(hidden_size, factor_dim) + self.anchor_out_head = nn.Linear(hidden_size, factor_dim) + + def forward( + self, + token_embeddings: torch.Tensor, + candidate_log_probs: torch.Tensor, + pass_hidden: torch.Tensor, + anchor_hidden: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor]: + """Return start [B, C] and transition [B, S-1, C, C] logits. + + Inputs are candidate embeddings [B, S, C, H], vocabulary-normalized + log probabilities [B, S, C], draft states [B, S, H], and the projected + target row that predicted the anchor token [B, H]. S excludes the anchor. + """ + batch, slots, topk = candidate_log_probs.shape + if not 1 <= slots < self.block_size or topk != self.candidate_topk: + raise ValueError("LiLiCorr candidate lattice does not match its trained dimensions") + log_probs = candidate_log_probs.float() + ranks = torch.arange(topk, device=log_probs.device, dtype=torch.float32) + rank_fraction = (ranks / max(topk - 1, 1)).expand_as(log_probs) + is_top1 = (ranks == 0).float().expand_as(log_probs) + features = torch.stack( + ( + log_probs, + log_probs.exp(), + log_probs - log_probs.amax(-1, keepdim=True), + rank_fraction, + is_top1, + ), + dim=-1, + ).to(pass_hidden.dtype) + hidden = self.token_proj(token_embeddings) + self.pass_hidden_proj(pass_hidden).unsqueeze( + -2 + ) + hidden = hidden + self.feature_mlp(self.feature_norm(features)) + hidden = hidden + self.slot_embedding[:, :slots] + self.rank_embedding + hidden = hidden.reshape(batch, slots * topk, self.hidden_size) + slot_ids = torch.arange(slots, device=hidden.device).repeat_interleave(topk) + relative = slot_ids[:, None] - slot_ids[None, :] + bias = self.relative_slot_bias[:, relative + self.block_size - 1] + bias = bias + (relative == 0).unsqueeze(0) * self.same_slot_bias[:, None, None] + for layer in self.layers: + hidden = layer(hidden, bias.to(hidden.dtype).unsqueeze(0)) + hidden = self.output_norm(hidden).reshape(batch, slots, topk, self.hidden_size) + anchor = self.anchor_norm(self.context_proj(anchor_hidden)) + anchor_rows = anchor[:, None, None, :].expand_as(hidden) + factors = F.silu( + self.factor_input_proj(torch.cat((hidden, anchor_rows, hidden * anchor_rows), dim=-1)) + ) + outgoing = F.normalize(self.out_head(factors), dim=-1, eps=self.vector_eps) + incoming = F.normalize(self.in_head(factors), dim=-1, eps=self.vector_eps) + anchor_out = F.normalize(self.anchor_out_head(anchor), dim=-1, eps=self.vector_eps) + start = (anchor_out[:, None, :] * incoming[:, 0]).sum(-1) + pairs = outgoing[:, :-1] @ incoming[:, 1:].transpose(-1, -2) + return start.float() * self.logit_scale, pairs.float() * self.logit_scale + + +def lilicorr_proposal_logits(start: torch.Tensor, pairs: torch.Tensor) -> torch.Tensor: + """Return [B, S, C] proposal rows along the greedy candidate path. + + Subsequent sampling must use these same rows for both drawing and verification. + """ + rows = [start] + previous = start.argmax(-1) + for slot in range(pairs.shape[1]): + row = ( + pairs[:, slot] + .gather(1, previous[:, None, None].expand(-1, 1, pairs.shape[-1])) + .squeeze(1) + ) + rows.append(row) + previous = row.argmax(-1) + return torch.stack(rows, dim=1) + + +def _linear_quant_config(model_config: ModelConfig, name: str) -> QuantConfig | None: + """Resolve checkpoint module metadata, with global exclusions taking precedence.""" + candidates = (name, "model." + name) + config = model_config.quant_config + if config is not None and any( + config.is_module_excluded_from_quantization(n) for n in candidates + ): + return None + if model_config.quant_config_dict is not None: + return next( + ( + model_config.quant_config_dict[n] + for n in candidates + if n in model_config.quant_config_dict + ), + None, + ) + return config + + +def _load_linear( + weights: dict[str, torch.Tensor], + in_features: int, + out_features: int, + dtype: torch.dtype, + *, + bias: bool, + quant_config: QuantConfig | None, + device: torch.device, +) -> nn.Module: + algo = quant_config.quant_algo if quant_config is not None else None + scales = { + None: (), + QuantAlgo.FP8: ("weight_scale", "input_scale"), + QuantAlgo.NVFP4: ("weight_scale", "weight_scale_2", "input_scale"), + QuantAlgo.W4A16_NVFP4: ("weight_scale", "weight_scale_2"), + } + if algo not in scales: + raise NotImplementedError(f"Unsupported LiLiCorr projection quantization: {algo}") + required = {"weight", *scales[algo]} + if bias: + required.add("bias") + missing = required - weights.keys() + if missing: + raise ValueError(f"LiLiCorr linear layer is missing {sorted(missing)}") + if algo is None: + if weights["weight"].dtype not in ( + torch.float16, + torch.bfloat16, + torch.float32, + torch.float64, + ): + raise ValueError("Packed or FP8 LiLiCorr weights require quantization metadata") + layer = nn.Linear(in_features, out_features, bias=bias, dtype=dtype, device=device) + layer.load_state_dict(weights, strict=True) + else: + layer = Linear( + in_features, out_features, bias=bias, dtype=dtype, quant_config=quant_config + ).to(device) + layer.load_weights([weights]) + return layer + + +def _normalize_lilicorr_name(name: str) -> str: + name = name.replace("feature_mlp.0.", "feature_norm.") + name = name.replace("feature_mlp.1.", "feature_mlp.up_proj.") + name = name.replace("feature_mlp.3.", "feature_mlp.down_proj.") + return name.replace(".mlp.0.", ".mlp.up_proj.").replace(".mlp.2.", ".mlp.down_proj.") + + +def _normalize_lilicorr_weights(weights: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]: + """Accept sequential and named-MLP exports without changing their tensor semantics.""" + normalized = {} + for name, tensor in weights.items(): + name = _normalize_lilicorr_name(name) + if name in ("slot_embedding", "rank_embedding") and tensor.ndim == 5: + tensor = tensor.squeeze(1) + if name in normalized: + raise ValueError(f"Duplicate LiLiCorr parameter {name}") + normalized[name] = tensor + return normalized + + +class LiLiCorrForCausalLM(DFlashForCausalLM): + _uses_lilicorr = True + + def __init__( + self, draft_config: ModelConfig, *, dflash_attention_backend: str = "AUTO" + ) -> None: + super().__init__(draft_config, dflash_attention_backend=dflash_attention_backend) + # Gate and up have independently calibrated FP4 scales. Keep their + # packed projections separate so neither calibration is discarded. + for layer_idx, layer in enumerate(self.model.layers): + mlp = layer.mlp + if isinstance(mlp, GatedMLP) and not mlp.split_gate_up: + layer.mlp = GatedMLP( + hidden_size=mlp.hidden_size, + intermediate_size=mlp.intermediate_size, + bias=mlp.gate_up_proj.has_bias, + activation=mlp.activation, + dtype=self.config.torch_dtype, + config=self.model_config, + layer_idx=layer_idx, + split_gate_up=True, + swiglu_limit=mlp.swiglu_limit, + swiglu_alpha=mlp.swiglu_alpha, + swiglu_beta=mlp.swiglu_beta, + ) + for name, module in layer.mlp.named_modules(): + if isinstance(module, Linear): + module.quant_config = _linear_quant_config( + self.model_config, f"layers.{layer_idx}.mlp.{name}" + ) + module._weights_created = False + module.create_weights() + settings = getattr(self.config, "dflash_config", {}) + fields = ( + "hidden_size", + "num_layers", + "num_heads", + "mlp_ratio", + "candidate_topk", + "factor_dim", + "vector_eps", + "logit_scale", + ) + missing = [f"lilicorr_{field}" for field in fields if f"lilicorr_{field}" not in settings] + if missing: + raise ValueError(f"LiLiCorr checkpoint is missing config fields {missing}") + if self.block_size is None: + raise ValueError("LiLiCorr checkpoint requires dflash_config.block_size or block_size") + if getattr(self.config, "is_causal", False) or settings.get("causal", False): + raise ValueError("LiLiCorr requires non-causal draft attention") + self.lilicorr = LiLiCorrHead( + model_hidden_size=self.config.hidden_size, + block_size=self.block_size, + rms_norm_eps=self.config.rms_norm_eps, + **{field: settings[f"lilicorr_{field}"] for field in fields}, + ) + self.has_own_lm_head = bool(getattr(self.config, "has_own_lm_head", False)) + if not 0 <= self._dflash2_conv_taps <= self.block_size: + raise ValueError("LiLiCorr conv_kernel_size must be between 0 and block_size") + if self._dflash2_conv_taps and self._dflash2_conv_group_size < 1: + raise ValueError("LiLiCorr grouped convolution requires conv_group_size > 0") + logger.info( + f"LiLiCorr enabled: block_size={self.block_size}, " + f"candidate_topk={self.lilicorr.candidate_topk}, own_lm_head={self.has_own_lm_head}" + ) + + def load_weights( + self, + weights: dict[str, torch.Tensor], + weight_mapper: BaseWeightMapper | None = None, + **kwargs: object, + ) -> None: + weights = {name.removeprefix("model."): value for name, value in weights.items()} + raw_head_weights = { + name.removeprefix("lilicorr."): value + for name, value in weights.items() + if name.startswith("lilicorr.") + } + head_weights = _normalize_lilicorr_weights(raw_head_weights) + source_names = { + _normalize_lilicorr_name(name): "lilicorr." + name.rsplit(".", 1)[0] + for name in raw_head_weights + } + device = self.model.norm.weight.device + self.lilicorr.to(device=device, dtype=self.config.torch_dtype) + consumed = set() + for name, layer in list(self.lilicorr.named_modules()): + if isinstance(layer, nn.Linear): + prefix = name + "." + values = { + key[len(prefix) :]: value + for key, value in head_weights.items() + if key.startswith(prefix) + } + loaded = _load_linear( + values, + layer.in_features, + layer.out_features, + self.config.torch_dtype, + bias=layer.bias is not None, + quant_config=_linear_quant_config( + self.model_config, source_names.get(name + ".weight", "lilicorr." + name) + ), + device=device, + ) + self.lilicorr.set_submodule(name, loaded) + consumed.update(prefix + key for key in values) + remaining = {name: value for name, value in head_weights.items() if name not in consumed} + expected = { + name + for name, _ in self.lilicorr.named_parameters() + if not any(name.startswith(key.rsplit(".", 1)[0] + ".") for key in consumed) + } + if remaining.keys() != expected: + raise ValueError( + f"LiLiCorr head weight mismatch: missing={sorted(expected - remaining.keys())}, " + f"unexpected={sorted(remaining.keys() - expected)}" + ) + self.lilicorr.load_state_dict(remaining, strict=False) + weights = { + name: value for name, value in weights.items() if not name.startswith("lilicorr.") + } + if self.has_own_lm_head: + head = { + name.removeprefix("lm_head."): value + for name, value in weights.items() + if name.startswith("lm_head.") + } + if "weight" not in head: + raise ValueError("LiLiCorr has_own_lm_head requires checkpoint lm_head weights") + self.lm_head.load_weights([head]) + weights = { + name: value for name, value in weights.items() if not name.startswith("lm_head.") + } + if self._dflash2_conv_taps: + weights = self._load_dflash2_weights(weights, load_selector=False) + super().load_weights(weights, weight_mapper=weight_mapper, **kwargs) + # The backbone loads partially to leave shared embeddings and the output + # head alone. Complete its accumulated quantization scales before use. + for module in self.model.modules(): + if isinstance(module, Linear): + module.process_weights_after_loading() + + def _load_target_projection(self, weights: dict[str, torch.Tensor]) -> None: + projection = { + name.removeprefix("fc."): value + for name, value in weights.items() + if name.startswith("fc.") + } + if "weight" not in projection: + raise ValueError("LiLiCorr target projection is missing fc.weight") + weight = projection["weight"] + quant_config = _linear_quant_config(self.model_config, "fc") + in_features = weight.shape[1] + if ( + weight.dtype == torch.uint8 + and quant_config is not None + and quant_config.quant_algo in (QuantAlgo.NVFP4, QuantAlgo.W4A16_NVFP4) + ): + in_features *= 2 + self.fc = _load_linear( + projection, + in_features, + self.config.hidden_size, + self.config.torch_dtype, + bias=False, + quant_config=quant_config, + device=self.model.norm.weight.device, + ) + for name in projection: + del weights["fc." + name] + + def project_target_hidden(self, hidden_states: torch.Tensor) -> torch.Tensor: + return self.hidden_norm(self.fc(hidden_states.to(self.config.torch_dtype))) + + def load_weights_from_target_model(self, target_model: nn.Module) -> None: + embedding = target_model.model.embed_tokens + if ( + embedding.embedding_dim != self.config.hidden_size + or embedding.num_embeddings != self.config.vocab_size + ): + raise ValueError("LiLiCorr requires matching target embedding width and vocabulary") + self.draft_model_full.model.embed_tokens = embedding + if not self.has_own_lm_head: + self.draft_model_full.lm_head = target_model.lm_head + self.lm_head = target_model.lm_head + + def select_lilicorr_path( + self, + candidate_ids: torch.Tensor, + candidate_log_probs: torch.Tensor, + draft_hidden: torch.Tensor, + anchor_hidden: torch.Tensor, + ) -> torch.Tensor: + embeddings = self.model.embed_tokens(candidate_ids.reshape(-1)).reshape( + *candidate_ids.shape, self.config.hidden_size + ) + start, pairs = self.lilicorr(embeddings, candidate_log_probs, draft_hidden, anchor_hidden) + return lilicorr_proposal_logits(start, pairs) diff --git a/tensorrt_llm/_torch/speculative/dflash.py b/tensorrt_llm/_torch/speculative/dflash.py index d8097dffbd58..a1f9eca1b99a 100644 --- a/tensorrt_llm/_torch/speculative/dflash.py +++ b/tensorrt_llm/_torch/speculative/dflash.py @@ -46,6 +46,27 @@ from ...llmapi.llm_args import DFlashDecodingConfig +def dflash_context_dtype(draft_model: nn.Module) -> torch.dtype: + """Use projection activations for the context cache, not packed weight storage.""" + projection = getattr(draft_model, "fc", None) + if projection is None: + return torch.bfloat16 + return getattr(projection, "dtype", projection.weight.dtype) + + +def last_accepted_hidden(projected_rows: torch.Tensor, num_accepted: torch.Tensor) -> torch.Tensor: + """Select the target row that predicted each newly committed anchor. + + Rows are [requests, verification tokens, hidden], after the target projection + and hidden normalization: hidden_norm(fc(target_hidden)). Accepted counts + include the target's bonus token, so its predicting row is at count minus one. + """ + indices = (num_accepted - 1).clamp_min(0).long() + return projected_rows.gather( + 1, indices[:, None, None].expand(-1, 1, projected_rows.shape[-1]) + ).squeeze(1) + + def compute_dflash_ctx_buffer_bytes( max_batch_size: int, max_ctx_len: int, @@ -328,7 +349,7 @@ def validate_dflash_ctx_buffer_budget( max_ctx = min(_ctx_candidates) if _ctx_candidates else 8192 num_kv_heads_per_rank = (num_kv_heads + tp_size - 1) // tp_size # Config-time approximations of the real allocation's inputs: the lazy - # allocation reads draft_model.fc.weight.dtype and _draft_block_width() + # allocation reads dflash_context_dtype() and _draft_block_width() # (both subclass-overridable), whereas here dtype_bytes is derived from # config["torch_dtype"] and the block width is hardcoded to K + 1. dtype_bytes = 4 if draft_config.get("torch_dtype") in ("float32", "float") else 2 @@ -784,6 +805,8 @@ def set_draft_model(self, draft_model) -> None: "DFlash 2 candidate selection requires a shared draft/target vocab " "(d2t vocab mapping is not supported)." ) + if self._d2t is not None and getattr(draft_model, "lilicorr", None) is not None: + raise NotImplementedError("LiLiCorr requires a shared draft/target vocabulary") def _check_ctx_arena_fits(self, capacity, num_slots, L, nkv, hd, dtype, kv_factor=2): """Fail with the arithmetic before allocating the drafter context arena. @@ -1035,7 +1058,7 @@ def _lazy_init_ctx_buffers( "the binding constraint, so requests past it will draft nothing." ) - dtype = draft_model.fc.weight.dtype if hasattr(draft_model, "fc") else torch.bfloat16 + dtype = dflash_context_dtype(draft_model) # Reserve slot index max_batch as a scratch slot for padding/unknown # dummies; real requests only draw slots 0..max_batch-1, so dummy @@ -1698,7 +1721,16 @@ def _forward_impl( # DFlash 2: replace the independent per-position picks with one # coherent path through the block (absent for plain DFlash). - if getattr(draft_model, "has_candidate_selector", False): + if getattr(draft_model, "lilicorr", None) is not None: + gen_logits = self._apply_lilicorr( + draft_model, + gen_logits, + gen_hidden_states.reshape(num_gens, K, -1), + inputs["lilicorr_anchor"], + spec_metadata, + ) + vocab_size = gen_logits.shape[-1] + elif getattr(draft_model, "has_candidate_selector", False): gen_logits = self._apply_dflash2_selector( draft_model, gen_logits, @@ -1855,6 +1887,38 @@ def _apply_dflash2_selector( anchor_tokens.long(), ) + def _apply_lilicorr( + self, + draft_model, + gen_logits: torch.Tensor, + draft_hidden: torch.Tensor, + anchor_hidden: torch.Tensor, + spec_metadata, + ) -> torch.Tensor: + if anchor_hidden is None: + raise RuntimeError( + "LiLiCorr requires the projected target row that predicted the anchor" + ) + full_vocab = draft_model.config.vocab_size + candidate_ids, unary_logits, block_logits = self._dflash2_global_top_k( + gen_logits, spec_metadata, draft_model.lilicorr.candidate_topk, full_vocab + ) + # Candidate features use probabilities normalized over the entire vocabulary, + # including the mass outside top-k and on other tensor-parallel ranks. + log_normalizer = torch.logsumexp(gen_logits.float(), dim=-1, keepdim=True) + if gen_logits.shape[-1] != full_vocab: + from ..distributed.ops import allgather + + log_normalizer = torch.logsumexp( + allgather(log_normalizer, self.mapping, dim=-1), dim=-1, keepdim=True + ) + proposal = draft_model.select_lilicorr_path( + candidate_ids, unary_logits.float() - log_normalizer, draft_hidden, anchor_hidden + ) + block_logits.fill_(float("-inf")) + block_logits.scatter_(-1, candidate_ids, proposal.to(block_logits.dtype)) + return block_logits + def _dflash2_global_top_k( self, gen_logits: torch.Tensor, @@ -1951,6 +2015,7 @@ def prepare_1st_drafter_inputs( hidden_dim = ( spec_metadata.hidden_size if spec_metadata.hidden_size > 0 else hidden_states.shape[-1] ) + lilicorr_anchor = None if num_gens > 0: gen_num_accepted = num_accepted_tokens[num_contexts : num_contexts + num_gens] @@ -2012,6 +2077,11 @@ def prepare_1st_drafter_inputs( gen_hs = captured_hs[gen_start : gen_start + num_gens * total_tokens_per_req] gen_hs_to_project = gen_hs.reshape(-1, gen_hs.shape[-1]) projected_to_store = draft_model.project_target_hidden(gen_hs_to_project) + if getattr(draft_model, "lilicorr", None) is not None: + # The last accepted output is predicted by the corresponding input + # row; the newly committed token itself has no target row yet. + projected_rows = projected_to_store.reshape(num_gens, K_plus_1, -1) + lilicorr_anchor = last_accepted_hidden(projected_rows, gen_num_accepted) gen_num_accepted_long = gen_num_accepted.long() col_idx = self._ctx_len[slots].unsqueeze(1) + offsets_kp1.unsqueeze(0) write_mask = offsets_kp1.unsqueeze(0) < gen_num_accepted_long.unsqueeze(1) @@ -2092,6 +2162,7 @@ def prepare_1st_drafter_inputs( bonus = torch.empty(0, dtype=torch.long, device="cuda") return { + "lilicorr_anchor": lilicorr_anchor, "noise_embedding": noise_embedding, "query_positions": query_positions, "num_ctx_per_req": num_ctx_per_req_t, diff --git a/tensorrt_llm/usage/architecture_allowlist.py b/tensorrt_llm/usage/architecture_allowlist.py index ead3627edcd8..40b2ccccf036 100644 --- a/tensorrt_llm/usage/architecture_allowlist.py +++ b/tensorrt_llm/usage/architecture_allowlist.py @@ -65,6 +65,7 @@ "KimiK3ForConditionalGeneration", "KimiLinearForCausalLM", "LagunaForCausalLM", + "LiLiCorrDraftModel", "Llama4ForConditionalGeneration", "LlamaForCausalLM", "LlavaLlamaModel", diff --git a/tests/integration/test_lists/test-db/l0_h100.yml b/tests/integration/test_lists/test-db/l0_h100.yml index 31659d64e83c..445a53697858 100644 --- a/tests/integration/test_lists/test-db/l0_h100.yml +++ b/tests/integration/test_lists/test-db/l0_h100.yml @@ -62,6 +62,7 @@ l0_h100: # function's leaves on Hopper, keep the rest of the sampler dir. - unittest/_torch/sampler -k "not test_speculative_d2h_parity_real_predictor" - unittest/_torch/speculative/test_eagle3.py + - unittest/_torch/speculative/test_lilicorr_loading.py - unittest/_torch/speculative/test_fused_sampling_op.py - unittest/_torch/speculative/test_rejection_buffers_guard.py - unittest/_torch/speculative/test_sa_hybrid_state_promotion.py diff --git a/tests/unittest/_torch/speculative/hw_agnostic/test_lilicorr.py b/tests/unittest/_torch/speculative/hw_agnostic/test_lilicorr.py new file mode 100644 index 000000000000..b07570003159 --- /dev/null +++ b/tests/unittest/_torch/speculative/hw_agnostic/test_lilicorr.py @@ -0,0 +1,463 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from types import SimpleNamespace +from unittest.mock import patch + +import pytest +import torch +from torch import nn + +from tensorrt_llm._torch.model_config import ModelConfig +from tensorrt_llm._torch.models import modeling_dflash as dflash +from tensorrt_llm._torch.models import modeling_lilicorr as lili +from tensorrt_llm._torch.models.modeling_dflash import DFlashForCausalLM +from tensorrt_llm._torch.speculative.dflash import ( + DFlashWorker, + dflash_context_dtype, + last_accepted_hidden, +) +from tensorrt_llm.models.modeling_utils import QuantConfig +from tensorrt_llm.quantization.mode import QuantAlgo + +HEAD_CONFIG = dict( + model_hidden_size=32, + hidden_size=8, + num_layers=2, + num_heads=2, + mlp_ratio=2, + block_size=4, + candidate_topk=3, + factor_dim=8, + rms_norm_eps=1e-5, + vector_eps=1e-4, + logit_scale=8, +) + + +def _head() -> lili.LiLiCorrHead: + head = lili.LiLiCorrHead(**HEAD_CONFIG).eval() + for name, parameter in head.named_parameters(): + values = torch.arange(parameter.numel(), dtype=torch.float32).reshape(parameter.shape) + parameter.data.copy_(torch.sin(values * 0.17 + sum(name.encode()) * 0.013) * 0.2) + return head + + +def test_head_matches_reference_scores() -> None: + """Compare all start/transition potentials with a frozen Model Optimizer reference.""" + head = _head() + embeddings = torch.linspace(-1.2, 1.3, 512).reshape(16, 32) + ids = torch.tensor([[[2, 5, 9], [6, 3, 1], [4, 7, 8]]]) + logs = torch.tensor([[[-0.4, -1.5, -3.0], [-0.6, -1.7, -2.7], [-0.2, -2.1, -3.4]]]) + hidden = torch.linspace(-0.7, 0.9, 96).reshape(1, 3, 32) + anchor = torch.linspace(-0.5, 0.4, 32).reshape(1, 32) + start, pairs = head(embeddings[ids], logs, hidden, anchor) + expected_start = torch.tensor([[-2.00565147, -2.00631571, -2.01010752]]) + expected_pairs = torch.tensor( + [ + [ + [ + [1.16437435, 1.16812551, 1.15701485], + [1.16656375, 1.17030466, 1.15921211], + [1.17068172, 1.17445230, 1.16334319], + ], + [ + [1.14806044, 1.14459574, 1.12698877], + [1.14481199, 1.14134955, 1.12373281], + [1.14475155, 1.14125943, 1.12357390], + ], + ] + ] + ) + torch.testing.assert_close(start, expected_start, atol=2e-5, rtol=2e-5) + torch.testing.assert_close(pairs, expected_pairs, atol=2e-5, rtol=2e-5) + + +def test_candidate_walk_follows_the_selected_predecessor() -> None: + start = torch.tensor([[1.0, 4.0], [5.0, 1.0]]) + pairs = torch.tensor( + [ + [[[8.0, 0.0], [2.0, 7.0]], [[4.0, 3.0], [6.0, 1.0]]], + [[[1.0, 9.0], [8.0, 0.0]], [[7.0, 2.0], [3.0, 8.0]]], + ] + ) + expected = torch.tensor( + [[[1.0, 4.0], [2.0, 7.0], [6.0, 1.0]], [[5.0, 1.0], [1.0, 9.0], [3.0, 8.0]]] + ) + torch.testing.assert_close(lili.lilicorr_proposal_logits(start, pairs), expected) + torch.testing.assert_close(lili.lilicorr_proposal_logits(start, pairs[:, :0]), start[:, None]) + + +@pytest.mark.parametrize("tp_size", [1, 4]) +def test_worker_normalizes_full_vocabulary(tp_size: int) -> None: + torch.manual_seed(21) + logits = torch.randn(2, 3, 32) + shards = logits.chunk(tp_size, dim=-1) + expected_values, expected_ids = logits.topk(3, dim=-1) + for rank, shard in enumerate(shards): + worker = DFlashWorker.__new__(DFlashWorker) + worker.mapping = SimpleNamespace(tp_size=tp_size, tp_rank=rank, enable_attention_dp=False) + worker._d2t = None + + def gather(local: torch.Tensor, mapping: object, dim: int) -> torch.Tensor: + if local.shape[-1] == 1: + payloads = [x.logsumexp(-1, keepdim=True) for x in shards] + else: + payloads = [] + for index, values in enumerate(shards): + top, ids = values.topk(3, dim=-1) + payloads.append( + torch.stack(((ids + index * shard.shape[-1]).float(), top), -1).flatten(-2) + ) + torch.testing.assert_close(local, payloads[rank]) + return torch.cat(payloads, dim=dim) + + def select( + ids: torch.Tensor, probs: torch.Tensor, hidden: torch.Tensor, anchor: torch.Tensor + ) -> torch.Tensor: + torch.testing.assert_close(ids, expected_ids) + torch.testing.assert_close(probs, logits.log_softmax(-1).gather(-1, ids)) + return expected_values + + model = SimpleNamespace( + config=SimpleNamespace(vocab_size=32), + lilicorr=SimpleNamespace(candidate_topk=3), + select_lilicorr_path=select, + ) + with patch("tensorrt_llm._torch.distributed.ops.allgather", side_effect=gather): + actual = worker._apply_lilicorr( + model, + shard.clone(), + torch.zeros(2, 3, 8), + torch.zeros(2, 8), + SimpleNamespace(draft_vocab_size=32, vocab_size=32), + ) + torch.testing.assert_close(actual.gather(-1, expected_ids), expected_values) + assert (actual > -torch.inf).sum() == expected_ids.numel() + + +def test_projection_cache_uses_activation_dtype() -> None: + packed = SimpleNamespace(weight=torch.zeros(8, 8, dtype=torch.uint8), dtype=torch.bfloat16) + assert dflash_context_dtype(SimpleNamespace(fc=packed)) == torch.bfloat16 + assert dflash_context_dtype(SimpleNamespace(fc=nn.Linear(8, 8))) == torch.float32 + + +def test_partial_acceptance_selects_normalized_target_row() -> None: + model = lili.LiLiCorrForCausalLM.__new__(lili.LiLiCorrForCausalLM) + nn.Module.__init__(model) + model.config = SimpleNamespace(torch_dtype=torch.float32) + model.fc = nn.Linear(64, 32, bias=False) + model.fc.weight.data.copy_(torch.linspace(-0.3, 0.5, 2048).reshape(32, 64)) + model.hidden_norm = lili.LiLiCorrRMSNorm(32, 1e-5) + model.hidden_norm.weight.data.copy_(torch.linspace(0.5, 1.5, 32)) + captured = torch.sin(torch.arange(3 * 4 * 64).float() * 0.17).reshape(3, 4, 64) + counts = torch.tensor([1, 3, 4]) + + projected = model.project_target_hidden(captured) + actual = last_accepted_hidden(projected, counts) + raw = captured[torch.arange(3), counts - 1] @ model.fc.weight.T + expected = raw * torch.rsqrt(raw.square().mean(-1, keepdim=True) + 1e-5) + expected = expected * model.hidden_norm.weight + torch.testing.assert_close(actual, expected) + assert not torch.allclose(actual, raw) + + +def test_worker_rejects_unsupported_layout_before_collectives() -> None: + worker = DFlashWorker.__new__(DFlashWorker) + worker.mapping = SimpleNamespace(tp_size=2, tp_rank=0, enable_attention_dp=False) + worker._d2t = None + model = SimpleNamespace( + config=SimpleNamespace(vocab_size=16), lilicorr=SimpleNamespace(candidate_topk=3) + ) + logits = torch.zeros(1, 3, 9) + with patch("tensorrt_llm._torch.distributed.ops.allgather") as gather: + with pytest.raises(NotImplementedError, match="plain TP column shard"): + worker._apply_lilicorr( + model, + logits, + torch.zeros(1, 3, 32), + torch.zeros(1, 32), + SimpleNamespace(draft_vocab_size=16, vocab_size=16), + ) + gather.assert_not_called() + + +@pytest.fixture +def lightweight_backbone(monkeypatch: pytest.MonkeyPatch) -> None: + class Backbone(nn.Module): + def __init__(self, config: ModelConfig) -> None: + super().__init__() + self.model = nn.Module() + self.model.layers = nn.ModuleList() + self.model.norm = nn.LayerNorm(32) + self.lm_head = nn.Linear(32, 16, bias=False) + + monkeypatch.setattr(dflash, "get_model_architecture", lambda config: (Backbone, None)) + monkeypatch.setattr(dflash, "get_dflash_flash_attention", lambda: None) + + +def _draft_config() -> ModelConfig: + return ModelConfig( + pretrained_config=SimpleNamespace( + architectures=["LiLiCorrDraftModel"], + hidden_size=32, + vocab_size=16, + torch_dtype=torch.float32, + rms_norm_eps=1e-5, + num_hidden_layers=0, + dflash_config={"block_size": 4, **{"lilicorr_" + k: v for k, v in HEAD_CONFIG.items()}}, + ) + ) + + +@pytest.mark.usefixtures("lightweight_backbone") +@pytest.mark.parametrize("marker", ["architecture", "projector_type", "lilicorr_enabled", "plain"]) +def test_checkpoint_markers_route_to_the_correct_drafter(marker: str) -> None: + config = _draft_config() + if marker == "architecture": + config.pretrained_config.block_size = config.pretrained_config.dflash_config.pop( + "block_size" + ) + config.pretrained_config.dflash_config.update(conv_kernel_size=2, conv_group_size=8) + else: + config.pretrained_config.architectures = ["Qwen3ForCausalLM"] + if marker == "projector_type": + config.pretrained_config.dflash_config[marker] = "lilicorr" + elif marker == "lilicorr_enabled": + config.pretrained_config.dflash_config[marker] = True + assert dflash.declares_lilicorr(config.pretrained_config) == (marker != "plain") + model = dflash._build_dflash_draft( + SimpleNamespace(spec_config=SimpleNamespace(attention_backend="VANILLA")), + config, + None, + None, + ) + expected_type = DFlashForCausalLM if marker == "plain" else lili.LiLiCorrForCausalLM + assert type(model) is expected_type + assert not model.is_dflash2 + if marker == "architecture": + assert model.block_size == 4 + assert model._dflash2_conv_taps == 2 and model.candidate_selector is None + + +def test_laguna_lilicorr_is_rejected() -> None: + config = _draft_config() + config.pretrained_config.architectures = ["LagunaForCausalLM", "LiLiCorrDraftModel"] + with pytest.raises(NotImplementedError, match="generic GQA DFlash backbone"): + dflash._build_dflash_draft( + SimpleNamespace(spec_config=SimpleNamespace(attention_backend="VANILLA")), + config, + None, + None, + ) + + +@pytest.mark.usefixtures("lightweight_backbone") +@pytest.mark.parametrize("missing", ["dflash_config", "block_size"]) +def test_missing_checkpoint_settings_have_clear_errors(missing: str) -> None: + config = _draft_config() + if missing == "dflash_config": + del config.pretrained_config.dflash_config + message = "missing config fields" + else: + del config.pretrained_config.dflash_config["block_size"] + message = "requires dflash_config.block_size or block_size" + with pytest.raises(ValueError, match=message): + lili.LiLiCorrForCausalLM(config, dflash_attention_backend="VANILLA") + + +@pytest.mark.usefixtures("lightweight_backbone") +@pytest.mark.parametrize( + "settings, message", + [ + ({"conv_kernel_size": -1}, "conv_kernel_size"), + ({"conv_kernel_size": 5}, "conv_kernel_size"), + ({"conv_kernel_size": 2, "conv_group_size": 0}, "conv_group_size"), + ({"causal": True}, "non-causal"), + ], +) +def test_invalid_convolution_and_causal_settings(settings: dict, message: str) -> None: + config = _draft_config() + config.pretrained_config.dflash_config.update(settings) + with pytest.raises(ValueError, match=message): + lili.LiLiCorrForCausalLM(config, dflash_attention_backend="VANILLA") + + +@pytest.mark.usefixtures("lightweight_backbone") +@pytest.mark.parametrize("is_causal", [None, False, True]) +def test_lilicorr_attention_is_noncausal_with_symmetric_windows(is_causal: bool | None) -> None: + config = _draft_config() + config.pretrained_config.is_causal = is_causal + config.pretrained_config.layer_types = ["sliding_attention", "full_attention"] + config.pretrained_config.sliding_window = 8 + config.pretrained_config.use_sliding_window = True + if is_causal: + with pytest.raises(ValueError, match="non-causal"): + lili.LiLiCorrForCausalLM(config, dflash_attention_backend="VANILLA") + return + model = lili.LiLiCorrForCausalLM(config, dflash_attention_backend="VANILLA") + assert model._get_attention_mask_args(0) == (False, (7, 7)) + assert model._get_attention_mask_args(1) == (False, (-1, -1)) + + +@pytest.mark.usefixtures("lightweight_backbone") +@pytest.mark.parametrize("algo", [None, QuantAlgo.NVFP4, QuantAlgo.W4A16_NVFP4]) +def test_projection_width_comes_from_checkpoint(algo: QuantAlgo | None) -> None: + config = _draft_config() + config.quant_config = QuantConfig(quant_algo=algo) + model = lili.LiLiCorrForCausalLM(config, dflash_attention_backend="VANILLA") + assert model.target_layer_ids is None + packed = algo in (QuantAlgo.NVFP4, QuantAlgo.W4A16_NVFP4) + weight = ( + torch.zeros(32, 48, dtype=torch.uint8) + if packed + else torch.linspace(-0.5, 0.5, 32 * 96).reshape(32, 96) + ) + weights = {"fc.weight": weight} + if algo is None: + model._load_target_projection(weights) + assert model.fc.in_features == 96 + rows = torch.randn(2, 96) + torch.testing.assert_close(model.fc(rows), rows @ weight.T) + else: + with patch.object(lili, "_load_linear", return_value=nn.Identity()) as load: + model._load_target_projection(weights) + assert load.call_args.args[1:3] == (96, 32) + assert not weights + + +def test_checkpoint_loading_preserves_metadata_scales_and_own_head( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Mock native GEMMs; exercise checkpoint routing, finalization and sharing.""" + + class RecordingLinear(nn.Module): + def __init__( + self, *args: object, quant_config: QuantConfig | None = None, **kwargs: object + ) -> None: + super().__init__() + self.quant_config, self.loaded, self.scale = quant_config, {}, None + + def create_weights(self) -> None: + pass + + def load_weights(self, weights: list[dict]) -> None: + self.loaded = weights[0] + + def process_weights_after_loading(self) -> None: + self.scale = self.loaded.get("weight_scale_2") + + class MLP(nn.Module): + def __init__(self, **kwargs: object) -> None: + super().__init__() + self.__dict__.update( + dict( + hidden_size=32, + intermediate_size=64, + split_gate_up=False, + activation=torch.nn.functional.silu, + layer_idx=None, + swiglu_limit=None, + swiglu_alpha=None, + swiglu_beta=None, + ) + | kwargs + ) + self.gate_up_proj = SimpleNamespace(has_bias=False) + if self.split_gate_up: + self.gate_proj, self.up_proj = RecordingLinear(), RecordingLinear() + + def initialize(self: nn.Module, config: ModelConfig, **kwargs: object) -> None: + nn.Module.__init__(self) + self.model_config, self.config = config, config.pretrained_config + self.target_layer_ids, self.block_size, self._dflash2_conv_taps = [0, 1], 4, 0 + self.model = nn.Module() + self.model.norm = nn.LayerNorm(32) + self.model.layers = nn.ModuleList([nn.Module(), nn.Module()]) + for layer in self.model.layers: + layer.mlp = MLP() + self.lm_head = RecordingLinear() + self.draft_model_full = SimpleNamespace(model=self.model, lm_head=self.lm_head) + + def load_backbone(self: nn.Module, weights: dict, **kwargs: object) -> None: + self._load_target_projection(weights) + for layer_idx, layer in enumerate(self.model.layers): + for name in ("gate", "up"): + getattr(layer.mlp, name + "_proj").load_weights( + [ + { + "weight_scale_2": weights[ + f"layers.{layer_idx}.mlp.{name}_proj.weight_scale_2" + ] + } + ] + ) + + monkeypatch.setattr(lili, "Linear", RecordingLinear) + monkeypatch.setattr(lili, "GatedMLP", MLP) + monkeypatch.setattr(DFlashForCausalLM, "__init__", initialize) + monkeypatch.setattr(DFlashForCausalLM, "load_weights", load_backbone) + config = ModelConfig( + pretrained_config=SimpleNamespace( + hidden_size=32, + vocab_size=16, + torch_dtype=torch.bfloat16, + rms_norm_eps=1e-5, + has_own_lm_head=True, + dflash_config={"lilicorr_" + k: v for k, v in HEAD_CONFIG.items()}, + ), + quant_config=QuantConfig( + quant_algo=QuantAlgo.W4A16_NVFP4, + group_size=16, + exclude_modules=["lilicorr.feature_mlp*"], + ), + quant_config_dict={ + "fc": QuantConfig(quant_algo=QuantAlgo.W4A16_NVFP4, group_size=32), + "lilicorr.feature_mlp.1": QuantConfig(quant_algo=QuantAlgo.FP8), + "layers.0.mlp.gate_proj": QuantConfig(quant_algo=QuantAlgo.W4A16_NVFP4, group_size=32), + "layers.0.mlp.up_proj": QuantConfig(quant_algo=QuantAlgo.W4A16_NVFP4, group_size=16), + "layers.1.mlp.gate_proj": QuantConfig(quant_algo=QuantAlgo.W4A16_NVFP4, group_size=16), + "layers.1.mlp.up_proj": QuantConfig(quant_algo=QuantAlgo.W4A16_NVFP4, group_size=32), + }, + ) + model = lili.LiLiCorrForCausalLM(config) + source_head = _head().state_dict() + weights = { + "lilicorr." + + k.replace("feature_norm.", "feature_mlp.0.").replace( + "feature_mlp.up_proj.", "feature_mlp.1." + ): v.unsqueeze(1) if k in ("slot_embedding", "rank_embedding") else v + for k, v in source_head.items() + } + weights.update( + { + "fc.weight": torch.zeros(32, 32, dtype=torch.uint8), + "fc.weight_scale": torch.ones(32, 2).to(torch.float8_e4m3fn), + "fc.weight_scale_2": torch.tensor(0.125), + "lm_head.weight": torch.ones(16, 32), + "layers.0.mlp.gate_proj.weight_scale_2": torch.tensor(0.25), + "layers.0.mlp.up_proj.weight_scale_2": torch.tensor(0.5), + "layers.1.mlp.gate_proj.weight_scale_2": torch.tensor(0.75), + "layers.1.mlp.up_proj.weight_scale_2": torch.tensor(1.0), + } + ) + model.load_weights(weights) + assert model.fc.quant_config.group_size == 32 + assert isinstance(model.lilicorr.feature_mlp.up_proj, nn.Linear) + for name, value in model.lilicorr.state_dict().items(): + torch.testing.assert_close(value, source_head[name].to(torch.bfloat16)) + mlp = model.model.layers[0].mlp + assert mlp.layer_idx == 0 + assert mlp.split_gate_up and mlp.gate_proj.scale == 0.25 and mlp.up_proj.scale == 0.5 + assert mlp.gate_proj.quant_config.group_size == 32 and mlp.up_proj.quant_config.group_size == 16 + second_mlp = model.model.layers[1].mlp + assert second_mlp.layer_idx == 1 + assert second_mlp.gate_proj.quant_config.group_size == 16 + assert second_mlp.up_proj.quant_config.group_size == 32 + own_head = model.lm_head + embedding = nn.Embedding(16, 32) + model.load_weights_from_target_model( + SimpleNamespace(model=SimpleNamespace(embed_tokens=embedding), lm_head=object()) + ) + assert model.model.embed_tokens is embedding and model.lm_head is own_head + torch.testing.assert_close(own_head.loaded["weight"], weights["lm_head.weight"]) diff --git a/tests/unittest/_torch/speculative/test_lilicorr_loading.py b/tests/unittest/_torch/speculative/test_lilicorr_loading.py new file mode 100644 index 000000000000..265ccb4683ce --- /dev/null +++ b/tests/unittest/_torch/speculative/test_lilicorr_loading.py @@ -0,0 +1,35 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import pytest +import torch + +from tensorrt_llm._torch.models.modeling_lilicorr import _load_linear +from tensorrt_llm._torch.modules.linear import Linear +from tensorrt_llm.models.modeling_utils import QuantConfig +from tensorrt_llm.quantization.mode import QuantAlgo + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="FP8 loading requires CUDA") +def test_fp8_projection_loads_real_weights_and_scales() -> None: + weight = torch.linspace(-1, 1, 64 * 32, device="cuda").reshape(32, 64) + weights = { + "weight": weight.to(torch.float8_e4m3fn), + "weight_scale": torch.tensor(0.25, device="cuda"), + "input_scale": torch.tensor(0.5, device="cuda"), + } + layer = _load_linear( + weights, + 64, + 32, + torch.bfloat16, + bias=False, + quant_config=QuantConfig(quant_algo=QuantAlgo.FP8), + device=torch.device("cuda"), + ) + assert isinstance(layer, Linear) + assert layer.weight.shape == (32, 64) + torch.testing.assert_close(layer.weight.float(), weights["weight"].float()) + torch.testing.assert_close(layer.weight_scale, weights["weight_scale"]) + torch.testing.assert_close(layer.input_scale, weights["input_scale"]) + torch.testing.assert_close(layer.inv_input_scale, 1 / weights["input_scale"])