diff --git a/modelopt/torch/speculative/config.py b/modelopt/torch/speculative/config.py index c3f42a5ad14..18a841c25f2 100644 --- a/modelopt/torch/speculative/config.py +++ b/modelopt/torch/speculative/config.py @@ -146,7 +146,19 @@ class DFlashConfig(ModeloptBaseConfig): dflash_use_torch_compile: bool = ModeloptField( default=True, - description="Whether to use torch.compile on DFlash forward/loss methods.", + description=( + "Whether to torch.compile the draft decoder stack and the DSpark TVD chunk in " + "training. Compiles once, on the first step." + ), + ) + + dflash_use_flex_attention: bool = ModeloptField( + default=False, + description=( + "Compute the draft's attention with FlexAttention over a block-sparse BlockMask " + "instead of SDPA with a dense [B, 1, Q, KV] mask, skipping fully-masked tiles. " + "Requires torch >= 2.5." + ), ) dflash_swa_window_size: int | None = ModeloptField( diff --git a/modelopt/torch/speculative/dflash/dflash_model.py b/modelopt/torch/speculative/dflash/dflash_model.py index 3ce06afeda8..e4cb8753668 100644 --- a/modelopt/torch/speculative/dflash/dflash_model.py +++ b/modelopt/torch/speculative/dflash/dflash_model.py @@ -52,4 +52,5 @@ def modify(self, config): self.dflash_draft_attention = config.dflash_draft_attention self.dflash_attention_sink = config.dflash_attention_sink self.dflash_init_checkpoint = config.dflash_init_checkpoint + self.dflash_use_flex_attention = config.dflash_use_flex_attention self.dflash_export_rope_scaling = config.dflash_export_rope_scaling diff --git a/modelopt/torch/speculative/plugins/dflash_flex_attention.py b/modelopt/torch/speculative/plugins/dflash_flex_attention.py new file mode 100644 index 00000000000..2bacc8e1984 --- /dev/null +++ b/modelopt/torch/speculative/plugins/dflash_flex_attention.py @@ -0,0 +1,128 @@ +# 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. + +"""Block-sparse FlexAttention path for the DFlash/DSpark draft. + +Query block ``b`` with anchor ``a_b`` sees only the context prefix ``kv < a_b`` and its own +draft block, so most of the dense ``[B, 1, Q, KV]`` mask is empty. SDPA computes all of it, +on its memory-efficient fallback (the fused kernels reject arbitrary masks and cap +``head_dim`` at 256); FlexAttention takes the mask as a predicate and skips empty tiles. + +K/V are repeated to the query head count rather than using ``enable_gqa=True``, whose +backward is ~10x slower. +""" + +import torch + +__all__ = ["build_draft_block_mask", "flex_attention_forward", "is_block_mask"] + +# head_dim > 256 cannot use FlexAttention's default tiles (SMEM overflow); 32x32 with +# two pipeline stages is the only combination measured to both fit and run. +_LARGE_HEAD_DIM_KERNEL_OPTIONS = { + "BLOCK_M": 32, + "BLOCK_N": 32, + "BLOCK_M1": 32, + "BLOCK_N1": 32, + "BLOCK_M2": 32, + "BLOCK_N2": 32, + "num_stages": 2, + "num_warps": 4, +} +_MAX_DEFAULT_TILE_HEAD_DIM = 256 + +_PINNED_TILE_MASK_BLOCK_SIZE = 64 +_DEFAULT_TILE_MASK_BLOCK_SIZE = 128 + + +def _mask_block_size(head_dim): + """BlockMask block size, which must be divisible by the kernel's BLOCK_M/BLOCK_N. + + The pinned 32x32 tiles allow the finer 64; the autotuner may pick tiles up to 128. + """ + if head_dim > _MAX_DEFAULT_TILE_HEAD_DIM: + return _PINNED_TILE_MASK_BLOCK_SIZE + return _DEFAULT_TILE_MASK_BLOCK_SIZE + + +_flex_attention_compiled = None +_create_block_mask_compiled = None + + +def _flex_ops(): + """Resolve and compile the FlexAttention entry points once per process.""" + global _flex_attention_compiled, _create_block_mask_compiled + if _flex_attention_compiled is None: + from torch.nn.attention.flex_attention import create_block_mask, flex_attention + + # dynamic=False: shapes are fixed within a run, and dynamic shapes slow the kernel. + _flex_attention_compiled = torch.compile(flex_attention, dynamic=False) + _create_block_mask_compiled = torch.compile(create_block_mask, dynamic=False) + return _flex_attention_compiled, _create_block_mask_compiled + + +def is_block_mask(mask) -> bool: + """True if ``mask`` is a FlexAttention ``BlockMask`` rather than a dense tensor.""" + if mask is None or torch.is_tensor(mask): + return False + try: + from torch.nn.attention.flex_attention import BlockMask + except ImportError: + return False + return isinstance(mask, BlockMask) + + +def build_draft_block_mask(mask_mod, bsz, q_len, kv_len, device, head_dim): + """Block-sparse ``BlockMask`` of the draft attention's ``mask_mod``.""" + _, create_block_mask = _flex_ops() + return create_block_mask( + mask_mod, + bsz, + None, + q_len, + kv_len, + device=device, + BLOCK_SIZE=_mask_block_size(head_dim), + ) + + +def _repeat_kv(x, n_rep): + """HF's ``repeat_kv``: [B, n_kv, S, D] -> [B, n_kv * n_rep, S, D].""" + if n_rep == 1: + return x + b, h, s, d = x.shape + return x[:, :, None].expand(b, h, n_rep, s, d).reshape(b, h * n_rep, s, d) + + +def flex_attention_forward(query, key, value, block_mask, scaling): + """FlexAttention with the draft's BlockMask. Returns ``[B, q_len, n_heads, head_dim]``. + + The layout matches what HF's ``sdpa_attention_forward`` returns so the caller's + ``reshape(bsz, q_len, -1)`` is unchanged. + """ + flex_attention, _ = _flex_ops() + head_dim = query.shape[-1] + n_rep = query.shape[1] // key.shape[1] + kernel_options = ( + _LARGE_HEAD_DIM_KERNEL_OPTIONS if head_dim > _MAX_DEFAULT_TILE_HEAD_DIM else None + ) + attn_output = flex_attention( + query, + _repeat_kv(key, n_rep), + _repeat_kv(value, n_rep), + block_mask=block_mask, + scale=scaling, + kernel_options=kernel_options, + ) + return attn_output.transpose(1, 2).contiguous() diff --git a/modelopt/torch/speculative/plugins/hf_dflash.py b/modelopt/torch/speculative/plugins/hf_dflash.py index cd596f773d9..7fe9bd7d23e 100644 --- a/modelopt/torch/speculative/plugins/hf_dflash.py +++ b/modelopt/torch/speculative/plugins/hf_dflash.py @@ -96,6 +96,7 @@ _FINAL_NORM_PATHS, _LM_HEAD_PATHS, ) +from .modeling_final_norm import _maybe_apply_base_final_norm logger = logging.getLogger(__name__) @@ -125,6 +126,39 @@ def _multimodal_forward_kwargs(model_kwargs: dict) -> dict: } +def _gather_rows(x, positions): + """``x[b, positions[b, ...]]`` for ``x`` [B, seq, D] and ``positions`` [B, ...].""" + idx = positions.reshape(positions.shape[0], -1, 1).expand(-1, -1, x.size(-1)) + return torch.gather(x, 1, idx).reshape(*positions.shape, x.size(-1)) + + +def _draft_mask_mod(seq_len, anchor_positions, block_keep_mask, block_size, window, causal): + """The draft's attention visibility rule, as a FlexAttention ``mask_mod``. + + Elementwise ops and indexing only, so the one rule feeds ``create_block_mask`` and, + evaluated on broadcast index grids, the dense SDPA mask. + """ + # Indexed inside mask_mod, which create_block_mask runs under vmap: keep them integral. + anchors = anchor_positions.to(torch.int32) + keep = block_keep_mask.to(torch.bool) + + def mask_mod(b, h, q_idx, kv_idx): + q_block = q_idx // block_size + anchor = anchors[b, q_block] + is_ctx = kv_idx < seq_len + ctx_ok = is_ctx & (kv_idx < anchor) + if window is not None: + # Window from the query's real position (anchor + position in block). + ctx_ok = ctx_ok & (kv_idx > anchor + (q_idx % block_size) - window) + draft_ok = (~is_ctx) & (q_block == (kv_idx - seq_len) // block_size) + if causal: + # Block-causal: block position i sees draft positions <= i. + draft_ok = draft_ok & (((kv_idx - seq_len) % block_size) <= (q_idx % block_size)) + return (ctx_ok | draft_ok) & keep[b, q_block] + + return mask_mod + + def _dpace_position_weights( confidences: torch.Tensor, alpha: float, valid_mask: torch.Tensor | None = None ) -> torch.Tensor: @@ -430,6 +464,8 @@ def modify(self, config): self.is_quantized = False self._num_anchors = self.dflash_num_anchors + if self.dflash_use_torch_compile: + self.dflash_module.compile_body() def _build_draft_module(self, dflash_config): """Build the draft module. Subclasses override to use an augmented module.""" @@ -539,6 +575,9 @@ def _sample_anchor_positions(self, seq_len, loss_mask, device): Returns (anchor_positions [B, N], block_keep_mask [B, N]). + ``N`` is fixed by the config and ``seq_len``, not by the batch, so the compiled + (``dynamic=False``) draft attention sees one shape and never recompiles. + TODO: Fix the random seed per epoch (change between epochs) so that anchor positions are deterministic within an epoch. This would allow caching the derived masks and position IDs across steps while preserving the same data augmentation @@ -551,27 +590,29 @@ def _sample_anchor_positions(self, seq_len, loss_mask, device): valid = loss_mask[:, : max_anchor + 1] > 0.5 valid_counts = valid.sum(dim=1) - max_n = min(num_anchors, int(valid_counts.max().item()) - 1) - if max_n <= 0: - # No valid anchors — return empty - anchors = torch.zeros(bsz, 1, dtype=torch.long, device=device) - keep = torch.zeros(bsz, 1, dtype=torch.bool, device=device) - return anchors, keep + # Also bounded by max_anchor + 1 so short sequences cannot ask for more columns. + max_n = min(num_anchors, max_anchor + 1) + + # The data-dependent bound on sampled anchors, kept on-device so no shape depends on it. + cap = (valid_counts.max() - 1).clamp(min=0, max=max_n) indices = torch.arange(max_anchor + 1, device=device).unsqueeze(0).expand(bsz, -1) - masked_indices = torch.where(valid, indices, torch.tensor(seq_len + 1, device=device)) + fill = torch.tensor(seq_len + 1, device=device) + masked_indices = torch.where(valid, indices, fill) random_vals = torch.rand(bsz, max_anchor + 1, device=device) random_vals = torch.where(valid, random_vals, torch.tensor(2.0, device=device)) _, sorted_idx = random_vals.sort(dim=1) gathered = torch.gather(masked_indices, 1, sorted_idx) - anchors = gathered[:, :max_n].sort(dim=1).values - keep = torch.arange(max_n, device=device).unsqueeze(0) < valid_counts.unsqueeze(1).clamp( - max=max_n - ) + # Blank past `cap` before sorting: sorting first would pull anchors from beyond the + # bound into the row and silently change which positions are trained on. + cols = torch.arange(max_n, device=device).unsqueeze(0) + anchors = torch.where(cols < cap, gathered[:, :max_n], fill).sort(dim=1).values + + keep = cols < torch.minimum(valid_counts, cap).unsqueeze(1) anchors = torch.where(keep, anchors, torch.tensor(0, dtype=torch.long, device=device)) return anchors, keep @@ -609,7 +650,7 @@ def _build_position_ids(self, seq_len, anchor_positions, device): def _build_draft_attention_mask( self, seq_len, anchor_positions, block_keep_mask, n_blocks, dtype, device, window=None ): - """Build SDPA attention mask: context (causal) + draft (per ``dflash_draft_attention``). + """Build the draft attention mask: context (causal) + draft (per ``dflash_draft_attention``). When ``window`` is not None, all layers use sliding-window attention: each draft query only sees context positions within ``window`` tokens before its own position. @@ -621,44 +662,42 @@ def _build_draft_attention_mask( ``"bidirectional"`` (default, MiMo-style) lets every query see the whole block, while ``"causal"`` restricts a query at block position ``i`` to draft positions ``<= i`` so the block is modelled autoregressively. + + Returns a dense additive SDPA mask, or a ``BlockMask`` under + ``dflash_use_flex_attention``; both evaluate the one rule in ``_draft_mask_mod``. """ bsz = anchor_positions.shape[0] - block_size = self.dflash_block_size - q_len = n_blocks * block_size + q_len = n_blocks * self.dflash_block_size kv_len = seq_len + q_len + mask_mod = _draft_mask_mod( + seq_len, + anchor_positions, + block_keep_mask, + self.dflash_block_size, + window, + causal=self.dflash_draft_attention == "causal", + ) - q_indices = torch.arange(q_len, device=device).view(1, 1, -1, 1) - kv_indices = torch.arange(kv_len, device=device).view(1, 1, 1, -1) - q_block_ids = q_indices // block_size - - anchor_exp = anchor_positions.view(bsz, 1, n_blocks, 1).repeat_interleave(block_size, dim=2) + if self.dflash_use_flex_attention: + from .dflash_flex_attention import build_draft_block_mask - # Context: kv < S and kv < anchor - mask_ctx = (kv_indices < seq_len) & (kv_indices < anchor_exp) + return build_draft_block_mask( + mask_mod, + bsz, + q_len, + kv_len, + device, + head_dim=self.dflash_module.layers[0].self_attn.head_dim, + ) - # Sliding window on the context: keep only context kv whose real position is within - # `window` tokens before the query's real position (anchor + position-in-block). - if window is not None: - q_real_pos = anchor_exp + (q_indices % block_size) # [B, 1, q_len, 1] - mask_ctx = mask_ctx & (kv_indices > q_real_pos - window) - # Draft: kv >= S and same block - is_draft = kv_indices >= seq_len - kv_block_ids = (kv_indices - seq_len) // block_size - mask_draft = is_draft & (q_block_ids == kv_block_ids) - if self.dflash_draft_attention == "causal": - # Autoregressive within the block: query at block position i sees draft - # positions <= i only. Compare positions *within* the block so the term is - # independent of which block the query belongs to. - kv_pos_in_block = (kv_indices - seq_len) % block_size - mask_draft = mask_draft & (kv_pos_in_block <= (q_indices % block_size)) - # Valid block - valid_block = block_keep_mask.view(bsz, 1, n_blocks, 1).repeat_interleave(block_size, dim=2) - - final_mask = (mask_ctx | mask_draft) & valid_block # [B, 1, Q, KV] - - # Convert bool mask to float additive mask for SDPA + visible = mask_mod( + torch.arange(bsz, device=device).view(-1, 1, 1, 1), + None, + torch.arange(q_len, device=device).view(1, 1, -1, 1), + torch.arange(kv_len, device=device).view(1, 1, 1, -1), + ) # [B, 1, Q, KV] attn_mask = torch.zeros(bsz, 1, q_len, kv_len, device=device, dtype=dtype) - attn_mask.masked_fill_(~final_mask, torch.finfo(dtype).min) + attn_mask.masked_fill_(~visible, torch.finfo(dtype).min) return attn_mask def _build_generate_swa_mask(self, ctx_len, bsz, dtype, device): @@ -693,6 +732,40 @@ def _build_generate_swa_mask(self, ctx_len, bsz, dtype, device): attn_mask.masked_fill_(~keep, torch.finfo(dtype).min) return attn_mask + @torch.no_grad() + def _teacher_logits(self, base_outputs, positions, token_ids=None): + """Base-model logits at ``positions`` ([B, ...] sequence indices) -> [B, ..., vocab]. + + Only the requested rows of the base hidden go through the final norm and lm_head, so + no full-sequence logits are built from it; ``token_ids`` ([B, ..., k]) narrows the + projection to just those vocab entries. + """ + if base_outputs.logits is not None: + if token_ids is None: + return _gather_rows(base_outputs.logits, positions) + batch = torch.arange(positions.shape[0], device=positions.device) + batch = batch.view(-1, *[1] * positions.dim()) + return base_outputs.logits[batch, positions.unsqueeze(-1), token_ids] + if base_outputs.base_hidden is None: + raise ValueError( + "This objective needs the base model's distribution, but the batch carries " + "neither its logits nor its final hidden states (base_model_hidden_states)." + ) + lm_head = self._base_model_lm_head + # A producer can store the hidden states in a wider dtype than the target's weights. + # Cast before the final norm too, which online training runs in the target's dtype. + rows = _gather_rows(base_outputs.base_hidden, positions).to(lm_head.weight.dtype) + rows = _maybe_apply_base_final_norm( + rows, + {"base_hidden_prenorm": base_outputs.base_hidden_prenorm}, + self._base_model_norm, + ) + if token_ids is None: + return lm_head(rows) + logits = torch.einsum("...h,...kh->...k", rows, lm_head.weight[token_ids]) + bias = getattr(lm_head, "bias", None) + return logits if bias is None else logits + bias[token_ids] + def _compute_loss( self, logits, @@ -700,7 +773,7 @@ def _compute_loss( anchor_positions, block_keep_mask, loss_mask, - base_logits=None, + *, draft_hidden=None, base_outputs=None, return_terms=False, @@ -713,9 +786,10 @@ def _compute_loss( anchor_positions: Anchor positions per block [B, N]. block_keep_mask: Valid block mask [B, N]. loss_mask: Token-level loss mask [B, seq_len]. - 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. + base_outputs: The step's ``DFlashBaseModelOutput``: the KD teacher under + ``dflash_self_logit_distillation``, and read by variants. 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. @@ -755,12 +829,17 @@ def _compute_loss( flat_logits = logits.view(-1, logits.size(-1)) flat_targets = target_ids.view(-1) + kd = self.dflash_self_logit_distillation + if kd and base_outputs is None: + raise ValueError( + "dflash_self_logit_distillation distills from base_outputs, but none was " + "passed; an override of _compute_loss must forward it." + ) # Non-KD loss is per-token cross-entropy; compute it once (grad enabled) so the # D-PACE confidences below can reuse it instead of a second CE pass. The KD path - # (base_logits is not None) optimizes KL, so its confidences need a dedicated - # no_grad CE pass. + # optimizes KL, so its confidences need a dedicated no_grad CE pass. loss_per_token = None - if base_logits is None: + if not kd: loss_per_token = F.cross_entropy(flat_logits, flat_targets, reduction="none") # Block-position loss weighting: dynamic D-PACE weights or static exponential decay. @@ -790,15 +869,11 @@ def _compute_loss( valid_count = flat_weights.sum() + 1e-6 if valid_count > 1.0: - if base_logits is not None: + if kd: # KD loss: teacher logits for token anchor+k are at position anchor+k-1 teacher_indices = (safe_label_indices - 1).clamp(min=0) - teacher_logits = torch.gather( - base_logits.unsqueeze(1).expand(-1, n_blocks, -1, -1), - 2, - teacher_indices.unsqueeze(-1).expand(-1, -1, -1, base_logits.size(-1)), - ) - flat_teacher = teacher_logits.reshape(-1, base_logits.size(-1)).detach() + flat_teacher = self._teacher_logits(base_outputs, teacher_indices) + flat_teacher = flat_teacher.reshape(flat_logits.shape) target_soft = torch.softmax(flat_teacher, dim=-1) draft_logsoft = torch.log_softmax(flat_logits, dim=-1) kd_loss = -(target_soft * draft_logsoft).sum(dim=-1) @@ -882,15 +957,14 @@ def forward( # 1. Run base model → extract target hidden states if self.dflash_offline: assert "base_model_outputs" in kwargs - # When the loss needs them (see _needs_base_logits), from_offline_dict reconstructs - # base logits from the captured hidden (final norm re-applied as needed) when the producer didn't supply - # them, and raises if anything needed for that is missing. - base_outputs = DFlashBaseModelOutput.from_offline_dict( - kwargs["base_model_outputs"], - self._base_model_norm, - self._base_model_lm_head, - need_logits=self._needs_base_logits, - ) + base_outputs = DFlashBaseModelOutput.from_offline_dict(kwargs["base_model_outputs"]) + # Fail at the entry point rather than in the loss, after the draft has run. + if ( + self._needs_base_logits + and base_outputs.logits is None + and base_outputs.base_hidden is None + ): + raise KeyError("base_model_hidden_states") target_hidden = base_outputs.target_hidden else: # Multimodal models need the top-level conditional-generation forward so their @@ -996,7 +1070,6 @@ def forward( anchor_positions, block_keep_mask, loss_mask, - base_outputs.logits if self.dflash_self_logit_distillation else None, draft_hidden=hidden, base_outputs=base_outputs, ) diff --git a/modelopt/torch/speculative/plugins/hf_dflash2.py b/modelopt/torch/speculative/plugins/hf_dflash2.py index 625bc622688..6060923b546 100644 --- a/modelopt/torch/speculative/plugins/hf_dflash2.py +++ b/modelopt/torch/speculative/plugins/hf_dflash2.py @@ -205,7 +205,7 @@ def _compute_loss( anchor_positions, block_keep_mask, loss_mask, - base_logits=None, + *, draft_hidden=None, base_outputs=None, ): @@ -222,7 +222,6 @@ def _compute_loss( anchor_positions, block_keep_mask, loss_mask, - base_logits, draft_hidden=draft_hidden, base_outputs=base_outputs, return_terms=True, diff --git a/modelopt/torch/speculative/plugins/hf_dspark.py b/modelopt/torch/speculative/plugins/hf_dspark.py index bf877f1ee5d..7c81f4cc478 100644 --- a/modelopt/torch/speculative/plugins/hf_dspark.py +++ b/modelopt/torch/speculative/plugins/hf_dspark.py @@ -78,7 +78,12 @@ __all__ = ["HFDSparkModel"] -def _tvd_per_token(final_logits, teacher_logits, chunk_size=1024): +def _tvd_chunk(a, b): + """Per-token TVD for one row chunk: ``(softmax(a) - softmax(b)).abs().sum(-1)``.""" + return (torch.softmax(a.float(), dim=-1) - torch.softmax(b.float(), dim=-1)).abs().sum(dim=-1) + + +def _tvd_per_token(final_logits, teacher_logits, chunk_size=1024, chunk_fn=None): """Total-variation distance ||softmax(a)-softmax(b)||_1 / ... per token, memory-lean. Materializing both [N, vocab] float32 softmax tensors at once OOMs at large @@ -86,16 +91,17 @@ def _tvd_per_token(final_logits, teacher_logits, chunk_size=1024): gradient-checkpoint each chunk so the wide softmaxes are recomputed in backward rather than held — peak memory ~ chunk_size*vocab instead of N*vocab. The math is identical to ``(softmax(final)-softmax(teacher)).abs().sum(-1)``. - """ - - def _chunk(a, b): - return ( - (torch.softmax(a.float(), dim=-1) - torch.softmax(b.float(), dim=-1)).abs().sum(dim=-1) - ) + Chunks with ``Tensor.split``, not slicing (each slice's backward zero-fills a full + [N, vocab] tensor) nor ``torch.chunk`` (different chunk shapes, so recompiles). + """ + _chunk = chunk_fn or _tvd_chunk outs = [] - for i in range(0, final_logits.size(0), chunk_size): - a, b = final_logits[i : i + chunk_size], teacher_logits[i : i + chunk_size] + for a, b in zip( + final_logits.split(chunk_size, dim=0), + teacher_logits.split(chunk_size, dim=0), + strict=True, + ): if torch.is_grad_enabled() and a.requires_grad: outs.append(torch.utils.checkpoint.checkpoint(_chunk, a, b, use_reentrant=False)) else: @@ -134,6 +140,12 @@ def modify(self, config): "dflash_confidence_head_alpha > 0 but the confidence head was not built; " "set dflash_architecture_config.use_confidence_head=true." ) + # Compiling fuses the chunk's six vocab-wide elementwise passes. + self._tvd_chunk_fn = ( + torch.compile(_tvd_chunk, dynamic=False, fullgraph=True) + if self.dflash_use_torch_compile + else _tvd_chunk + ) def get_exporter(self): """Get the exporter for the DSpark draft model.""" @@ -176,13 +188,14 @@ def _compute_dspark_loss( anchor_positions, block_keep_mask, loss_mask, - target_model_logits, + base_outputs, ): """Compute the three-term DSpark loss (CE + TVD + confidence BCE) and metrics. Uses next-token (shift_label) alignment: block position k predicts the token at anchor+k+1; the aligned target distribution is the base model's own - next-token distribution at position anchor+k (= label index - 1). + next-token distribution at position anchor+k (= label index - 1), read from + ``base_outputs``. """ bsz, seq_len = input_ids.shape bs = self.dflash_block_size @@ -222,28 +235,31 @@ def _compute_dspark_loss( flat_weights = weight_mask.reshape(-1) valid_count = flat_weights.sum() + 1e-6 - # Aligned target distribution: base-model logits that predict token anchor+k+1 - # sit at position anchor+k (= label index - 1). - teacher_indices = (safe_label_indices - 1).clamp(min=0) - teacher_logits = torch.gather( - target_model_logits.unsqueeze(1).expand(-1, n_blocks, -1, -1), - 2, - teacher_indices.unsqueeze(-1).expand(-1, -1, -1, vocab), - ) - flat_teacher = teacher_logits.reshape(-1, vocab).detach() - if valid_count <= 1.0: - loss = flat_final.sum() * 0.0 + # Touch every draft parameter, as forward()'s early return does, so DDP with + # find_unused_parameters=False still sees the confidence head's gradient. + loss = ( + flat_final.sum() * 0.0 + sum(p.sum() for p in self.dflash_module.parameters()) * 0.0 + ) metrics = {"ce_loss": 0.0, "l1_loss": 0.0, "confidence_loss": 0.0, "base_accuracy": 0.0} return loss, 0.0, metrics + # Aligned target distribution: base-model logits that predict token anchor+k+1 + # sit at position anchor+k (= label index - 1). + teacher_indices = (safe_label_indices - 1).clamp(min=0) + flat_teacher = self._teacher_logits(base_outputs, teacher_indices).reshape(-1, vocab) + # Term 1: cross-entropy on the corrected (final) logits. ce_per_token = F.cross_entropy(flat_final, flat_targets, reduction="none") ce_loss = (ce_per_token * flat_weights).sum() / valid_count # Term 2: total-variation distance between the corrected draft and target. # Chunked + checkpointed to avoid materializing two [N, vocab] softmaxes at once. - l1_per_token = _tvd_per_token(flat_final, flat_teacher) + l1_per_token = _tvd_per_token( + flat_final, + flat_teacher, + chunk_fn=self._tvd_chunk_fn, + ) l1_loss = (l1_per_token * flat_weights).sum() / valid_count # Term 3: confidence head BCE against the analytical accept rate c* = 1 - 0.5*TVD. @@ -264,19 +280,21 @@ def _compute_dspark_loss( with torch.no_grad(): eval_count = binary_eval_mask.sum() + 1e-6 keep = binary_eval_mask > 0.5 - accuracy = ( - ((flat_final.argmax(dim=-1) == flat_targets) & keep).sum().float() / eval_count - ).item() - base_accuracy = ( - ((flat_base.argmax(dim=-1) == flat_targets) & keep).sum().float() / eval_count - ).item() + acc = ((flat_final.argmax(dim=-1) == flat_targets) & keep).sum().float() / eval_count + base_acc = ( + (flat_base.argmax(dim=-1) == flat_targets) & keep + ).sum().float() / eval_count + # One device sync for all five scalars instead of one per .item(). + acc_v, base_acc_v, ce_v, l1_v, conf_v = torch.stack( + [acc, base_acc, ce_loss.detach(), l1_loss.detach(), confidence_loss.detach()] + ).tolist() metrics = { - "ce_loss": ce_loss.detach().item(), - "l1_loss": l1_loss.detach().item(), - "confidence_loss": float(confidence_loss.detach().item()), - "base_accuracy": base_accuracy, + "ce_loss": ce_v, + "l1_loss": l1_v, + "confidence_loss": conf_v, + "base_accuracy": base_acc_v, } - return loss, accuracy, metrics + return loss, acc_v, metrics def forward( self, @@ -322,37 +340,27 @@ def forward( f"Adjust training_seq_len or use padding." ) - # 1. Target hidden states AND target-model logits (DSpark's L1/confidence - # terms both need the base model's next-token distribution). + # 1. Target hidden states, plus what the TVD/confidence terms read the base + # distribution from (the loss projects only the rows it uses). if self.dflash_offline: assert "base_model_outputs" in kwargs - # Reconstruct base logits through the shared DFlash offline path so the base - # final norm is re-applied when the producer captured a pre-(final-)norm hidden - # (vLLM streaming) — feeding an un-normed hidden straight to lm_head would make a - # corrupt distillation target. DSpark always needs the base distribution (its - # TVD/confidence terms), so need_logits=True unconditionally. - base_outputs = DFlashBaseModelOutput.from_offline_dict( - kwargs["base_model_outputs"], - self._base_model_norm, - self._base_model_lm_head, - need_logits=True, - ) - target_hidden = base_outputs.target_hidden - target_model_logits = base_outputs.logits + base_outputs = DFlashBaseModelOutput.from_offline_dict(kwargs["base_model_outputs"]) else: # Call the inner base model directly (NOT super().forward(), which during - # training runs the full DFlash pipeline). Compute target-model logits via - # the lm_head — DSpark's TVD/confidence terms need the base distribution. + # training runs the full DFlash pipeline). with torch.no_grad(): base_out = self._base_model( input_ids=input_ids, attention_mask=attention_mask, output_hidden_states=True, ) - target_model_logits = self._base_model_lm_head(base_out.last_hidden_state) offset = 1 selected = [base_out.hidden_states[lid + offset] for lid in self.target_layer_ids] - target_hidden = torch.cat(selected, dim=-1) # [B, seq, num_layers * H] + base_outputs = DFlashBaseModelOutput( + target_hidden=torch.cat(selected, dim=-1), # [B, seq, num_layers * H] + base_hidden=base_out.last_hidden_state, + ) + target_hidden = base_outputs.target_hidden # 2. Build loss mask (same convention as DFlash/Domino). if labels is not None: @@ -411,7 +419,7 @@ def forward( anchor_positions, block_keep_mask, loss_mask, - target_model_logits, + base_outputs, ) return ModelOutput(loss=loss, logits=None, train_acc=[[accuracy]], dspark_metrics=metrics) diff --git a/modelopt/torch/speculative/plugins/hf_lilicorr.py b/modelopt/torch/speculative/plugins/hf_lilicorr.py index b6d8b2f2d8c..42f6a8bbed5 100644 --- a/modelopt/torch/speculative/plugins/hf_lilicorr.py +++ b/modelopt/torch/speculative/plugins/hf_lilicorr.py @@ -237,7 +237,7 @@ def _compute_loss( anchor_positions, block_keep_mask, loss_mask, - base_logits=None, + *, draft_hidden=None, base_outputs=None, ): @@ -247,13 +247,12 @@ def _compute_loss( quantity that tracks acceptance length — rather than the backbone's per-token accuracy, which is reported as ``origin_accuracy`` in the metrics instead. - The two tensors this needs beyond the shared signature -- the target-layer hidden - states it anchors on, and the target's logits its penalty and calibration terms read -- - both come from ``base_outputs``, the container ``HFDFlashModel.forward`` already - builds for them. + The two things this needs beyond the shared signature -- the target-layer hidden + states it anchors on, and the target distribution its penalty and calibration + terms read -- both come from ``base_outputs``, the container + ``HFDFlashModel.forward`` already builds for them. """ target_hidden = base_outputs.target_hidden if base_outputs is not None else None - target_logits = base_outputs.logits if base_outputs is not None else None if draft_hidden is None or target_hidden is None: raise ValueError( "LiLiCorr requires draft_hidden and target_hidden in _compute_loss: the " @@ -267,7 +266,7 @@ def _compute_loss( anchor_positions, block_keep_mask, loss_mask, - base_logits, + base_outputs=base_outputs, ) target_ids, slot_mask = self._block_targets( input_ids, anchor_positions, block_keep_mask, loss_mask @@ -279,7 +278,7 @@ def _compute_loss( anchor_positions=anchor_positions, draft_hidden=draft_hidden, target_hidden=target_hidden, - target_logits=target_logits, + base_outputs=base_outputs, ) # The identity `loss == origin_loss + lilicorr_loss` is the cheap parity check # on this objective, so both halves are reported next to the total. @@ -303,7 +302,7 @@ def _compute_lilicorr_loss( anchor_positions, draft_hidden, target_hidden, - target_logits, + base_outputs, ): """Score the candidate lattice and evaluate the LiLiCorr objective. @@ -442,7 +441,7 @@ def _compute_lilicorr_loss( candidate_ids=candidate_ids, gt_indices=gt_indices, anchor_positions=anchor_positions, - target_logits=target_logits, + base_outputs=base_outputs, supervised=supervised, denominator=denominator, ) @@ -454,7 +453,7 @@ def _compute_lilicorr_loss( node_potentials=node_potentials, candidate_ids=candidate_ids, anchor_positions=anchor_positions, - target_logits=target_logits, + base_outputs=base_outputs, supervised=supervised, denominator=denominator, ) @@ -500,7 +499,7 @@ def _distractor_penalty( candidate_ids, gt_indices, anchor_positions, - target_logits, + base_outputs, supervised, denominator, ): @@ -521,7 +520,7 @@ def _distractor_penalty( candidate_target_logits = self._candidate_target_logits( candidate_ids=candidate_ids, anchor_positions=anchor_positions, - target_logits=target_logits, + base_outputs=base_outputs, requested_by="dflash_lilicorr_w_pen", ) # w_j is 0 at the ground truth itself, and 0 for any candidate the target scores @@ -534,37 +533,29 @@ def _distractor_penalty( return (per_slot * supervised).sum() / denominator def _candidate_target_logits( - self, *, candidate_ids, anchor_positions, target_logits, requested_by + self, *, candidate_ids, anchor_positions, base_outputs, requested_by ): """The target's logits for exactly the candidates in play: ``[blocks, slots, k]``.""" - if target_logits is None: + if base_outputs.logits is None and base_outputs.base_hidden is None: raise ValueError( f"{requested_by} > 0 requires the target model's logits. Online training " - "reads them from the target; offline training reconstructs them from the " + "reads them from the target; offline training projects them from the " "captured final hidden state." ) batch_blocks, num_slots, topk = candidate_ids.shape bsz, n_blocks = anchor_positions.shape - device = candidate_ids.device - - target_seq_len = target_logits.shape[1] + seq_len = base_outputs.target_hidden.shape[1] # The target's next-token logits that predict the token at anchor+1+s sit at - # position anchor+s. `sample_index` maps each flattened block back to its batch - # row, so (sample, position) addresses one target logit vector per slot. - sample_index = torch.arange(bsz, device=device).repeat_interleave(n_blocks) - slot_offsets = torch.arange(num_slots, device=device).view(1, -1) - positions = (anchor_positions.reshape(batch_blocks, 1) + slot_offsets).clamp( - min=0, max=target_seq_len - 1 - ) - sample_expanded = sample_index.view(batch_blocks, 1, 1).expand( - batch_blocks, num_slots, topk - ) - positions_expanded = positions.view(batch_blocks, num_slots, 1).expand( - batch_blocks, num_slots, topk + # position anchor+s. Only the [blocks, slots, k] candidate logits are computed, + # never whole vocabulary rows. + slot_offsets = torch.arange(num_slots, device=candidate_ids.device) + positions = (anchor_positions.unsqueeze(-1) + slot_offsets).clamp(min=0, max=seq_len - 1) + candidate_target_logits = self._teacher_logits( + base_outputs, + positions, + token_ids=candidate_ids.reshape(bsz, n_blocks, num_slots, topk), ) - # Advanced indexing reads only the [blocks, slots, k] scalars in play rather - # than gathering whole vocabulary rows. Teacher weights, hence detached. - return target_logits[sample_expanded, positions_expanded, candidate_ids].detach().float() + return candidate_target_logits.reshape(batch_blocks, num_slots, topk).float() def _calibration_loss( self, @@ -572,7 +563,7 @@ def _calibration_loss( node_potentials, candidate_ids, anchor_positions, - target_logits, + base_outputs, supervised, denominator, ): @@ -586,7 +577,7 @@ def _calibration_loss( candidate_target_logits = self._candidate_target_logits( candidate_ids=candidate_ids, anchor_positions=anchor_positions, - target_logits=target_logits, + base_outputs=base_outputs, requested_by="dflash_lilicorr_w_cal", ) target_probs = F.softmax(candidate_target_logits, dim=-1) diff --git a/modelopt/torch/speculative/plugins/hf_streaming_dataset.py b/modelopt/torch/speculative/plugins/hf_streaming_dataset.py index ad6082eb71a..fc9fa2f765b 100644 --- a/modelopt/torch/speculative/plugins/hf_streaming_dataset.py +++ b/modelopt/torch/speculative/plugins/hf_streaming_dataset.py @@ -545,14 +545,16 @@ def _fetch(self, sample: dict) -> EagleFetchPayload | None: time.sleep(0.0002) agent.release_xfer_handle(h) hidden_states = view.clone() # copy out before /done so the gen check brackets the read - # /done frees the slot + reports valid; valid=False -> ring lapped us mid-read, bytes - # stale -> resample. A failed /done can't prove staleness, so default valid=True. + # /done frees the slot and reports whether the ring lapped us mid-read (stale bytes). + # Fail closed: a failed /done cannot prove the bytes are ours, and a resample is cheap. try: valid = self._http_rdma.get( f"http://{host}:{port}/done", params={"req_id": rid} ).json()["valid"] - except Exception: - valid = True + except Exception as exc: + # Kept apart from the lap warning: these point at the sidecar, not ring pressure. + warn_rank_0(f"[streaming] /done failed for {sample['cid']} ({exc!r}); resampling") + return None if not valid: warn_rank_0(f"[streaming] slot lapped mid-read for {sample['cid']}; resampling") return None diff --git a/modelopt/torch/speculative/plugins/modeling_dflash.py b/modelopt/torch/speculative/plugins/modeling_dflash.py index cc1a9b5277a..34abc990327 100644 --- a/modelopt/torch/speculative/plugins/modeling_dflash.py +++ b/modelopt/torch/speculative/plugins/modeling_dflash.py @@ -56,8 +56,6 @@ from transformers.models.qwen3.modeling_qwen3 import repeat_kv from transformers.models.qwen3.modeling_qwen3 import rotate_half as _rotate_half -from .modeling_final_norm import _maybe_apply_base_final_norm - __all__ = ["DFlashBaseModelOutput", "DFlashModule", "build_target_layer_ids"] @@ -116,46 +114,31 @@ def _get_sink_attention_fn(): @dataclass class DFlashBaseModelOutput: - """Output container for base model forward pass in DFlash training.""" + """What a DFlash-family draft takes from the base model in a training step. + + ``logits`` holds the full base logits when the producer already has them (DFlash's + online CausalLM forward, or an offline batch carrying ``base_model_logits``); + otherwise ``base_hidden`` stays unprojected, and ``HFDFlashModel._teacher_logits`` + projects only the rows a loss reads. + """ target_hidden: torch.Tensor # concatenated hidden states from target layers [B, seq, N*H] - logits: torch.Tensor | None = None # base model logits [B, seq, vocab] + base_hidden: torch.Tensor | None = None # base final hidden, lm_head's input [B, seq, H] + base_hidden_prenorm: bool = False # base_hidden was captured before the final norm + logits: torch.Tensor | None = None # base logits [B, seq, vocab], when handed over as such @classmethod - def from_offline_dict( - cls, d: dict, base_model_norm=None, base_model_lm_head=None, need_logits=False - ): + def from_offline_dict(cls, d: dict): """Construct from a dict of pre-computed base model outputs (offline training). ``aux_hidden_states`` is required — missing it raises KeyError at the entry point rather than producing a cryptic failure deeper in the forward. - - When ``need_logits`` (self-logit-distillation) and the producer didn't supply - ``base_model_logits``, logits are reconstructed from the captured final hidden via - ``base_model_lm_head`` — first re-applying the base final norm when the producer captured - a pre-(final-)norm hidden (``base_hidden_prenorm``), so the reconstruction is correct - regardless of capture format. Anything missing on that path raises rather than silently - yielding None logits: no ``base_model_lm_head`` (ValueError), no captured hidden - (KeyError), or a pre-norm hidden with no ``base_model_norm`` (feeding an un-normed hidden - to lm_head would be a corrupt distillation target). """ - logits = d.get("base_model_logits") - if need_logits and logits is None: - if base_model_lm_head is None: - raise ValueError( - "need_logits=True but base_model_lm_head is None; cannot reconstruct logits." - ) - out_hiddens = d.get("base_model_hidden_states") - if out_hiddens is None: - raise KeyError("base_model_hidden_states") - # A producer can store the hidden states in a wider dtype than the target's weights. - # Cast before the final norm too, which online training runs in the target's dtype. - out_hiddens = out_hiddens.to(base_model_lm_head.weight.dtype) - out_hiddens = _maybe_apply_base_final_norm(out_hiddens, d, base_model_norm) - logits = base_model_lm_head(out_hiddens) return cls( target_hidden=d["aux_hidden_states"], - logits=logits, + base_hidden=d.get("base_model_hidden_states"), + base_hidden_prenorm=bool(d.get("base_hidden_prenorm", False)), + logits=d.get("base_model_logits"), ) @@ -274,7 +257,24 @@ def forward(self, hidden_states, target_hidden, position_embeddings, attention_m cos, sin = position_embeddings q, k = apply_rotary_pos_emb(q, k, cos, sin) - if self.attention_sink_bias is not None: + from .dflash_flex_attention import flex_attention_forward, is_block_mask + + if is_block_mask(attention_mask): + if self.attention_sink_bias is not None: + # The sink is an extra softmax column the flex kernel does not have; running + # without it would train a sink-less draft that is exported as sink-enabled. + raise NotImplementedError( + "dflash_use_flex_attention does not support dflash_attention_sink; " + "unset one of them." + ) + dropout = 0.0 if not self.training else self.attention_dropout + if dropout: + raise ValueError( + "FlexAttention path does not support attention_dropout > 0 " + f"(got {dropout}); unset dflash_use_flex_attention." + ) + attn_output = flex_attention_forward(q, k, v, attention_mask, self.scaling) + elif self.attention_sink_bias is not None: if self.sliding_window is not None: # The eager sink path applies only the caller-supplied mask; a per-layer # window from config.layer_types would be silently dropped. DFlash windows @@ -410,6 +410,7 @@ def __init__(self, config): ) self.norm = _NORM_CLS(config.hidden_size, eps=config.rms_norm_eps) self._rotary_config = config # Used by _maybe_init_rotary_emb + self._compiled_body = None # set by compile_body() # Explicit weight init is needed because DFlashModule is instantiated via # mtsp.convert() AFTER the base model's post_init() has already run, so HF's @@ -437,9 +438,28 @@ def _init_weights(self, config): def forward(self, noise_embedding, target_hidden, position_ids, attention_mask=None): """Forward with feature fusion, KV injection, and position embeddings.""" + # Outside the compiled body: lazy rotary init mutates the module. + self._maybe_init_rotary_emb(device=noise_embedding.device) + return self._body()(noise_embedding, target_hidden, position_ids, attention_mask) + + def compile_body(self): + """Inductor-compile the draft stack; ``_body`` runs it in training only. + + ``dynamic=False`` relies on the pinned block count, while generation runs at varying + lengths and would recompile for each. + """ + self._compiled_body = torch.compile(self._forward_body, dynamic=False) + + def _body(self): + """The compiled draft stack while training, if ``compile_body`` ran; eager otherwise.""" + if self.training and self._compiled_body is not None: + return self._compiled_body + return self._forward_body + + def _forward_body(self, noise_embedding, target_hidden, position_ids, attention_mask): + """Feature fusion, rotary selection, the decoder stack, and the final norm.""" hidden_states = noise_embedding target_hidden = self.hidden_norm(self.fc(target_hidden)) - self._maybe_init_rotary_emb(device=hidden_states.device) position_embeddings = self.rotary_emb(hidden_states, position_ids) for layer in self.layers: diff --git a/tests/gpu/torch/speculative/plugins/test_hf_dflash.py b/tests/gpu/torch/speculative/plugins/test_hf_dflash.py index 0cb740c0e9f..65f62721f21 100644 --- a/tests/gpu/torch/speculative/plugins/test_hf_dflash.py +++ b/tests/gpu/torch/speculative/plugins/test_hf_dflash.py @@ -220,6 +220,7 @@ def _make_base_model_outputs(self, model, bsz): def test_offline_forward_returns_loss(self, offline_model): """Offline forward consumes precomputed base_model_outputs and returns a finite loss.""" + assert offline_model.dflash_self_logit_distillation # KD: teacher projected from hidden bsz = 2 input_ids = torch.randint(0, offline_model.config.vocab_size, (bsz, SEQ_LEN), device="cuda") attention_mask = torch.ones(bsz, SEQ_LEN, dtype=torch.long, device="cuda") @@ -234,23 +235,6 @@ def test_offline_forward_returns_loss(self, offline_model): assert output.loss.requires_grad assert torch.isfinite(output.loss).item() - def test_offline_forward_self_logit_distillation_recomputes_logits(self, offline_model): - """When base_model_logits is absent, self-distillation path computes them from hidden states.""" - assert offline_model.dflash_self_logit_distillation - bsz = 2 - input_ids = torch.randint(0, offline_model.config.vocab_size, (bsz, SEQ_LEN), device="cuda") - attention_mask = torch.ones(bsz, SEQ_LEN, dtype=torch.long, device="cuda") - base_model_outputs = self._make_base_model_outputs(offline_model, bsz) - - output = offline_model( - input_ids=input_ids, - attention_mask=attention_mask, - base_model_outputs=base_model_outputs, - ) - assert hasattr(output, "logits") - assert output.logits is not None - assert torch.isfinite(output.loss).item() - @pytest.mark.skipif( torch.cuda.device_count() < 2, reason="needs 2 GPUs to shard the base model across devices" @@ -309,3 +293,110 @@ def test_generate_gathers_sharded_hidden_states(self): assert base_token.device == input_ids.device assert draft_tokens.device == input_ids.device assert draft_tokens.shape == (1, 2) + + +class TestDFlashFlexAttentionGPU: + """FlexAttention path must match the dense-mask SDPA path it replaces. + + The model is wider than the rest of this file because FlexAttention's Triton templates + need head_dim >= 16 (``get_tiny_llama`` defaults to head_dim 2). + """ + + SEQ = 64 + BASE_KWARGS = { + "num_hidden_layers": 4, + "hidden_size": 128, + "num_attention_heads": 2, + "num_key_value_heads": 1, + "intermediate_size": 64, + "max_position_embeddings": 256, + "vocab_size": 64, + } + + @classmethod + def _model(cls, **cfg_overrides): + model = get_tiny_llama(**cls.BASE_KWARGS) + config = get_dflash_config() + config.update(cfg_overrides) + mtsp.convert(model, [("dflash", config)]) + return model.cuda().train() + + @classmethod + def _pair(cls, **cfg_overrides): + """A dense model and a flex model with identical draft weights.""" + dense = cls._model(**cfg_overrides) + flex = cls._model(dflash_use_flex_attention=True, **cfg_overrides) + flex.dflash_module.load_state_dict(dense.dflash_module.state_dict()) + return dense, flex + + @classmethod + def _inputs(cls, bsz=2): + input_ids = torch.randint(0, cls.BASE_KWARGS["vocab_size"], (bsz, cls.SEQ), device="cuda") + attention_mask = torch.ones(bsz, cls.SEQ, dtype=torch.long, device="cuda") + return input_ids, attention_mask + + def test_mask_builder_returns_block_mask(self): + """With the flag on, the mask builder hands back a BlockMask, not a dense tensor.""" + pytest.importorskip("torch.nn.attention.flex_attention") + from modelopt.torch.speculative.plugins.dflash_flex_attention import is_block_mask + + model = self._model(dflash_use_flex_attention=True) + mask = model._build_draft_attention_mask( + self.SEQ, + torch.tensor([[4, 8]], device="cuda"), + torch.tensor([[True, True]], device="cuda"), + 2, + torch.float32, + torch.device("cuda"), + window=None, + ) + assert is_block_mask(mask) + assert not torch.is_tensor(mask) + + @pytest.mark.parametrize("window", [None, 8]) + def test_matches_dense_mask_path(self, window): + """Loss agrees with the dense path to bf16 tolerance, with and without SWA.""" + pytest.importorskip("torch.nn.attention.flex_attention") + overrides = {} if window is None else {"dflash_swa_window_size": window} + dense, flex = self._pair(**overrides) + input_ids, attention_mask = self._inputs() + + # Anchors are resampled every forward: same seed, so only the kernel differs. + torch.manual_seed(1234) + out_dense = dense(input_ids=input_ids, attention_mask=attention_mask) + torch.manual_seed(1234) + out_flex = flex(input_ids=input_ids, attention_mask=attention_mask) + + torch.testing.assert_close(out_flex.loss, out_dense.loss, rtol=2e-2, atol=2e-2) + + def test_matches_dense_mask_path_with_invalid_blocks(self): + """Fully-masked query rows (invalid blocks), a softmax over nothing, match too.""" + pytest.importorskip("torch.nn.attention.flex_attention") + dense, flex = self._pair() + input_ids, attention_mask = self._inputs() + # Row 1 has few supervised positions, so its trailing blocks are invalid. + labels = input_ids.clone() + labels[1, : self.SEQ - BLOCK_SIZE] = -100 + + torch.manual_seed(7) + out_dense = dense(input_ids=input_ids, attention_mask=attention_mask, labels=labels) + torch.manual_seed(7) + out_flex = flex(input_ids=input_ids, attention_mask=attention_mask, labels=labels) + + assert torch.isfinite(out_flex.loss), "flex produced a non-finite loss" + torch.testing.assert_close(out_flex.loss, out_dense.loss, rtol=2e-2, atol=2e-2) + + def test_backward_produces_finite_grads(self): + """The flex path is differentiable and its grads are finite.""" + pytest.importorskip("torch.nn.attention.flex_attention") + flex = self._model(dflash_use_flex_attention=True) + input_ids, attention_mask = self._inputs() + + flex(input_ids=input_ids, attention_mask=attention_mask).loss.backward() + grads = [ + p.grad + for p in flex.dflash_module.parameters() + if p.requires_grad and p.grad is not None + ] + assert grads, "no draft gradients were produced" + assert all(torch.isfinite(g).all() for g in grads) diff --git a/tests/unit/torch/speculative/plugins/test_hf_dflash.py b/tests/unit/torch/speculative/plugins/test_hf_dflash.py index a7887bea37c..d32ff987705 100644 --- a/tests/unit/torch/speculative/plugins/test_hf_dflash.py +++ b/tests/unit/torch/speculative/plugins/test_hf_dflash.py @@ -39,6 +39,7 @@ import modelopt.torch.speculative.plugins.hf_dflash as hf_dflash from modelopt.torch.speculative.plugins.hf_dflash import ( DFlashAttention, + DFlashBaseModelOutput, DFlashModule, HFDFlashModel, _dpace_position_weights, @@ -265,28 +266,36 @@ def _converted_model(self, objective, **overrides): def test_compute_loss_dpace_branch(self): """Default dpace objective produces a finite loss and valid accuracy.""" - model = self._converted_model("dpace") + model = self._converted_model("dpace", dflash_self_logit_distillation=False) loss, acc = model._compute_loss(*self._make_inputs()) assert torch.isfinite(loss).item() and loss.item() > 0 assert 0.0 <= acc <= 1.0 def test_compute_loss_decay_branch(self): """The static-decay objective path also produces a finite loss.""" - model = self._converted_model("decay", dflash_loss_decay_factor=4.0) + model = self._converted_model( + "decay", dflash_loss_decay_factor=4.0, dflash_self_logit_distillation=False + ) loss, acc = model._compute_loss(*self._make_inputs()) assert torch.isfinite(loss).item() and loss.item() > 0 assert 0.0 <= acc <= 1.0 def test_compute_loss_dpace_kd_branch(self): - """dpace + KD (base_logits given): confidences use a dedicated no_grad CE pass.""" + """dpace + KD: confidences use a dedicated no_grad CE pass.""" vocab = 32 model = self._converted_model("dpace") inputs = self._make_inputs(vocab=vocab) - base_logits = torch.randn(1, SEQ_LEN, vocab) - loss, acc = model._compute_loss(*inputs, base_logits=base_logits) + teacher = DFlashBaseModelOutput(None, logits=torch.randn(1, SEQ_LEN, vocab)) + loss, acc = model._compute_loss(*inputs, base_outputs=teacher) assert torch.isfinite(loss).item() assert 0.0 <= acc <= 1.0 + def test_kd_without_base_outputs_raises(self): + """KD is decided by the config, so a missing teacher cannot fall back to CE.""" + model = self._converted_model("dpace") + with pytest.raises(ValueError, match="base_outputs"): + model._compute_loss(*self._make_inputs()) + class TestDFlashSaveRestore: """Test DFlash model save and restore.""" @@ -415,6 +424,67 @@ def test_no_sliding_window_without_config(self): assert attn.sliding_window is None +class TestDraftMaskRule: + """The dense mask and the FlexAttention BlockMask evaluate one rule, ``_draft_mask_mod``.""" + + SEQ_LEN, BLOCK_SIZE = 16, 4 + + def _model(self, attention="bidirectional"): + model = get_tiny_llama(num_hidden_layers=4) + config = get_dflash_config(block_size=self.BLOCK_SIZE) + config["dflash_draft_attention"] = attention + mtsp.convert(model, [("dflash", config)]) + return model + + def _visible(self, model, anchors, keep, window=None): + mask = model._build_draft_attention_mask( + self.SEQ_LEN, anchors, keep, anchors.shape[1], torch.float32, "cpu", window=window + ) + return mask > torch.finfo(torch.float32).min / 2 + + @pytest.mark.parametrize("attention", ["bidirectional", "causal"]) + def test_matches_a_hand_written_mask(self, attention): + """Context strictly before each anchor, then the query's own block, nothing else.""" + anchors = (6, 9) + visible = self._visible( + self._model(attention), torch.tensor([anchors]), torch.tensor([[True, True]]) + ) + n_q = len(anchors) * self.BLOCK_SIZE + expected = torch.zeros(n_q, self.SEQ_LEN + n_q, dtype=torch.bool) + for block, anchor in enumerate(anchors): + for i in range(self.BLOCK_SIZE): + row = block * self.BLOCK_SIZE + i + expected[row, :anchor] = True + start = self.SEQ_LEN + block * self.BLOCK_SIZE + expected[ + row, start : start + (i + 1 if attention == "causal" else self.BLOCK_SIZE) + ] = True + assert torch.equal(visible[0, 0], expected) + + def test_dropped_block_sees_nothing(self): + visible = self._visible( + self._model(), torch.tensor([[6, 9]]), torch.tensor([[True, False]]) + ) + assert visible[0, 0, : self.BLOCK_SIZE].any(dim=-1).all() + assert not visible[0, 0, self.BLOCK_SIZE :].any() + + @pytest.mark.parametrize("attention", ["bidirectional", "causal"]) + @pytest.mark.parametrize("window", [None, 6]) + def test_dense_mask_matches_the_vmapped_rule(self, attention, window): + """Broadcast evaluation (dense) and vmap evaluation (flex) of the rule agree.""" + from torch.nn.attention.flex_attention import create_mask + + anchors = torch.tensor([[2, 7, 11], [5, 0, 9]]) + keep = torch.tensor([[True, True, False], [True, False, True]]) + mask_mod = hf_dflash._draft_mask_mod( + self.SEQ_LEN, anchors, keep, self.BLOCK_SIZE, window, causal=attention == "causal" + ) + q_len = anchors.shape[1] * self.BLOCK_SIZE + vmapped = create_mask(mask_mod, 2, 1, q_len, self.SEQ_LEN + q_len, device="cpu") + dense = self._visible(self._model(attention), anchors, keep, window) + assert torch.equal(dense, vmapped) + + class TestDFlashSwaMask: """Test all-layer non-causal sliding-window attention mask (MiMo-style).""" @@ -1142,3 +1212,276 @@ def test_gradients_are_bit_identical_with_and_without(self): assert grads[False] and grads[False].keys() == grads[True].keys() for name, grad in grads[False].items(): assert torch.equal(grad, grads[True][name]), name + + +def _legacy_sample_anchor_positions(model, seq_len, loss_mask, device): + """Verbatim copy of the pre-static implementation, kept as the semantic reference.""" + bs = model.dflash_block_size + bsz = loss_mask.shape[0] + max_anchor = max(seq_len - bs, 0) + num_anchors = getattr(model, "_num_anchors", 512) + + valid = loss_mask[:, : max_anchor + 1] > 0.5 + valid_counts = valid.sum(dim=1) + max_n = min(num_anchors, int(valid_counts.max().item()) - 1) + + if max_n <= 0: + return ( + torch.zeros(bsz, 1, dtype=torch.long, device=device), + torch.zeros(bsz, 1, dtype=torch.bool, device=device), + ) + + indices = torch.arange(max_anchor + 1, device=device).unsqueeze(0).expand(bsz, -1) + masked_indices = torch.where(valid, indices, torch.tensor(seq_len + 1, device=device)) + random_vals = torch.rand(bsz, max_anchor + 1, device=device) + random_vals = torch.where(valid, random_vals, torch.tensor(2.0, device=device)) + _, sorted_idx = random_vals.sort(dim=1) + gathered = torch.gather(masked_indices, 1, sorted_idx) + anchors = gathered[:, :max_n].sort(dim=1).values + keep = torch.arange(max_n, device=device).unsqueeze(0) < valid_counts.unsqueeze(1).clamp( + max=max_n + ) + anchors = torch.where(keep, anchors, torch.tensor(0, dtype=torch.long, device=device)) + return anchors, keep + + +class TestAnchorSamplingStaticShape: + """n_blocks is fixed by the config, and the anchors it picks are still the legacy ones.""" + + DEVICE = torch.device("cpu") + + @staticmethod + def _model(num_anchors): + model = get_tiny_llama(num_hidden_layers=4) + config = get_dflash_config(block_size=BLOCK_SIZE) + config["dflash_num_anchors"] = num_anchors + config["dflash_self_logit_distillation"] = False # the padding test calls the CE loss + mtsp.convert(model, [("dflash", config)]) + return model + + @staticmethod + def _loss_mask(lengths, seq_len=SEQ_LEN): + mask = torch.zeros(len(lengths), seq_len) + for row, n in enumerate(lengths): + mask[row, :n] = 1.0 + return mask + + # Each case puts the legacy bound somewhere different: under the cap, ragged across + # rows, at 1, and in the two degenerate cases the old code special-cased. + @pytest.mark.parametrize( + "lengths", + [(13, 13), (13, 5), (9, 9), (5, 3), (2, 1), (1, 1), (0, 0), (13, 0)], + ) + def test_matches_legacy_sampling(self, lengths): + """Same seed in, same anchors and same keep mask out -- bitwise.""" + model = self._model(num_anchors=8) + loss_mask = self._loss_mask(lengths) + + torch.manual_seed(1234) + legacy_anchors, legacy_keep = _legacy_sample_anchor_positions( + model, SEQ_LEN, loss_mask, self.DEVICE + ) + torch.manual_seed(1234) + anchors, keep = model._sample_anchor_positions(SEQ_LEN, loss_mask, self.DEVICE) + + n_old = legacy_keep.shape[1] + assert keep.shape[1] >= n_old + assert torch.equal(keep[:, :n_old], legacy_keep) + assert torch.equal(anchors[:, :n_old], legacy_anchors) + # Everything past the legacy bound is inert padding. + assert not keep[:, n_old:].any() + assert not anchors[:, n_old:].any() + + def test_sort_is_truncated_before_it_widens(self): + """The surplus columns must not pull unsampled anchors in among the kept ones.""" + model = self._model(num_anchors=8) + # One long row, so the legacy bound (valid_counts.max() - 1) sits below the static + # width and there really are surplus columns to get wrong. + loss_mask = self._loss_mask((6, 6)) + torch.manual_seed(7) + legacy_anchors, legacy_keep = _legacy_sample_anchor_positions( + model, SEQ_LEN, loss_mask, self.DEVICE + ) + torch.manual_seed(7) + anchors, keep = model._sample_anchor_positions(SEQ_LEN, loss_mask, self.DEVICE) + assert legacy_keep.shape[1] < keep.shape[1], "fixture no longer exercises truncation" + kept = anchors[keep] + assert torch.equal(kept, legacy_anchors[legacy_keep]) + + def test_shape_is_independent_of_batch_contents(self): + """The whole point: one shape, therefore one compile.""" + model = self._model(num_anchors=8) + shapes = { + model._sample_anchor_positions(SEQ_LEN, self._loss_mask(lengths), self.DEVICE)[0].shape + for lengths in [(13, 13), (13, 5), (9, 9), (5, 3), (2, 1), (1, 1), (0, 0)] + } + assert len(shapes) == 1, f"n_blocks still varies with the batch: {shapes}" + + def test_shape_is_min_of_num_anchors_and_sequence(self): + anchors, _ = self._model(num_anchors=4)._sample_anchor_positions( + SEQ_LEN, self._loss_mask((13, 13)), self.DEVICE + ) + assert anchors.shape[1] == 4, "num_anchors should bind here" + anchors, _ = self._model(num_anchors=512)._sample_anchor_positions( + SEQ_LEN, self._loss_mask((13, 13)), self.DEVICE + ) + assert anchors.shape[1] == SEQ_LEN - BLOCK_SIZE + 1, "the sequence should bind here" + + def test_sampling_does_not_sync_on_the_batch(self): + """Anchor sampling must not read the batch back to the host.""" + import ast + import inspect + import textwrap + + tree = ast.parse( + textwrap.dedent(inspect.getsource(hf_dflash.HFDFlashModel._sample_anchor_positions)) + ) + called = { + node.func.attr + for node in ast.walk(tree) + if isinstance(node, ast.Call) and isinstance(node.func, ast.Attribute) + } + assert not called & {"item", "tolist"}, f"host sync in anchor sampling: {called}" + + def test_trailing_padding_blocks_do_not_change_the_loss(self): + """The padding the static shape introduces is weightless in the loss and accuracy. + + Exact in math, but the reduction runs over more elements, so compare to tolerance. + """ + model = self._model(num_anchors=8) + vocab, bsz, n_blocks, pad = 32, 1, 2, 3 + torch.manual_seed(0) + input_ids = torch.randint(0, vocab, (bsz, SEQ_LEN)) + loss_mask = torch.ones(bsz, SEQ_LEN) + logits = torch.randn(bsz, n_blocks * BLOCK_SIZE, vocab) + anchors = torch.tensor([[0, BLOCK_SIZE]])[:, :n_blocks] + keep = torch.ones(bsz, n_blocks) + + base_loss, base_acc = model._compute_loss(logits, input_ids, anchors, keep, loss_mask) + padded_loss, padded_acc = model._compute_loss( + torch.cat([logits, torch.randn(bsz, pad * BLOCK_SIZE, vocab)], dim=1), + input_ids, + torch.cat([anchors, torch.zeros(bsz, pad, dtype=anchors.dtype)], dim=1), + torch.cat([keep, torch.zeros(bsz, pad)], dim=1), + loss_mask, + ) + torch.testing.assert_close(padded_loss, base_loss) + assert padded_acc == pytest.approx(base_acc) + + +class TestTeacherLogits: + """The base distribution is projected only at the rows, and vocab entries, asked for.""" + + @staticmethod + def _model(): + model = get_tiny_llama(num_hidden_layers=4) + mtsp.convert(model, [("dflash", get_dflash_config())]) + return model + + @staticmethod + def _positions(bsz=2): + torch.manual_seed(0) + return torch.randint(0, SEQ_LEN, (bsz, 3, BLOCK_SIZE)) + + @staticmethod + def _rows(full, positions): + return full[torch.arange(positions.shape[0]).view(-1, 1, 1), positions] + + @staticmethod + def _hidden(model, bsz=2): + dtype = model._base_model_lm_head.weight.dtype + return torch.randn(bsz, SEQ_LEN, model.config.hidden_size, dtype=dtype) + + def test_matches_the_full_sequence_projection(self): + model = self._model() + hidden = self._hidden(model) + positions = self._positions() + got = model._teacher_logits(DFlashBaseModelOutput(None, base_hidden=hidden), positions) + want = self._rows(model._base_model_lm_head(hidden), positions) + torch.testing.assert_close(got, want) + + def test_prenorm_hidden_gets_the_base_final_norm(self): + model = self._model() + hidden = self._hidden(model) + positions = self._positions() + outputs = DFlashBaseModelOutput(None, base_hidden=hidden, base_hidden_prenorm=True) + want = self._rows(model._base_model_lm_head(model._base_model_norm(hidden)), positions) + torch.testing.assert_close(model._teacher_logits(outputs, positions), want) + + def test_token_ids_pick_entries_of_the_full_rows(self): + model = self._model() + outputs = DFlashBaseModelOutput(None, base_hidden=self._hidden(model)) + positions = self._positions() + token_ids = torch.randint(0, model.config.vocab_size, (*positions.shape, 5)) + rows = model._teacher_logits(outputs, positions) + torch.testing.assert_close( + model._teacher_logits(outputs, positions, token_ids), rows.gather(-1, token_ids) + ) + + def test_handed_over_logits_are_gathered_not_recomputed(self): + model = self._model() + logits = torch.randn(2, SEQ_LEN, model.config.vocab_size) + outputs = DFlashBaseModelOutput(None, logits=logits) + positions = self._positions() + token_ids = torch.randint(0, model.config.vocab_size, (*positions.shape, 5)) + want = self._rows(logits, positions) + assert torch.equal(model._teacher_logits(outputs, positions), want) + assert torch.equal( + model._teacher_logits(outputs, positions, token_ids), want.gather(-1, token_ids) + ) + + def test_missing_base_distribution_raises(self): + model = self._model() + with pytest.raises(ValueError, match="base_model_hidden_states"): + model._teacher_logits(DFlashBaseModelOutput(None), self._positions()) + + def test_offline_dict_keeps_the_hidden_unprojected(self): + d = { + "aux_hidden_states": torch.randn(1, SEQ_LEN, 8), + "base_model_hidden_states": torch.randn(1, SEQ_LEN, 4), + "base_hidden_prenorm": True, + } + outputs = DFlashBaseModelOutput.from_offline_dict(d) + assert outputs.logits is None + assert outputs.base_hidden is d["base_model_hidden_states"] + assert outputs.base_hidden_prenorm + + def test_kd_loss_matches_full_logits(self): + """KD read from the base hidden equals KD from the same logits built in full.""" + model = self._model() + hidden = self._hidden(model, bsz=1) + inputs = TestDPaceLossIntegration._make_inputs(vocab=model.config.vocab_size) + full = DFlashBaseModelOutput(None, logits=model._base_model_lm_head(hidden)) + loss_full, _ = model._compute_loss(*inputs, base_outputs=full) + loss, _ = model._compute_loss( + *inputs, base_outputs=DFlashBaseModelOutput(None, base_hidden=hidden) + ) + torch.testing.assert_close(loss, loss_full) + + def test_offline_step_never_projects_the_full_sequence(self): + """Offline/streaming with KD on, lm_head only ever sees draft rows and teacher rows.""" + model = get_tiny_llama(num_hidden_layers=4) + model.config.num_orig_hidden_layers = 4 + mtsp.convert(model, [("dflash", get_dflash_config(offline=True))]) + assert model.dflash_self_logit_distillation + model.train() + seen = [] + model._base_model_lm_head.register_forward_hook( + lambda _module, args, _out: seen.append(tuple(args[0].shape)) + ) + bsz, hidden = 2, model.config.hidden_size + dtype = next(model.dflash_module.parameters()).dtype + base_model_outputs = { + "aux_hidden_states": torch.randn( + bsz, SEQ_LEN, len(model.target_layer_ids) * hidden, dtype=dtype + ), + "base_model_hidden_states": torch.randn(bsz, SEQ_LEN, hidden, dtype=dtype), + } + input_ids = torch.randint(1, model.config.vocab_size, (bsz, SEQ_LEN)) + model( + input_ids=input_ids, + attention_mask=torch.ones_like(input_ids), + base_model_outputs=base_model_outputs, + ).loss.backward() + assert len(seen) == 2, seen # the draft's logits, then the KD teacher rows + assert (bsz, SEQ_LEN, hidden) not in seen, seen diff --git a/tests/unit/torch/speculative/plugins/test_hf_dflash2.py b/tests/unit/torch/speculative/plugins/test_hf_dflash2.py index 931a4a3ee46..4a1112eb174 100644 --- a/tests/unit/torch/speculative/plugins/test_hf_dflash2.py +++ b/tests/unit/torch/speculative/plugins/test_hf_dflash2.py @@ -352,11 +352,15 @@ def test_selector_loss_increases_total_loss(self): 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. + Guards the selector's target alignment: with the gold token off by a position the + loss still falls, but the backbone's top-k rarely holds that token, so coverage + stays low. """ model = get_tiny_llama(num_hidden_layers=4) mtsp.convert(model, [("dflash", _get_dflash2_config())]) + # fp32: CPUs without AVX-512 have no fast bf16 matmul, and the draft keeps Qwen3's + # default 22016-wide MLP. + model.float() model.train() input_ids, attention_mask, labels = _make_batch(model.dflash_config.vocab_size) @@ -369,6 +373,7 @@ def test_overfits_a_single_batch(self): assert out.train_acc[0][0] > 0.9 assert out["selector_metrics"]["selector_accuracy"].item() > 0.9 + assert out["selector_metrics"]["selector_coverage"].item() > 0.9 class TestCandidateSelectorAlignment: diff --git a/tests/unit/torch/speculative/plugins/test_hf_dspark.py b/tests/unit/torch/speculative/plugins/test_hf_dspark.py index 3686788ad90..3afed95fdbd 100644 --- a/tests/unit/torch/speculative/plugins/test_hf_dspark.py +++ b/tests/unit/torch/speculative/plugins/test_hf_dspark.py @@ -37,8 +37,9 @@ 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_dspark import HFDSparkModel +from modelopt.torch.speculative.plugins.hf_dspark import HFDSparkModel, _tvd_chunk, _tvd_per_token from modelopt.torch.speculative.plugins.modeling_dflash import ( + DFlashBaseModelOutput, DFlashModule, build_target_layer_ids, repeat_kv, @@ -770,3 +771,184 @@ def test_export_round_trips_explicit_ids(self, tmp_path): with open(tmp_path / "exp" / "config.json") as f: cfg = json.load(f) assert cfg["dflash_config"]["target_layer_ids"] == [0, 7] + + +class TestTvdPerTokenChunking: + """``_tvd_per_token`` must be chunk-size-invariant, and must chunk via ``split``. + + Calls the helper directly: the model fixtures never reach a second chunk. + """ + + @staticmethod + def _run(chunk_size, n=12, vocab=32): + """Forward + backward through ``_tvd_per_token`` from a fixed seed.""" + torch.manual_seed(0) + final = torch.randn(n, vocab, requires_grad=True) + teacher = torch.randn(n, vocab) + out = _tvd_per_token(final, teacher, chunk_size=chunk_size) + # Weight the rows unequally, so a cat that reassembles the chunks in the wrong + # order cannot cancel out in the reduction. + (out * torch.arange(1, n + 1, dtype=out.dtype)).sum().backward() + return out.detach(), final.grad.detach() + + # 12 rows: 5 and 7 leave a ragged last chunk (5+5+2, 7+5); 13 and 1024 exceed n. + @pytest.mark.parametrize("chunk_size", [1, 2, 3, 5, 7, 11, 12, 13, 1024]) + def test_chunk_size_invariant(self, chunk_size): + ref_out, ref_grad = self._run(12) # one chunk == the unchunked reference + out, grad = self._run(chunk_size) + assert torch.equal(out, ref_out), f"TVD value changed at chunk_size={chunk_size}" + assert torch.equal(grad, ref_grad), f"TVD grad changed at chunk_size={chunk_size}" + + def test_chunks_via_split_not_slice(self): + """Pin the split itself: slicing gives identical values, only the graph differs.""" + final = torch.randn(8, 4, requires_grad=True) + teacher = torch.randn(8, 4) + out = _tvd_per_token(final, teacher, chunk_size=2) + + # `alive` keeps visited nodes referenced: `.next_functions` returns fresh wrappers, + # and a freed wrapper's id() can be reused, which would truncate the walk. + seen, visited, alive, stack = set(), set(), [], [out.grad_fn] + while stack: + fn = stack.pop() + if fn is None or id(fn) in visited: + continue + visited.add(id(fn)) + alive.append(fn) + seen.add(type(fn).__name__) + stack.extend(nxt for nxt, _ in fn.next_functions) + + assert any(name.startswith("SplitBackward") for name in seen), ( + f"expected a SplitBackward node in the graph, saw: {sorted(seen)}" + ) + assert not any(name.startswith("SliceBackward") for name in seen), ( + f"chunking regressed to per-chunk slicing, saw: {sorted(seen)}" + ) + + def test_no_grad_path_matches_grad_path(self): + """The ``requires_grad=False`` branch skips checkpointing; values must not move.""" + torch.manual_seed(0) + final = torch.randn(12, 32) + teacher = torch.randn(12, 32) + with torch.no_grad(): + plain = _tvd_per_token(final, teacher, chunk_size=5) + grad_out = _tvd_per_token(final.requires_grad_(True), teacher, chunk_size=5) + assert torch.equal(plain, grad_out.detach()) + + +def _draft_args(model, n_blocks=2, bsz=1): + """Build the draft module's inputs with the model's own helpers, as training does.""" + m = model.dflash_module + dt = m.fc.weight.dtype # the draft carries the base model's dtype, not fp32 + torch.manual_seed(0) + input_ids = torch.randint(1, model.dflash_config.vocab_size, (bsz, SEQ_LEN)) + anchors = (torch.arange(n_blocks).unsqueeze(0) * BLOCK_SIZE).expand(bsz, -1).contiguous() + keep = torch.ones(bsz, n_blocks, dtype=torch.bool) + noise = model._build_noise_embedding(input_ids, anchors, keep, n_blocks).to(dt) + target = torch.randn(bsz, SEQ_LEN, m.fc.in_features, dtype=dt) + pos = model._build_position_ids(SEQ_LEN, anchors, input_ids.device) + mask = model._build_draft_attention_mask( + SEQ_LEN, anchors, keep, n_blocks, dt, input_ids.device, window=None + ) + return input_ids, anchors, keep, (noise, target, pos, mask) + + +def _dspark_model(use_compile=False, **cfg_kwargs): + model = get_tiny_llama(num_hidden_layers=4) + cfg = _get_dspark_config(**cfg_kwargs) + cfg["dflash_use_torch_compile"] = use_compile + mtsp.convert(model, [("dflash", cfg)]) + model.train() + return model + + +class TestDraftStackCompile: + """The draft stack is Inductor-compiled only when asked for, and only while training.""" + + def test_flag_off_keeps_the_eager_body(self): + m = _dspark_model(use_compile=False).dflash_module + assert m._body() == m._forward_body + + def test_eval_keeps_the_eager_body(self): + """Generation runs at varying lengths, which dynamic=False would recompile for.""" + m = _dspark_model(use_compile=True).dflash_module + m.eval() + assert m._body() == m._forward_body + + def test_flag_reaches_the_tvd_chunk(self): + assert _dspark_model(use_compile=False)._tvd_chunk_fn is _tvd_chunk + assert _dspark_model(use_compile=True)._tvd_chunk_fn is not _tvd_chunk + + def test_compiled_matches_eager(self): + """Same weights, same inputs, both bodies -- same hidden states.""" + # fp32: in bf16, Inductor and eager round at different points, too far apart to + # compare tightly. + model = _dspark_model(use_compile=True).float() + m = model.dflash_module + _, _, _, args = _draft_args(model) + m._maybe_init_rotary_emb(device=args[0].device) # normally done by forward() + + with torch.no_grad(): + ref = m._forward_body(*args) + got = m._body()(*args) + assert m._body() != m._forward_body, "fixture did not actually compile" + # Inductor may reassociate, so not bitwise. + torch.testing.assert_close(got, ref, rtol=1e-4, atol=1e-4) + + +class TestDdpGradientCoverage: + """Every draft parameter must get a gradient on every batch, degenerate ones included. + + Needed for ddp_find_unused_parameters=false; the confidence head is the one at risk, + since it is not behind final_logits. + """ + + @staticmethod + def _model(): + return _dspark_model(use_confidence_head=True, confidence_alpha=1.0) + + @staticmethod + def _ungraded(model): + return [ + n + for n, p in model.dflash_module.named_parameters() + if p.requires_grad and p.grad is None + ] + + def test_batch_with_no_valid_anchor(self): + """Nothing is a training target, so forward() returns before building the draft.""" + model = self._model() + torch.manual_seed(0) + input_ids = torch.randint(1, model.dflash_config.vocab_size, (2, SEQ_LEN)) + out = model( + input_ids=input_ids, + attention_mask=torch.ones_like(input_ids), + labels=torch.full_like(input_ids, -100), + ) + out.loss.backward() + assert not self._ungraded(model), f"no gradient for {self._ungraded(model)}" + + def test_loss_branch_with_zero_total_weight(self): + """Anchors exist but every label position is masked, so the three terms are skipped.""" + model = self._model() + m = model.dflash_module + n_blocks, bsz = 2, 1 + input_ids, anchors, _, args = _draft_args(model, n_blocks=n_blocks, bsz=bsz) + hidden = m(*args) + vocab = model.dflash_config.vocab_size + backbone_logits = torch.randn(bsz, n_blocks * BLOCK_SIZE, vocab, dtype=hidden.dtype) + base_outputs = DFlashBaseModelOutput(args[1]) # no base distribution: not read here + final_logits, confidence_logits = model._apply_markov_head( + hidden, backbone_logits, input_ids, anchors, n_blocks + ) + loss, _, _ = model._compute_dspark_loss( + backbone_logits, + final_logits, + confidence_logits, + input_ids, + anchors, + torch.ones(bsz, n_blocks), # blocks are kept ... + torch.zeros(bsz, SEQ_LEN), # ... but no label position carries weight + base_outputs, + ) + loss.backward() + assert not self._ungraded(model), f"no gradient for {self._ungraded(model)}" diff --git a/tests/unit/torch/speculative/plugins/test_hf_lilicorr.py b/tests/unit/torch/speculative/plugins/test_hf_lilicorr.py index d1402987550..0d132f9ce49 100644 --- a/tests/unit/torch/speculative/plugins/test_hf_lilicorr.py +++ b/tests/unit/torch/speculative/plugins/test_hf_lilicorr.py @@ -38,7 +38,7 @@ from modelopt.torch.speculative.config import DFLASH_DEFAULT_CFG from modelopt.torch.speculative.plugins.hf_dflash import HFDFlashModel from modelopt.torch.speculative.plugins.hf_lilicorr import HFLiLiCorrModel -from modelopt.torch.speculative.plugins.modeling_dflash import DFlashModule +from modelopt.torch.speculative.plugins.modeling_dflash import DFlashBaseModelOutput, DFlashModule from modelopt.torch.speculative.plugins.modeling_lilicorr import LiLiCorrModule BLOCK_SIZE = 4 @@ -347,7 +347,9 @@ def test_calibration_loss_is_minimal_when_the_head_matches_the_target(self): inputs = { "candidate_ids": candidate_ids, "anchor_positions": torch.tensor([[0, BLOCK_SIZE]]), - "target_logits": target_logits, + "base_outputs": DFlashBaseModelOutput( + target_hidden=torch.zeros(1, SEQ_LEN, 1), logits=target_logits + ), } at_target = model._candidate_target_logits(**inputs, requested_by="test") diff --git a/tests/unit/torch/speculative/plugins/test_hf_streaming_dataset.py b/tests/unit/torch/speculative/plugins/test_hf_streaming_dataset.py index 00b5a1f6ab5..fed02d88b12 100644 --- a/tests/unit/torch/speculative/plugins/test_hf_streaming_dataset.py +++ b/tests/unit/torch/speculative/plugins/test_hf_streaming_dataset.py @@ -430,6 +430,49 @@ def test_lapped_slot_is_treated_as_miss(monkeypatch): ds[0] +def _first_done_raising_handler(seq, n_layers, hidden, calls): + """Sidecar handler whose first /done raises a transport error; ``calls`` logs the paths.""" + inner = _rdma_sidecar_handler(seq, n_layers, hidden) + + def handler(request: httpx.Request) -> httpx.Response: + calls.append(request.url.path) + if request.url.path == "/done" and calls.count("/done") == 1: + raise httpx.ConnectError("simulated sidecar failure") + return inner(request) + + return handler + + +def test_failed_done_resamples_instead_of_returning_the_read(monkeypatch): + """The discard is a resample, not a hard failure: the next entry is fetched and returned.""" + seq, n_layers, hidden = 8, 3, 16 + calls: list[str] = [] + _mock_rdma( + monkeypatch, + _first_done_raising_handler(seq, n_layers, hidden, calls), + ) + + ds = EagleVllmStreamingDataset( + entries=[ + {"conversation_id": f"c-{i}", "messages": [{"role": "user", "content": "x"}]} + for i in range(2) + ], + tokenizer=_tokenizer_returning(seq), + config=EagleVllmStreamingConfig( + server_urls="http://mock:8000", + model="mock-model", + max_seq_len=seq, + fail_after_consecutive_skips=100, + ), + ) + + batch = ds[0] + assert batch["base_model_hidden_states"].shape == (seq, hidden) + # Two prompts posted: the first read was thrown away, the second is what came back. + assert calls.count("/v1/completions") == 2 + assert calls.count("/done") == 2 + + def test_oversize_server_response_raises(monkeypatch): """If the server captured more tokens than max_seq_len (its connector max_tokens > our recv buffer), reading would silently truncate the slice; fail loud instead so the