From f88735069be4125d140c486a0ef83ddba41d9c7e Mon Sep 17 00:00:00 2001 From: Yihan Wang Date: Thu, 20 Aug 2026 02:52:33 -0700 Subject: [PATCH 1/2] [None][refactor] add Vanilla sparse attention primitives Signed-off-by: Yihan Wang --- .../_torch/attention_backend/vanilla.py | 94 +++++++++++++++++++ .../attention/test_vanilla_attention.py | 23 +++++ 2 files changed, 117 insertions(+) diff --git a/tensorrt_llm/_torch/attention_backend/vanilla.py b/tensorrt_llm/_torch/attention_backend/vanilla.py index ac1ef7aa6d0f..52da0a7d33bf 100644 --- a/tensorrt_llm/_torch/attention_backend/vanilla.py +++ b/tensorrt_llm/_torch/attention_backend/vanilla.py @@ -348,6 +348,57 @@ def _single_request_attn_forward(self, enable_gqa=True, ) + @staticmethod + def _single_token_sparse_attn_forward(q: torch.Tensor, + key_states: torch.Tensor, + value_states: torch.Tensor, + sparse_indices: torch.Tensor, + indices_block_size: int, + qk_scale: float) -> torch.Tensor: + """Run one-token GQA over selected sparse units.""" + if indices_block_size <= 0: + raise ValueError("indices_block_size must be positive") + if sparse_indices.ndim != 2: + raise ValueError( + "sparse_indices must have shape [num_kv_heads, topk]") + num_kv_heads = key_states.shape[1] + if sparse_indices.shape[0] != num_kv_heads: + raise ValueError( + "sparse_indices and key_states must have matching KV heads") + if q.shape[0] % num_kv_heads != 0: + raise ValueError("Query heads must be divisible by KV heads") + if torch.any(sparse_indices < -1): + raise ValueError("Sparse indices may only use -1 as padding") + + heads_per_kv = q.shape[0] // num_kv_heads + unit_offsets = torch.arange(indices_block_size, + device=sparse_indices.device) + output = torch.empty_like(q, dtype=torch.float32) + for kv_head in range(num_kv_heads): + token_indices = ( + sparse_indices[kv_head, :, None] * indices_block_size + + unit_offsets).flatten() + token_indices = token_indices[(token_indices >= 0) + & + (token_indices < key_states.shape[0])] + if token_indices.numel() == 0: + raise ValueError("Every KV head must select at least one token") + token_indices = torch.unique(token_indices).to(torch.long) + selected_keys = key_states.index_select( + 0, token_indices)[:, kv_head].to(torch.float32) + selected_values = value_states.index_select( + 0, token_indices)[:, kv_head].to(torch.float32) + head_start = kv_head * heads_per_kv + head_end = head_start + heads_per_kv + output[head_start:head_end] = F.scaled_dot_product_attention( + q[None, head_start:head_end, None].to(torch.float32), + selected_keys[None, None], + selected_values[None, None], + scale=qk_scale, + enable_gqa=True, + ).view(heads_per_kv, -1) + return output + def _single_request_forward(self, q, k, @@ -646,6 +697,49 @@ def _mla_forward_generation(self, fused_q: torch.Tensor, return torch.cat(outputs, dim=0) + @staticmethod + def _selected_mla_attention( + query: torch.Tensor, + selected_latent: torch.Tensor, + *, + value_dim: int, + scale: float, + attention_sink: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + """Compute absorbed MLA over selected latent rows.""" + if query.ndim != 2 or selected_latent.ndim != 2: + raise ValueError("Selected MLA query and latent rows must be 2D") + if selected_latent.shape[0] == 0: + raise ValueError("Selected MLA requires at least one latent row") + if query.shape[1] != selected_latent.shape[1]: + raise ValueError( + "Selected MLA query and latent dimensions must match, got " + f"{query.shape[1]} and {selected_latent.shape[1]}") + if value_dim <= 0 or value_dim > selected_latent.shape[1]: + raise ValueError( + f"Selected MLA value_dim must be in [1, {selected_latent.shape[1]}], " + f"got {value_dim}") + + scores = torch.matmul(query.float(), selected_latent.float().T) * scale + scores = scores.float() + if attention_sink is None: + probabilities = F.softmax(scores, dim=-1).to(query.dtype) + else: + if attention_sink.shape != (query.shape[0], ): + raise ValueError( + "Selected MLA attention sink must have shape " + f"[{query.shape[0]}], got {tuple(attention_sink.shape)}") + sink = attention_sink.to(device=query.device, + dtype=torch.float32).unsqueeze(1) + max_score = torch.maximum(scores.amax(dim=-1, keepdim=True), sink) + score_exp = torch.exp(scores - max_score) + denominator = score_exp.sum(dim=-1, keepdim=True) + denominator += torch.exp(sink - max_score) + probabilities = (score_exp / denominator).to(query.dtype) + + return torch.matmul(probabilities, + selected_latent[:, :value_dim].to(query.dtype)) + def _mla_forward_context(self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, metadata: VanillaAttentionMetadata, diff --git a/tests/unittest/_torch/attention/test_vanilla_attention.py b/tests/unittest/_torch/attention/test_vanilla_attention.py index 7de13815c359..10dec61ded40 100644 --- a/tests/unittest/_torch/attention/test_vanilla_attention.py +++ b/tests/unittest/_torch/attention/test_vanilla_attention.py @@ -22,6 +22,29 @@ class TestVanillaAttention(unittest.TestCase): + def test_sparse_gqa_deduplicates_blocks(self): + result = VanillaAttention._single_token_sparse_attn_forward( + torch.zeros(1, 1), + torch.zeros(4, 1, 1), + torch.tensor([0.0, 0.0, 10.0, 10.0]).reshape(4, 1, 1), + torch.tensor([[0, 0, 1]], dtype=torch.int32), + indices_block_size=2, + qk_scale=1.0, + ) + + torch.testing.assert_close(result, torch.tensor([[5.0]])) + + def test_selected_mla_attention_sink(self): + result = VanillaAttention._selected_mla_attention( + torch.zeros(1, 2), + torch.tensor([[2.0, 3.0]]), + value_dim=1, + scale=1.0, + attention_sink=torch.zeros(1), + ) + + torch.testing.assert_close(result, torch.tensor([[1.0]])) + def test_kv_cache_manager_v2_sliding_window_eviction(self): device = torch.device("cuda") dtype = torch.bfloat16 From 6982f237c25ed7629a5e38b2f4fbdcdc2e8d34cb Mon Sep 17 00:00:00 2001 From: Yihan Wang Date: Thu, 20 Aug 2026 18:46:46 -0700 Subject: [PATCH 2/2] [None][feat] add DeepSeek V4 Vanilla sparse attention Signed-off-by: Yihan Wang --- .../sparse/deepseek_v4/__init__.py | 3 + .../sparse/deepseek_v4/module.py | 2 + .../sparse/deepseek_v4/vanilla_backend.py | 560 ++++++++++++++++++ .../attention_backend/sparse/registry.py | 4 + .../_torch/attention/backend_capability.py | 16 +- .../unittest/_torch/attention/backend_case.py | 407 ++++++++++++- .../_torch/attention/model_attn_config.py | 50 +- .../test_deepseek_v4_sparse_mla.py | 423 ++++++++----- .../deepseek_v4/test_deepseek_v4_vanilla.py | 309 ++++++++++ .../sparse/test_sparse_mla_forward.py | 145 +---- .../attention/test_attention_backends.py | 47 +- 11 files changed, 1684 insertions(+), 282 deletions(-) create mode 100644 tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/vanilla_backend.py create mode 100644 tests/unittest/_torch/attention/sparse/deepseek_v4/test_deepseek_v4_vanilla.py diff --git a/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/__init__.py b/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/__init__.py index a822bd90bd40..9b98b819d514 100644 --- a/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/__init__.py +++ b/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/__init__.py @@ -18,6 +18,7 @@ from .indexer import DeepseekV4Indexer from .metadata import DeepseekV4TrtllmAttentionMetadata from .params import DeepseekV4AttentionType, DeepSeekV4MetadataParams, DeepSeekV4Params +from .vanilla_backend import DeepseekV4VanillaAttention, DeepseekV4VanillaIndexer __all__ = [ "DeepSeekV4MetadataParams", @@ -27,4 +28,6 @@ "DeepseekV4Indexer", "DeepseekV4TrtllmAttention", "DeepseekV4TrtllmAttentionMetadata", + "DeepseekV4VanillaAttention", + "DeepseekV4VanillaIndexer", ] diff --git a/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/module.py b/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/module.py index 2c09926c9835..2d4436e8a6c6 100644 --- a/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/module.py +++ b/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/module.py @@ -1153,6 +1153,7 @@ def _indexer_branch(): if self.apply_rotary_emb: assert ctx_position_ids is not None k_pe_ctx = self.apply_rope(q_ctx, k_pe_ctx, ctx_position_ids) + latent_cache_ctx = torch.cat([compressed_kv_ctx, k_pe_ctx], dim=-1) context_o_lora_bmm_input = forward_context_sparse_attn( self, @@ -1181,6 +1182,7 @@ def _indexer_branch(): if self.apply_rotary_emb: assert gen_position_ids is not None k_pe_gen = self.apply_rope(q_gen, k_pe_gen, gen_position_ids) + latent_cache_gen = torch.cat([compressed_kv_gen, k_pe_gen], dim=-1) generation_o_lora_bmm_input = forward_generation_sparse_attn( self, diff --git a/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/vanilla_backend.py b/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/vanilla_backend.py new file mode 100644 index 000000000000..40fce9f9d52a --- /dev/null +++ b/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/vanilla_backend.py @@ -0,0 +1,560 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Vanilla correctness backend for DeepSeek-V4 sparse MLA.""" + +import math +from dataclasses import replace +from typing import TYPE_CHECKING, Optional + +import torch +import torch.nn.functional as F + +from tensorrt_llm._torch.attention_backend.interface import ( + AttentionForwardArgs, + AttentionInputType, + MLAParams, + PositionalEmbeddingParams, + merge_attention_forward_args, +) +from tensorrt_llm._torch.attention_backend.vanilla import VanillaAttention +from tensorrt_llm.models.modeling_utils import QuantConfig + +from ..dsa.indexer import HAS_FAST_HADAMARD +from .indexer import DeepseekV4Indexer +from .metadata import DeepseekV4TrtllmAttentionMetadata +from .params import DeepseekV4AttentionType, DeepSeekV4Params + +if TYPE_CHECKING: + from tensorrt_llm.llmapi.llm_args import SparseAttentionConfig + + +class DeepseekV4VanillaIndexer: + """FP32 golden for DeepSeek-V4 indexer scoring.""" + + def __init__(self, indexer: DeepseekV4Indexer) -> None: + self.indexer = indexer + + @property + def uses_fp4(self) -> bool: + return self.indexer.indexer_k_dtype == "fp4" + + @staticmethod + def _hadamard(x: torch.Tensor) -> torch.Tensor: + width = x.shape[-1] + y = x.float().clone() + stride = 1 + while stride < width: + indices = torch.arange(width, device=x.device) + lower = (indices & stride) == 0 + upper = indices ^ stride + a = y[..., lower].clone() + b = y[..., upper[lower]].clone() + y[..., lower] = a + b + y[..., upper[lower]] = a - b + stride <<= 1 + return (y * width**-0.5).to(x.dtype) + + @staticmethod + def _mxfp4_quant_dequant( + x: torch.Tensor, + *, + round_ties_to_even: bool, + min_amax: float, + ) -> torch.Tensor: + block_size = 32 + fp4_values = torch.tensor( + [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0], + device=x.device, + dtype=torch.float32, + ) + blocks = x.float().reshape(-1, x.shape[-1] // block_size, block_size) + scaled = torch.clamp(blocks.abs().amax(dim=-1, keepdim=True), min=min_amax) / 6.0 + bits = scaled.contiguous().view(torch.int32) + log2_ceil = ((bits >> 23) & 0xFF) - 127 + (bits & 0x7FFFFF).ne(0) + scale = ((log2_ceil + 127) << 23).view(torch.float32) + + normalized = (blocks / scale).abs().clamp(max=6.0) + thresholds = torch.tensor([0.0, 0.25, 0.75, 1.25, 1.75, 2.5, 3.5, 5.0], device=x.device) + value_indices = torch.zeros_like(normalized, dtype=torch.uint8) + for level in range(1, len(thresholds)): + take_upper = normalized > thresholds[level] + if round_ties_to_even: + take_upper |= (normalized == thresholds[level]) & (level % 2 == 0) + value_indices = torch.where( + take_upper, torch.full_like(value_indices, level), value_indices + ) + values = fp4_values[value_indices.to(torch.long)] + values = torch.where(torch.signbit(blocks), -values, values) + return (values * scale).reshape_as(x) + + def _prepare_fp4(self, x: torch.Tensor, *, is_query: bool) -> torch.Tensor: + if HAS_FAST_HADAMARD: + x = self._hadamard(x) + return self._mxfp4_quant_dequant( + x, + round_ties_to_even=not is_query, + min_amax=1.0e-12 if is_query else 6.0 * torch.finfo(torch.float32).tiny, + ) + + def project_query( + self, + qr: torch.Tensor, + position_ids: torch.Tensor, + freqs_cis: torch.Tensor, + ) -> torch.Tensor: + num_tokens = qr.shape[0] + q = F.linear(qr, self.indexer.wq_b.weight) + q = q.view(num_tokens, self.indexer.n_heads, self.indexer.head_dim).unsqueeze(0) + self._apply_rope(q[..., -self.indexer.rope_dim :], freqs_cis[position_ids.long()]) + q = q.squeeze(0) + return self._prepare_fp4(q, is_query=True) if self.uses_fp4 else q + + @staticmethod + def _apply_rope(x: torch.Tensor, freqs_cis: torch.Tensor) -> None: + complex_x = torch.view_as_complex(x.float().unflatten(-1, (-1, 2))) + freqs_cis = freqs_cis.view(1, complex_x.size(1), 1, complex_x.size(-1)) + x.copy_(torch.view_as_real(complex_x * freqs_cis).flatten(-2)) + + def prepare_key(self, k: torch.Tensor) -> torch.Tensor: + return self._prepare_fp4(k, is_query=False) if self.uses_fp4 else k + + def token_weights(self, hidden_states: torch.Tensor) -> torch.Tensor: + weights = F.linear(hidden_states, self.indexer.weights_proj.weight) + return weights.float() * (self.indexer.n_heads**-0.5) + + def scores( + self, + q: torch.Tensor, + k: torch.Tensor, + weights: torch.Tensor, + ) -> torch.Tensor: + head_scores = torch.einsum("hd,kd->hk", q.float(), k.float()) + head_scores = F.relu(head_scores) * self.indexer.softmax_scale + return (head_scores * weights.float().unsqueeze(-1)).sum(dim=0) + + @staticmethod + def topk(scores: torch.Tensor, topk: int) -> torch.Tensor: + indices = torch.full((topk,), -1, dtype=torch.int32, device=scores.device) + if scores.numel() > 0: + valid_topk = min(topk, scores.numel()) + indices[:valid_topk] = torch.topk(scores.float(), valid_topk).indices.to(torch.int32) + return indices + + +class DeepseekV4VanillaAttention(VanillaAttention): + """PyTorch golden for DeepSeek-V4 dual-pool selected attention. + + Selection and compression remain owned by DeepSeek-V4. The Vanilla backend + consumes caller-provided ratio-4 compressed indices, reads the SWA and + compressed pools using their native local index spaces, and reuses + :meth:`VanillaAttention._selected_mla_attention` for the attention math. + """ + + Metadata = DeepseekV4TrtllmAttentionMetadata + + def __init__( + self, + layer_idx: int, + num_heads: int, + head_dim: int, + num_kv_heads: Optional[int] = None, + quant_config: Optional[QuantConfig] = None, + q_scaling: Optional[float] = None, + pos_embd_params: Optional[PositionalEmbeddingParams] = None, + mla_params: Optional[MLAParams] = None, + skip_create_weights_in_init: bool = False, + attention_chunk_size: Optional[int] = None, + sparse_attention_config: Optional["SparseAttentionConfig"] = None, + sparse_params: Optional[DeepSeekV4Params] = None, + dtype: Optional[torch.dtype] = None, + aux_stream: Optional[torch.cuda.Stream] = None, + **kwargs, + ) -> None: + del skip_create_weights_in_init, attention_chunk_size, dtype, aux_stream + if sparse_attention_config is None: + sparse_attention_config = sparse_params + if sparse_attention_config is None: + raise ValueError( + "sparse_attention_config or sparse_params is required for " + "DeepseekV4VanillaAttention" + ) + if sparse_params is None: + sparse_params = sparse_attention_config.to_sparse_params() + if mla_params is None: + raise ValueError("DeepSeek-V4 attention requires MLA parameters") + + mla_params = replace( + mla_params, + v_head_dim=head_dim, + rope_append=False, + ) + super().__init__( + layer_idx, + num_heads, + head_dim, + num_kv_heads=num_kv_heads, + quant_config=quant_config, + q_scaling=q_scaling, + pos_embd_params=pos_embd_params, + mla_params=mla_params, + sparse_params=sparse_params, + **kwargs, + ) + + compress_ratios = sparse_attention_config.compress_ratios + if layer_idx >= len(compress_ratios): + raise ValueError( + f"DeepSeek-V4 layer {layer_idx} has no compression ratio in " + f"a {len(compress_ratios)}-layer configuration" + ) + self.sparse_attention_config = sparse_attention_config + self.compress_ratio = compress_ratios[layer_idx] + self.window_size = sparse_attention_config.window_size + if self.compress_ratio not in (1, 4, 128): + raise ValueError( + "DeepSeek-V4 Vanilla attention supports compression ratios " + f"1, 4, and 128, got {self.compress_ratio}" + ) + + def mla_rope_generation( + self, + q: Optional[torch.Tensor], + q_pe: Optional[torch.Tensor], + latent_cache: torch.Tensor, + metadata: DeepseekV4TrtllmAttentionMetadata, + cu_q_seqlens: torch.Tensor, + cu_kv_seqlens: torch.Tensor, + fmha_scheduler_counter: torch.Tensor, + mla_bmm1_scale: Optional[torch.Tensor], + mla_bmm2_scale: Optional[torch.Tensor], + quant_q_buffer: Optional[torch.Tensor], + out_scale: Optional[torch.Tensor] = None, + kv_norm_weight: Optional[torch.Tensor] = None, + kv_norm_eps: float = 1e-6, + precomputed_cu_seqlens: bool = False, + precomputed_fmha_scheduler: bool = False, + kv_only: bool = False, + kv_done_elsewhere: bool = False, + quant_scale_qkv: Optional[torch.Tensor] = None, + ) -> None: + """No-op counterpart of the fused TRTLLM generation preparation. + + Vanilla does not fuse RoPE or require FMHA scheduler buffers. The MLA + module applies RoPE before invoking this method. + """ + del ( + q, + q_pe, + latent_cache, + metadata, + cu_q_seqlens, + cu_kv_seqlens, + fmha_scheduler_counter, + mla_bmm1_scale, + mla_bmm2_scale, + quant_q_buffer, + out_scale, + kv_norm_weight, + kv_norm_eps, + precomputed_cu_seqlens, + precomputed_fmha_scheduler, + kv_only, + kv_done_elsewhere, + quant_scale_qkv, + ) + + @staticmethod + def _phase_sequence_range( + metadata: DeepseekV4TrtllmAttentionMetadata, + attention_input_type: AttentionInputType, + ) -> tuple[int, int]: + if attention_input_type == AttentionInputType.context_only: + return 0, metadata.num_contexts + if attention_input_type == AttentionInputType.generation_only: + return metadata.num_contexts, metadata.num_seqs + raise ValueError( + "DeepSeek-V4 Vanilla attention requires a context-only or generation-only input" + ) + + @staticmethod + def _page_locations( + block_table: torch.Tensor, + local_indices: torch.Tensor, + tokens_per_block: int, + ) -> tuple[torch.Tensor, torch.Tensor]: + if local_indices.ndim != 1: + raise ValueError( + f"DeepSeek-V4 local cache indices must be 1D, got {tuple(local_indices.shape)}" + ) + + block_indices = torch.div(local_indices, tokens_per_block, rounding_mode="floor") + page_indices = block_table.index_select(0, block_indices).to(torch.long) + offsets = torch.remainder(local_indices, tokens_per_block) + return page_indices, offsets + + def _gather_paged_rows( + self, + cache: torch.Tensor, + block_table: torch.Tensor, + local_indices: torch.Tensor, + tokens_per_block: int, + ) -> torch.Tensor: + if local_indices.numel() == 0: + return cache.new_empty((0, self.head_dim)) + page_indices, offsets = self._page_locations( + block_table, + local_indices.to(device=block_table.device, dtype=torch.long), + tokens_per_block, + ) + return cache[page_indices, offsets, : self.head_dim] + + def _store_swa_rows( + self, + cache: torch.Tensor, + block_table: torch.Tensor, + positions: torch.Tensor, + rows: torch.Tensor, + tokens_per_block: int, + ) -> None: + if positions.numel() == 0: + return + page_indices, offsets = self._page_locations( + block_table, + positions.to(device=block_table.device, dtype=torch.long), + tokens_per_block, + ) + cache[page_indices, offsets, : self.head_dim] = rows.to(cache.dtype) + + def _validate_inputs( + self, + q: torch.Tensor, + latent_cache: torch.Tensor, + topk_indices: Optional[torch.Tensor], + num_phase_tokens: int, + ) -> None: + expected_q_dim = self.num_heads * self.head_dim + if q.ndim != 2 or q.shape != (num_phase_tokens, expected_q_dim): + raise ValueError( + "DeepSeek-V4 query must have shape " + f"[{num_phase_tokens}, {expected_q_dim}], got {tuple(q.shape)}" + ) + if latent_cache.ndim != 2 or latent_cache.shape != (num_phase_tokens, self.head_dim): + raise ValueError( + "DeepSeek-V4 latent cache must have shape " + f"[{num_phase_tokens}, {self.head_dim}], got " + f"{tuple(latent_cache.shape)}" + ) + + if self.compress_ratio == 4: + if topk_indices is None: + raise ValueError( + "DeepSeek-V4 ratio-4 Vanilla attention requires caller-provided " + "compressed top-k indices" + ) + if topk_indices.ndim != 2 or topk_indices.shape[0] != num_phase_tokens: + raise ValueError( + "DeepSeek-V4 compressed top-k indices must have shape " + f"[{num_phase_tokens}, top_k], got {tuple(topk_indices.shape)}" + ) + if topk_indices.dtype != torch.int32: + raise ValueError( + "DeepSeek-V4 compressed top-k indices must have dtype int32, " + f"got {topk_indices.dtype}" + ) + if torch.any(topk_indices < -1): + raise ValueError("DeepSeek-V4 compressed top-k indices may only use -1 as padding") + + def _forward_sparse( + self, + q: torch.Tensor, + metadata: DeepseekV4TrtllmAttentionMetadata, + forward_args: AttentionForwardArgs, + ) -> torch.Tensor: + seq_start, seq_end = self._phase_sequence_range(metadata, forward_args.attention_input_type) + phase_seq_lens = metadata.seq_lens.tolist()[seq_start:seq_end] + phase_past_tokens = metadata.kv_cache_params.num_cached_tokens_per_seq[seq_start:seq_end] + num_phase_tokens = sum(phase_seq_lens) + + latent_cache = forward_args.latent_cache + if latent_cache is None: + raise ValueError("DeepSeek-V4 Vanilla attention requires latent_cache") + sparse_backend_args = forward_args.sparse_backend_args + topk_indices = sparse_backend_args.topk_indices if sparse_backend_args is not None else None + self._validate_inputs(q, latent_cache, topk_indices, num_phase_tokens) + if self.compress_ratio == 4: + assert topk_indices is not None + available_compressed = torch.tensor( + [ + (int(past) + token_idx + 1) // self.compress_ratio + for q_len, past in zip(phase_seq_lens, phase_past_tokens, strict=True) + for token_idx in range(q_len) + ], + dtype=topk_indices.dtype, + device=topk_indices.device, + ) + if torch.any(topk_indices >= available_compressed.unsqueeze(1)): + raise ValueError( + "DeepSeek-V4 compressed top-k index selects an " + "unavailable or future compressed entry" + ) + + cache_manager = metadata.kv_cache_manager + if cache_manager is None: + raise ValueError("DeepSeek-V4 Vanilla attention requires a KV cache manager") + if self.quant_config is not None and self.quant_config.layer_quant_mode.has_fp8_kv_cache(): + raise NotImplementedError( + "DeepSeek-V4 Vanilla attention does not support an FP8 KV cache" + ) + + swa_cache = cache_manager.get_buffers(self.layer_idx, DeepseekV4AttentionType.SWA) + local_layer_idx = cache_manager.layer_offsets[self.layer_idx] + swa_block_tables = metadata.sliding_block_tables[ + local_layer_idx, DeepseekV4AttentionType.SWA.value + ] + tokens_per_block = cache_manager.tokens_per_block + + compressed_cache = None + compressed_block_tables = None + compressed_tokens_per_block = 0 + if self.compress_ratio > 1: + compressed_cache = cache_manager.get_buffers( + self.layer_idx, DeepseekV4AttentionType.COMPRESS + ) + compressed_block_tables = metadata.compress_block_tables[self.compress_ratio] + compressed_tokens_per_block = cache_manager.compressed_block_sizes[self.layer_idx] + + attention_sink = forward_args.attention_sinks + if attention_sink is None: + attention_sink = getattr(self, "attn_sink", None) + if isinstance(attention_sink, torch.nn.Parameter): + attention_sink = attention_sink.data + + q = q.view(num_phase_tokens, self.num_heads, self.head_dim) + qk_head_dim = self.qk_nope_head_dim + self.qk_rope_head_dim + scale = 1.0 / ( + math.sqrt(qk_head_dim) * (self.q_scaling if self.q_scaling is not None else 1.0) + ) + + outputs = [] + token_offset = 0 + for phase_idx, (q_len, past) in enumerate( + zip(phase_seq_lens, phase_past_tokens, strict=True) + ): + seq_idx = seq_start + phase_idx + past = int(past) + block_table_swa = swa_block_tables[seq_idx] + latent_seq = latent_cache[token_offset : token_offset + q_len].to(q.dtype) + + past_window_start = max(0, past - self.window_size + 1) + past_positions = torch.arange( + past_window_start, + past, + device=block_table_swa.device, + dtype=torch.long, + ) + past_window = self._gather_paged_rows( + swa_cache, + block_table_swa, + past_positions, + tokens_per_block, + ).to(q.dtype) + + per_token_outputs = [] + for token_idx in range(q_len): + current_position = past + token_idx + swa_start = max(0, current_position - self.window_size + 1) + selected_parts = [] + if swa_start < past: + selected_parts.append(past_window[swa_start - past_window_start :]) + current_start = max(0, swa_start - past) + selected_parts.append(latent_seq[current_start : token_idx + 1]) + + if self.compress_ratio > 1: + if self.compress_ratio == 4: + assert topk_indices is not None + compressed_row = topk_indices[token_offset + token_idx] + compressed_indices = compressed_row[compressed_row >= 0].to( + device=block_table_swa.device, dtype=torch.long + ) + else: + num_compressed = (current_position + 1) // self.compress_ratio + compressed_indices = torch.arange( + num_compressed, + device=block_table_swa.device, + dtype=torch.long, + ) + + if compressed_indices.numel() > 0: + assert compressed_cache is not None + assert compressed_block_tables is not None + selected_parts.append( + self._gather_paged_rows( + compressed_cache, + compressed_block_tables[seq_idx], + compressed_indices, + compressed_tokens_per_block, + ).to(q.dtype) + ) + + selected_latent = torch.cat(selected_parts, dim=0) + per_token_outputs.append( + self._selected_mla_attention( + q[token_offset + token_idx], + selected_latent, + value_dim=self.v_head_dim, + scale=scale, + attention_sink=attention_sink, + ) + ) + + outputs.append( + torch.stack(per_token_outputs).reshape(q_len, self.num_heads * self.v_head_dim) + ) + + total_length = past + q_len + first_stored_position = max(past, total_length - self.window_size) + stored_positions = torch.arange( + first_stored_position, + total_length, + device=block_table_swa.device, + dtype=torch.long, + ) + stored_rows = latent_seq[first_stored_position - past :] + self._store_swa_rows( + swa_cache, + block_table_swa, + stored_positions, + stored_rows, + tokens_per_block, + ) + token_offset += q_len + + result = torch.cat(outputs, dim=0) + if forward_args.output is not None: + if forward_args.output.shape != result.shape: + raise ValueError( + "DeepSeek-V4 output buffer must have shape " + f"{tuple(result.shape)}, got {tuple(forward_args.output.shape)}" + ) + forward_args.output.copy_(result) + return forward_args.output + return result + + def forward( + self, + q: torch.Tensor, + k: Optional[torch.Tensor], + v: Optional[torch.Tensor], + metadata: DeepseekV4TrtllmAttentionMetadata, + forward_args: Optional[AttentionForwardArgs] = None, + **kwargs, + ) -> torch.Tensor: + forward_args = merge_attention_forward_args(forward_args, kwargs) + if k is not None or v is not None: + raise ValueError( + "DeepSeek-V4 Vanilla attention expects absorbed queries and " + "latent cache, not explicit K/V tensors" + ) + if metadata.multi_item_part_lens is not None: + raise ValueError("DeepSeek-V4 Vanilla attention does not support multi-item scoring") + return self._forward_sparse(q, metadata, forward_args) diff --git a/tensorrt_llm/_torch/attention_backend/sparse/registry.py b/tensorrt_llm/_torch/attention_backend/sparse/registry.py index 298462b0f409..702690e6cc4c 100644 --- a/tensorrt_llm/_torch/attention_backend/sparse/registry.py +++ b/tensorrt_llm/_torch/attention_backend/sparse/registry.py @@ -70,6 +70,10 @@ def get_vanilla_sparse_attn_attention_backend( if sparse_params.algorithm == "rocket": return RocketVanillaAttention + elif sparse_params.algorithm == "deepseek_v4": + from .deepseek_v4 import DeepseekV4VanillaAttention + + return DeepseekV4VanillaAttention elif sparse_params.algorithm == "minimax_m3": return _resolve_minimax_m3_backend_cls(sparse_params) else: diff --git a/tests/unittest/_torch/attention/backend_capability.py b/tests/unittest/_torch/attention/backend_capability.py index 19d9d7230dee..41472eb50f8c 100644 --- a/tests/unittest/_torch/attention/backend_capability.py +++ b/tests/unittest/_torch/attention/backend_capability.py @@ -20,7 +20,7 @@ # fp4_kv - NVFP4 KV cache (Blackwell only) # sliding_window - sliding-window attention via attention_window_size # no_cache - ragged/prefill forward with kv_cache_manager=None -# sparse - sparse-attention forward plumbing (degenerate regime here) +# sparse - sparse-attention forward plumbing # mla - multi-head latent attention # cross_attn - cross-attention (encoder-decoder) # kv_layouts - supported paged-cache block layouts ("NHD" / "HND") @@ -60,7 +60,7 @@ fp4_kv=False, sliding_window=True, no_cache=True, - sparse=False, + sparse=True, mla=True, cross_attn=True, kv_layouts=("NHD",), # reads the NHD get_buffers view @@ -86,7 +86,7 @@ def required_features(case) -> set: feats.add("sliding_window") if getattr(case, "cache", "paged") == "none": feats.add("no_cache") - if getattr(case, "sparse", "off") != "off": + if getattr(case, "sparse_attention_config", None) is not None: feats.add("sparse") if getattr(case, "is_mla", False): feats.add("mla") @@ -114,6 +114,16 @@ def unsupported_reason(backend: str, case) -> Optional[str]: if not caps.get(feat, False): return f"{backend} does not support feature '{feat}'" + sparse_config = getattr(case, "sparse_attention_config", None) + if sparse_config is not None: + algorithm = sparse_config.algorithm + if backend == "TRTLLM" and algorithm == "deepseek_v4": + # Selected sparse MLA uses trtllm-gen kernels that only ship for + # Blackwell (sm_100+). On Hopper, MLA generation falls back to the + # dense FlashMLA kernel, which has no sparse path. + if sm < 100: + return f"TRTLLM DeepSeek V4 requires sm>=100/Blackwell (have sm{sm})" + # KV-cache block layout: a case may request a specific layout (NHD/HND). A # backend that cannot store the cache that way is skipped (e.g. TRTLLM is # head-major HND only). The Vanilla golden always runs in its native NHD and diff --git a/tests/unittest/_torch/attention/backend_case.py b/tests/unittest/_torch/attention/backend_case.py index cc30f1e051e0..dc895d3a6ec6 100644 --- a/tests/unittest/_torch/attention/backend_case.py +++ b/tests/unittest/_torch/attention/backend_case.py @@ -28,6 +28,9 @@ PredefinedAttentionMask, RopeParams, ) +from tensorrt_llm._torch.attention_backend.sparse import get_sparse_attn_kv_cache_manager +from tensorrt_llm._torch.attention_backend.sparse.deepseek_v4 import DeepseekV4AttentionType +from tensorrt_llm._torch.attention_backend.sparse.params import SparseBackendForwardArgs from tensorrt_llm._torch.attention_backend.utils import create_attention, get_attention_backend from tensorrt_llm._torch.flashinfer_utils import IS_FLASHINFER_AVAILABLE from tensorrt_llm._torch.metadata import KVCacheParams @@ -35,7 +38,7 @@ from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager from tensorrt_llm._utils import str_dtype_to_torch, torch_dtype_to_binding from tensorrt_llm.functional import PositionEmbeddingType, RotaryScalingType -from tensorrt_llm.llmapi.llm_args import KvCacheConfig +from tensorrt_llm.llmapi.llm_args import KvCacheConfig, SparseAttentionConfig from tensorrt_llm.mapping import Mapping from tensorrt_llm.models.modeling_utils import QuantConfig from tensorrt_llm.quantization.mode import QuantAlgo @@ -92,7 +95,11 @@ class BackendCase: q_scaling: float = 1.0 page_size: int = 64 cache: str = "paged" # "paged" | "none" - sparse: str = "off" # "off" | "degenerate" + # User-facing sparse config, lowered into backend params, metadata params, and + # the sparse KV-cache manager exactly as in production. The selection unit and + # top-k are derived from it (properties below); the attention family is + # ``is_mla``. + sparse_attention_config: Optional[SparseAttentionConfig] = None # RoPE config: RopeParams kwargs (+ optional "is_neox"), or None to disable. rope: Optional[dict] = None # When True (and rope set), exercise TRTLLM's in-kernel fused RoPE: TRTLLM @@ -116,6 +123,7 @@ class BackendCase: # and latent-cache inputs: TRTLLM fuses RoPE while Vanilla/FlashInfer receive # the equivalent pre-rotated tensors. v_head_dim: Optional[int] = None + hidden_size: Optional[int] = None q_lora_rank: Optional[int] = None kv_lora_rank: Optional[int] = None qk_nope_head_dim: Optional[int] = None @@ -134,6 +142,28 @@ def nnz_q(self) -> int: def is_cross(self) -> bool: return self.seq_lens_kv is not None + @property + def is_sparse(self) -> bool: + return self.sparse_attention_config is not None + + @property + def sparse_topk(self) -> Optional[int]: + """Per-token selection budget (``index_topk``) from the sparse config.""" + cfg = self.sparse_attention_config + if cfg is None: + return None + return cfg.index_topk + + @property + def prompt_lens(self) -> List[int]: + """Original prompt lengths expected by fused generation kernels.""" + return [ + seq_len if i < self.num_contexts else cached_len + for i, (seq_len, cached_len) in enumerate( + zip(self.seq_lens, self.num_cached_tokens, strict=True) + ) + ] + @property def is_gen_only(self) -> bool: """A uniform pure-decode batch eligible for a captured CUDA graph. @@ -175,6 +205,11 @@ def kv_torch_dtype(self): @property def max_num_tokens(self) -> int: """Metadata token capacity needed by the largest packed tensor in the case.""" + if ( + self.sparse_attention_config is not None + and self.sparse_attention_config.algorithm == "deepseek_v4" + ): + return max(self.nnz_q, self.nnz_kv, *self.token_nums) return max(DEFAULT_MAX_NUM_TOKENS, self.nnz_q, self.nnz_kv, *self.token_nums) def to_dict(self) -> dict: @@ -201,6 +236,12 @@ def _rope_params_from_dict(d: dict) -> RopeParams: return RopeParams(**kwargs) +def _validate_sparse_case(case: BackendCase) -> None: + """Reject unsupported sparse contracts (called only for sparse cases).""" + if case.sparse_topk is None or case.sparse_topk <= 0: + raise ValueError("Sparse backend cases require a positive top-k") + + def _randn(gen: torch.Generator, dtype: torch.dtype, *shape) -> torch.Tensor: """Seeded random tensor on cuda in ``dtype`` (shared by all input builders).""" return torch.randn(*shape, generator=gen, device="cuda").to(dtype) @@ -275,10 +316,16 @@ def _build_kv_cache_manager(case: BackendCase, backend: str, kv_dtype: torch.dty # FlashInfer, the comparison validates the absorbed-MQA math regardless of the # RoPE values (RoPE correctness is covered by test_attention_mla.py). # --------------------------------------------------------------------------- -def _build_mla_kv_cache_manager(case: BackendCase, backend: str): +def _build_mla_kv_cache_manager( + case: BackendCase, + backend: str, + sparse_config=None, +): """A SELFKONLY KV cache for MLA: one latent head, head_dim kv_lora+qk_rope.""" d_latent = case.kv_lora_rank + case.qk_rope_head_dim - paged = BACKEND_CAPS[backend]["paged"] + # Sparse selected-attention tests deliberately exercise multiple pages in + # Vanilla too. Dense Vanilla keeps its historical single-block setup. + paged = case.is_sparse or BACKEND_CAPS[backend]["paged"] max_total = max(case.token_nums) if paged: tokens_per_block = case.page_size @@ -289,10 +336,12 @@ def _build_mla_kv_cache_manager(case: BackendCase, backend: str): num_blocks = case.num_seqs * pages_per_seq mapping = Mapping(world_size=1, tp_size=1, rank=0) cache_types = tensorrt_llm.bindings.internal.batch_manager.CacheType - cls = KVCacheManagerV2 if case.use_kv_cache_manager_v2 else KVCacheManager - return cls( - KvCacheConfig(max_tokens=num_blocks * tokens_per_block, enable_block_reuse=False), - cache_types.SELFKONLY, + kwargs = dict( + kv_cache_config=KvCacheConfig( + max_tokens=num_blocks * tokens_per_block, + enable_block_reuse=False, + ), + kv_cache_type=cache_types.SELFKONLY, num_layers=1, num_kv_heads=1, head_dim=d_latent, @@ -303,6 +352,21 @@ def _build_mla_kv_cache_manager(case: BackendCase, backend: str): dtype=torch_dtype_to_binding(case.compute_dtype), ) + if sparse_config is not None: + cls = get_sparse_attn_kv_cache_manager(sparse_config) + kwargs.update(sparse_attention_config=sparse_config) + if sparse_config.algorithm == "deepseek_v4": + kwargs.update( + compressor_dtype=tensorrt_llm.bindings.DataType.FLOAT, + vocab_size=129280, + max_input_len=max(case.seq_lens), + max_num_tokens=case.max_num_tokens, + ) + else: + cls = KVCacheManagerV2 if case.use_kv_cache_manager_v2 else KVCacheManager + + return cls(**kwargs) + def generate_mla_gen_inputs(case: BackendCase, seed: int = 0) -> Dict: """Random absorbed-MLA generation inputs (shared by all backends). @@ -331,6 +395,127 @@ def generate_mla_gen_inputs(case: BackendCase, seed: int = 0) -> Dict: ) +def _build_deepseek_v4_topk_indices( + case: BackendCase, + generator: torch.Generator, + compress_ratio: int, +) -> torch.Tensor: + """Build request-local indices in DeepSeek-V4's compressed-entry space.""" + topk = case.sparse_topk + if topk is None or topk <= 0: + raise ValueError("DeepSeek-V4 sparse cases require a positive top-k") + + indices = torch.full((case.nnz_q, topk), -1, dtype=torch.int32, device="cuda") + row = 0 + for cached_len, q_len in zip(case.num_cached_tokens, case.seq_lens, strict=True): + for token_idx in range(q_len): + num_compressed = (cached_len + token_idx + 1) // compress_ratio + selected_count = min(topk, num_compressed) + if selected_count == num_compressed: + selected = torch.arange(selected_count, dtype=torch.int32, device="cuda") + else: + selected = torch.randperm( + num_compressed, + generator=generator, + device="cuda", + )[:selected_count] + selected = torch.sort(selected).values.to(torch.int32) + indices[row, :selected_count] = selected + row += 1 + return indices + + +def generate_sparse_mla_inputs(case: BackendCase, seed: int = 0) -> Dict: + """Generate raw absorbed-MLA inputs plus backend-neutral sparse selections.""" + if not case.is_mla: + raise ValueError("This generator supports selected sparse MLA only") + gen = torch.Generator(device="cuda").manual_seed(seed) + selection_gen = torch.Generator(device="cuda").manual_seed(seed + 1) + cdt = case.compute_dtype + num_heads = case.num_heads + kv_lora_rank = case.kv_lora_rank + qk_rope_head_dim = case.qk_rope_head_dim + d_latent = kv_lora_rank + qk_rope_head_dim + + q_nope = _randn(gen, cdt, case.nnz_q, num_heads, kv_lora_rank) + q_pe = _randn(gen, cdt, case.nnz_q, num_heads, qk_rope_head_dim) + compressed_kv = _randn(gen, cdt, case.nnz_q, kv_lora_rank) + k_pe = _randn(gen, cdt, case.nnz_q, qk_rope_head_dim) + + pos_embd_params = _mla_context_pos_embd_params(case) + rope_params = pos_embd_params.rope + assert rope_params is not None + new_positions = make_position_ids(case.seq_lens, case.num_cached_tokens) + # Two input flavors, because RoPE happens in different places per phase: + # * generation: the absorbed path runs with skip_mla_rope_generation, so no + # backend ropes -- feed the RoPE'd (pre-formed) inputs to every backend. + # * context: the TRTLLM MLA context kernel ropes internally (no skip exists + # for context), so it must get RAW inputs and rope them once itself; Vanilla + # (which never ropes) still gets the RoPE'd inputs. + # q_pe rotates per head; k_pe is shared across heads. + rotated_q_pe = apply_rope( + q_pe.reshape(case.nnz_q, num_heads * qk_rope_head_dim), + new_positions, + rope_params, + qk_rope_head_dim, + is_neox=pos_embd_params.is_neox, + ).reshape(case.nnz_q, num_heads, qk_rope_head_dim) + fused_q = torch.cat((q_nope, rotated_q_pe), dim=-1).reshape(case.nnz_q, num_heads * d_latent) + fused_q_raw = torch.cat((q_nope, q_pe), dim=-1).reshape(case.nnz_q, num_heads * d_latent) + rotated_new_k_pe = apply_rope( + k_pe, + new_positions, + rope_params, + qk_rope_head_dim, + is_neox=pos_embd_params.is_neox, + ) + expected_new_latent = torch.cat((compressed_kv, rotated_new_k_pe), dim=-1) + # RoPE'd new-token latent: with skip_mla_rope_generation the backend appends it + # verbatim. The raw variant is roped in-kernel by the TRTLLM context path. + latent_cache = expected_new_latent + latent_cache_raw = torch.cat((compressed_kv, k_pe), dim=-1) + + cached_latent = [] + for cached_len in case.num_cached_tokens: + cached_compressed = _randn(gen, cdt, cached_len, kv_lora_rank) + cached_k_pe = _randn(gen, cdt, cached_len, qk_rope_head_dim) + if cached_len: + cached_positions = torch.arange(cached_len, dtype=torch.int32, device="cuda") + cached_k_pe = apply_rope( + cached_k_pe, + cached_positions, + rope_params, + qk_rope_head_dim, + is_neox=pos_embd_params.is_neox, + ) + cached_latent.append(torch.cat((cached_compressed, cached_k_pe), dim=-1)) + + sparse_config = case.sparse_attention_config + assert sparse_config is not None + sparse_params = sparse_config.to_sparse_params(layer_idx=None, pretrained_config=None) + compress_ratio = sparse_params.compress_ratios[0] + topk_indices = _build_deepseek_v4_topk_indices(case, selection_gen, compress_ratio) + compressed_latent = [ + _randn(gen, cdt, token_count // compress_ratio, d_latent) for token_count in case.token_nums + ] + + return dict( + fused_q=fused_q, + # The RoPE'd q_pe view is passed explicitly since the MLA RoPE step is + # skipped (skip_mla_rope_generation); it must match the fused_q pe slot. + q_pe=fused_q.view(case.nnz_q, num_heads, d_latent)[..., kv_lora_rank:], + latent_cache=latent_cache, + # Raw (un-RoPE'd) variants for the TRTLLM context path, which ropes in-kernel. + fused_q_raw=fused_q_raw, + q_pe_raw=q_pe, + latent_cache_raw=latent_cache_raw, + cached_latent=cached_latent, + compressed_latent=compressed_latent, + expected_new_latent=expected_new_latent, + topk_indices=topk_indices, + ) + + def _fill_mla_cache(mgr, layer_idx, request_ids, cached_latent, *, kv_layout="NHD"): """Write the per-request cached latent prefix into the MLA cache pool.""" if all(c.shape[0] == 0 for c in cached_latent): @@ -358,6 +543,50 @@ def _fill_mla_cache(mgr, layer_idx, request_ids, cached_latent, *, kv_layout="NH written += n +def _fill_deepseek_v4_cache( + mgr, + layer_idx: int, + request_ids: List[int], + cached_latent: List[torch.Tensor], + compressed_latent: List[torch.Tensor], +) -> None: + """Populate DeepSeek-V4's native SWA and compressed cache pools.""" + + def _write_rows( + buffer: torch.Tensor, + block_ids, + rows: torch.Tensor, + tokens_per_block: int, + ) -> None: + for token_idx, row in enumerate(rows): + block_idx = token_idx // tokens_per_block + offset = token_idx % tokens_per_block + buffer[block_ids[block_idx], offset, : row.shape[-1]].copy_(row.to(buffer.dtype)) + + swa_buffer = mgr.get_buffers(layer_idx, DeepseekV4AttentionType.SWA) + compressed_buffer = mgr.get_buffers(layer_idx, DeepseekV4AttentionType.COMPRESS) + compressed_tokens_per_block = mgr.compressed_block_sizes[layer_idx] + for request_id, swa_rows, compressed_rows in zip( + request_ids, + cached_latent, + compressed_latent, + strict=True, + ): + swa_blocks = mgr.get_cache_indices(request_id, layer_idx, DeepseekV4AttentionType.SWA) + _write_rows(swa_buffer, swa_blocks, swa_rows, mgr.tokens_per_block) + compressed_blocks = mgr.get_cache_indices( + request_id, + layer_idx, + DeepseekV4AttentionType.COMPRESS, + ) + _write_rows( + compressed_buffer, + compressed_blocks, + compressed_rows, + compressed_tokens_per_block, + ) + + def _kv_cache_tokens_per_block(buf: torch.Tensor, kv_layout: str) -> int: if kv_layout == "NHD": return buf.shape[2] @@ -455,6 +684,147 @@ def _assert_cache_contains_new_tokens( torch.testing.assert_close(actual, expected, atol=atol, rtol=rtol) +def _run_sparse_mla_backend( + case: BackendCase, + backend: str, + inputs: Dict, + *, + kv_layout: str, +) -> torch.Tensor: + """Run selected sparse MLA through production backend/config lowering.""" + sparse_config = case.sparse_attention_config + assert sparse_config is not None + sparse_params = sparse_config.to_sparse_params(layer_idx=None, pretrained_config=None) + sparse_metadata_params = sparse_config.to_sparse_metadata_params(pretrained_config=None) + AttentionCls = get_attention_backend(backend, sparse_params=sparse_params) + request_ids = list(range(case.num_seqs)) + d_latent = case.kv_lora_rank + case.qk_rope_head_dim + pos_embd_params = _mla_context_pos_embd_params(case) + mapping = Mapping(world_size=1, tp_size=1, rank=0) + attn = create_attention( + backend, + layer_idx=0, + num_heads=case.num_heads, + head_dim=d_latent, + num_kv_heads=1, + q_scaling=case.q_scaling, + pos_embd_params=pos_embd_params, + is_mla_enable=True, + q_lora_rank=case.q_lora_rank, + kv_lora_rank=case.kv_lora_rank, + qk_nope_head_dim=case.qk_nope_head_dim, + qk_rope_head_dim=case.qk_rope_head_dim, + v_head_dim=d_latent, + hidden_size=case.hidden_size, + predicted_tokens_per_seq=1, + sparse_params=sparse_params, + dtype=case.compute_dtype, + skip_create_weights_in_init=True, + ) + # Weights are skipped (selections are injected, not produced by an indexer); + # update_quant_config initializes the quant/FMHA state needed before forward. + attn.update_quant_config(None) + mgr = _build_mla_kv_cache_manager(case, backend, sparse_config) + + try: + mgr.add_dummy_requests(request_ids, case.token_nums) + _fill_deepseek_v4_cache( + mgr, + 0, + request_ids, + inputs["cached_latent"], + inputs["compressed_latent"], + ) + metadata = AttentionCls.Metadata( + num_contexts=case.num_contexts, + kv_cache_params=KVCacheParams( + use_cache=True, + num_cached_tokens_per_seq=case.num_cached_tokens, + ), + seq_lens=torch.tensor(case.seq_lens, dtype=torch.int), + max_num_requests=case.num_seqs, + max_num_tokens=case.max_num_tokens, + kv_cache_manager=mgr, + request_ids=request_ids, + prompt_lens=case.prompt_lens, + kv_layout=kv_layout, + mapping=mapping, + sparse_metadata_params=sparse_metadata_params, + ) + metadata.prepare() + + num_context_tokens = sum(case.seq_lens[: case.num_contexts]) + phases = [] + if case.num_contexts: + phases.append((AttentionInputType.context_only, slice(0, num_context_tokens))) + if case.num_contexts < case.num_seqs: + phases.append( + (AttentionInputType.generation_only, slice(num_context_tokens, case.nnz_q)) + ) + + outputs = [] + for attention_input_type, token_slice in phases: + # RoPE placement differs by phase (see generate_sparse_mla_inputs): + # * TRTLLM consumes raw Q/K and owns RoPE in both phases. + # * Vanilla consumes the equivalent pre-RoPE'd inputs. + # Generation explicitly runs TRTLLM's production RoPE/cache-append + # preparation before the attention forward. + if backend == "TRTLLM": + fused_q = inputs["fused_q_raw"][token_slice].clone() + q_pe = inputs["q_pe_raw"][token_slice].clone() + latent_cache = inputs["latent_cache_raw"][token_slice].clone() + else: + fused_q = inputs["fused_q"][token_slice].clone() + q_pe = inputs["q_pe"][token_slice] + latent_cache = inputs["latent_cache"][token_slice].clone() + topk_indices = inputs["topk_indices"][token_slice] + sparse_backend_args = SparseBackendForwardArgs(topk_indices=topk_indices) + forward_args = AttentionForwardArgs( + latent_cache=latent_cache, + q_pe=q_pe, + sparse_backend_args=sparse_backend_args, + attention_input_type=attention_input_type, + ) + if backend == "TRTLLM" and attention_input_type == AttentionInputType.generation_only: + num_seqs = metadata.num_seqs + cu_q_seqlens = torch.empty(num_seqs + 1, dtype=torch.int32, device=fused_q.device) + cu_kv_seqlens = torch.empty_like(cu_q_seqlens) + fmha_scheduler_counter = torch.empty(1, dtype=torch.uint32, device=fused_q.device) + attn.mla_rope_generation( + fused_q, + q_pe, + latent_cache, + metadata, + cu_q_seqlens, + cu_kv_seqlens, + fmha_scheduler_counter, + None, + None, + None, + ) + forward_args.cu_q_seqlens = cu_q_seqlens + forward_args.cu_kv_seqlens = cu_kv_seqlens + forward_args.fmha_scheduler_counter = fmha_scheduler_counter + out = attn.forward( + fused_q, + None, + None, + metadata, + forward_args=forward_args, + ) + if backend == "TRTLLM": + assert forward_args.sparse_runtime_params.sparse_attn_indices is not None + assert ( + forward_args.sparse_runtime_params.sparse_attn_indices_block_size + == sparse_params.indices_block_size + ) + outputs.append(out[0] if isinstance(out, tuple) else out) + + return torch.cat(outputs, dim=0)[: case.nnz_q].contiguous() + finally: + mgr.shutdown() + + def _run_mla_gen_backend( case, backend, inputs, *, kv_layout: str, cuda_graph=False ) -> torch.Tensor: @@ -755,6 +1125,9 @@ def _tolerances(case: "BackendCase", kv_dtype) -> tuple: dtype is bf16 its coarser mantissa compounds with the quant error, so the quantized atol gets extra headroom and the rtol relaxes to the bf16 rtol. """ + if case.is_sparse: + return 2e-1, 2e-2 + bf16 = case.compute_dtype == torch.bfloat16 if kv_dtype == torch.float8_e4m3fn: return (FP8_ATOL + BF16_ATOL, BF16_RTOL) if bf16 else (FP8_ATOL, RTOL) @@ -875,6 +1248,11 @@ def run_backend( the caller passes a layout the backend supports (gated by the capability matrix). MLA cases are dispatched to the absorbed-generation path. """ + if case.is_sparse: + if case.is_mla: + return _run_sparse_mla_backend(case, backend, inputs, kv_layout=kv_layout) + raise ValueError(f"Unsupported sparse contract: is_mla={case.is_mla}") + if case.is_mla: if case.is_context_only: return _run_mla_context_backend(case, backend, inputs, kv_layout=kv_layout) @@ -1064,7 +1442,13 @@ def run_case(case: BackendCase, *, seed: int = 0) -> Dict[str, torch.Tensor]: that want the raw tensors (e.g. the minimizer). """ is_mla = case.is_mla - if is_mla: + if case.is_sparse: + _validate_sparse_case(case) + if case.is_mla: + inputs = generate_sparse_mla_inputs(case, seed) + else: + raise ValueError(f"Unsupported sparse contract: is_mla={case.is_mla}") + elif is_mla: if case.is_context_only: inputs = generate_mla_context_inputs(case, seed) else: @@ -1097,8 +1481,9 @@ def run_case(case: BackendCase, *, seed: int = 0) -> Dict[str, torch.Tensor]: # A gen-only batch also exercises the captured-CUDA-graph path # (production replays a captured decode graph); it must still match the - # eager golden. - if case.is_gen_only: + # eager golden. The sparse runner rebuilds each request's logical cache on + # the host, which is not graph-capturable, so sparse cases are skipped. + if case.is_gen_only and not case.is_sparse: cg_out = run_backend( case, backend, diff --git a/tests/unittest/_torch/attention/model_attn_config.py b/tests/unittest/_torch/attention/model_attn_config.py index 82f52a460c88..8fe3f0d75da8 100644 --- a/tests/unittest/_torch/attention/model_attn_config.py +++ b/tests/unittest/_torch/attention/model_attn_config.py @@ -19,7 +19,9 @@ ``rope`` is one of ``None`` / ``"neox"`` / ``"gptj"``. ``mask`` is ``"causal"`` / ``"full"`` / ``"sliding"``. ``no_cache=True`` marks the -bidirectional, KV-cache-free DiT / encoder workloads. +bidirectional, KV-cache-free DiT / encoder workloads. Sparse models carry the +same user-facing sparse-attention config that production lowers independently +for the backend, metadata, and KV-cache manager. ID naming rule: - Use lowercase snake_case: @@ -46,14 +48,17 @@ Qwen2.5-VL with ``use_sliding_window=False``). - Vision/text encoders (SigLip/Radio/CLIP/Parakeet, MiniMax-VL tower) collapse onto the bidirectional MHA tuples already listed (e.g. 16x64 / 12x64 full). -- Sparse/DSA indexer attention (GLM-DSA, NSA, RocketKV) is a separate paradigm - validated under sparse/; there is no dense Vanilla golden for it. +- Sparse indexer correctness is validated under ``sparse/``. This model sweep + injects deterministic backend-neutral selections and validates the selected + attention path against the Vanilla golden. - Multimodal cross variants (Llama4-vision, Gemma4-MM) reuse the cross tuples. """ from dataclasses import dataclass from typing import Literal, Optional +from tensorrt_llm.llmapi.llm_args import DeepSeekV4SparseAttentionConfig, SparseAttentionConfig + AttentionPhase = Literal["ctx", "gen"] @@ -80,6 +85,18 @@ class ModelAttnConfig: qk_nope_head_dim: Optional[int] = None qk_rope_head_dim: Optional[int] = None v_head_dim: Optional[int] = None + hidden_size: Optional[int] = None + # User-facing sparse config, lowered by production `to_sparse_params()`. The + # sparse sweep derives its other parameters from this and `is_mla`. + sparse_attention_config: Optional[SparseAttentionConfig] = None + + @property + def sparse_topk(self) -> Optional[int]: + """Per-token selection budget (``index_topk``) from the sparse config.""" + cfg = self.sparse_attention_config + if cfg is None: + return None + return cfg.index_topk # --------------------------------------------------------------------------- @@ -555,6 +572,33 @@ class ModelAttnConfig: # MLA (DeepSeek-style absorbed latent attention). num_kv_heads == 1 latent head. # --------------------------------------------------------------------------- _MLA = [ + # DeepSeek-V4 ratio-4 layers combine a 128-token sliding window with + # compressed-entry selections. The unified sweep uses a reduced top-k to + # keep the backend parity case bounded; the native DeepSeek-V4 tests cover + # the production top-k and the ratio-1/128 layer variants. + ModelAttnConfig( + "deepseekv4_sparse_mla", + "DeepSeek-V4 ratio-4 layers", + num_heads=64, + num_kv_heads=1, + head_dim=512, + rope="gptj", + is_mla=True, + phases=("ctx", "gen"), + kv_lora_rank=448, + q_lora_rank=1536, + qk_nope_head_dim=128, + qk_rope_head_dim=64, + v_head_dim=512, + hidden_size=7168, + sparse_attention_config=DeepSeekV4SparseAttentionConfig( + index_n_heads=64, + index_head_dim=128, + index_topk=8, + compress_ratios=[4], + skip_indexer_for_short_seqs=False, + ), + ), # DeepSeek-V3: 128 Q heads, qk_nope=128, qk_rope=64, kv_lora=512, v=128. ModelAttnConfig( "deepseekv3_mla", diff --git a/tests/unittest/_torch/attention/sparse/deepseek_v4/test_deepseek_v4_sparse_mla.py b/tests/unittest/_torch/attention/sparse/deepseek_v4/test_deepseek_v4_sparse_mla.py index 77e1e9dcd5ff..dc168e094f72 100644 --- a/tests/unittest/_torch/attention/sparse/deepseek_v4/test_deepseek_v4_sparse_mla.py +++ b/tests/unittest/_torch/attention/sparse/deepseek_v4/test_deepseek_v4_sparse_mla.py @@ -38,6 +38,7 @@ DeepseekV4CacheManager, DeepseekV4TrtllmAttention, DeepseekV4TrtllmAttentionMetadata, + DeepseekV4VanillaAttention, ) from tensorrt_llm._torch.attention_backend.sparse.deepseek_v4.cache_manager import get_token_bytes from tensorrt_llm._torch.attention_backend.sparse.params import SparseBackendForwardArgs @@ -320,84 +321,6 @@ def _softmax_with_sink( return (num / denom).to(out_dtype) -def calculate_deepseek_v4_ref_ctx_sparse( - fused_q_rot: torch.Tensor, - latent_cache_ref: torch.Tensor, - compressed_ref_data: Optional[List[torch.Tensor]], - swa_window_size: int, - compressed_topk_indices: Optional[torch.Tensor], - seq_lens: List[int], - num_heads: int, - kv_lora_rank: int, - v_head_dim: int, - qk_nope_head_dim: int, - qk_rope_head_dim: int, - q_scaling: float, - compress_ratio: int, - attn_sink: Optional[torch.Tensor] = None, -): - """Per-token reference attention for DeepSeek-V4 context phase. - - For compress_ratio==1: only SWA tokens (causal window). - For compress_ratio==4: SWA tokens + indexer topk compressed tokens. - For compress_ratio==128: SWA tokens + all compressed tokens. - """ - fused_head_dim = kv_lora_rank + qk_rope_head_dim - bmm1_scale = 1 / (math.sqrt(qk_nope_head_dim + qk_rope_head_dim) * q_scaling) - - ref_results = [] - token_offset = 0 - for batch_idx, seq_len in enumerate(seq_lens): - per_token_outputs = [] - for token_idx in range(seq_len): - global_token_idx = token_offset + token_idx - pos = token_idx - - # Gather SWA KV - swa_start = max(0, pos - swa_window_size + 1) - swa_end = pos + 1 - swa_kv = latent_cache_ref[token_offset + swa_start : token_offset + swa_end] - - # Gather compressed KV - if compress_ratio > 1 and compressed_ref_data is not None: - if compress_ratio == 4 and compressed_topk_indices is not None: - indices_row = compressed_topk_indices[global_token_idx] - valid = indices_row[indices_row >= 0] - comp_kv = compressed_ref_data[batch_idx][valid.long()] - elif compress_ratio == 128: - num_comp = (pos + 1) // compress_ratio - comp_kv = compressed_ref_data[batch_idx][:num_comp] - else: - comp_kv = torch.empty( - 0, - swa_kv.shape[-1], - device=swa_kv.device, - dtype=swa_kv.dtype, - ) - - if comp_kv.numel() > 0: - all_kv = torch.cat([swa_kv, comp_kv], dim=0) - else: - all_kv = swa_kv - else: - all_kv = swa_kv - - # Compute attention - q_tok = fused_q_rot[global_token_idx].view(num_heads, fused_head_dim) - k_sel = all_kv.unsqueeze(0).expand(num_heads, -1, -1) - v_sel = all_kv[:, :v_head_dim].unsqueeze(0).expand(num_heads, -1, -1) - - attn_w = torch.matmul(q_tok.unsqueeze(1), k_sel.transpose(1, 2)) * bmm1_scale - attn_w = _softmax_with_sink(attn_w, attn_sink, fused_q_rot.dtype) - out = torch.matmul(attn_w, v_sel).squeeze(1) - per_token_outputs.append(out.reshape(1, num_heads * v_head_dim)) - - ref_results.append(torch.cat(per_token_outputs, dim=0)) - token_offset += seq_len - - return torch.cat(ref_results, dim=0) - - def _rotate_gen_inputs( fused_q: torch.Tensor, q_pe: torch.Tensor, @@ -604,6 +527,31 @@ def _create_pos_embd_params(scenario: Scenario) -> PositionalEmbeddingParams: ) +def _create_vanilla_layers( + layer_indices: List[int], + num_heads: int, + head_dim: int, + q_scaling: float, + pos_embd_params: PositionalEmbeddingParams, + mla_params: MLAParams, + sparse_config: DeepSeekV4SparseAttentionConfig, +) -> Dict[int, DeepseekV4VanillaAttention]: + return { + layer_idx: DeepseekV4VanillaAttention( + layer_idx=layer_idx, + num_heads=num_heads, + head_dim=head_dim, + num_kv_heads=1, + q_scaling=q_scaling, + pos_embd_params=pos_embd_params, + mla_params=mla_params, + sparse_attention_config=sparse_config, + skip_create_weights_in_init=True, + ) + for layer_idx in layer_indices + } + + @skip_pre_blackwell def test_deepseek_v4_sparse_mla_single_token_tp4_local_heads_repro(): """Reproduce the tp=4 Flash warmup sparse-MLA kernel shape. @@ -669,6 +617,15 @@ def test_deepseek_v4_sparse_mla_single_token_tp4_local_heads_repro(): skip_create_weights_in_init=True, ) layer.update_quant_config(None) + vanilla_layer = _create_vanilla_layers( + [layer_idx], + local_num_heads, + head_dim, + q_scaling, + pos_embd_params, + mla_params, + sparse_config, + )[layer_idx] attn_sink = torch.randn(local_num_heads, dtype=torch.float32, device=device).mul_(0.5) if not os.environ.get("DSV4_REPRO_NO_SINK"): layer.attn_sink = torch.nn.Parameter(attn_sink, requires_grad=False) @@ -794,27 +751,222 @@ def test_deepseek_v4_sparse_mla_single_token_tp4_local_heads_repro(): kv_lora_rank, qk_rope_head_dim, ) - ref_result = calculate_deepseek_v4_ref_ctx_sparse( + golden_sink = None if os.environ.get("DSV4_REPRO_NO_SINK") else attn_sink + ref_result = vanilla_layer.forward( fused_q_rot, - latent_cache_ref, None, - scenario.window_size, None, - context_lengths, - local_num_heads, - kv_lora_rank, - v_head_dim, - qk_nope_head_dim, - qk_rope_head_dim, - q_scaling, - ratio, - attn_sink=attn_sink, + attn_metadata, + attention_input_type=AttentionInputType.context_only, + latent_cache=latent_cache_ref, + attention_sinks=golden_sink, ) torch.testing.assert_close(result, ref_result, atol=0.2, rtol=0.02) cache_manager.shutdown() +@skip_pre_blackwell +def test_deepseek_v4_sparse_mla_vanilla_golden(): + """Validate TRTLLM sparse MLA against the Vanilla golden backend.""" + scenario = Scenario() + context_lengths = [14, 140] + device = torch.device("cuda") + dtype = scenario.dtype + num_heads = 16 + qk_rope_head_dim = scenario.qk_rope_head_dim + kv_lora_rank = scenario.kv_lora_rank - qk_rope_head_dim + head_dim = kv_lora_rank + qk_rope_head_dim + total_tokens = sum(context_lengths) + request_ids = list(range(len(context_lengths))) + torch.manual_seed(43) + + # The TRTLLM sparse kernel requires its padded top-k width to be a multiple + # of four. Ratio-128 metadata rounds the compressed width to a power of two, + # so use a max sequence length that gives at least four compressed slots. + max_seq_len = scenario.window_size * 3 + cache_manager, sparse_config = _create_cache_manager( + scenario, context_lengths, max_seq_len=max_seq_len + ) + requests = [ + LlmRequest( + request_id=request_id, + max_new_tokens=1, + input_tokens=list(range(context_length)), + sampling_config=SamplingConfig(), + is_streaming=False, + ) + for request_id, context_length in zip(request_ids, context_lengths, strict=True) + ] + for request in requests: + cache_manager.prepare_context(request) + cache_manager.resize_context(request, request.context_chunk_size) + + mla_params = MLAParams( + q_lora_rank=scenario.q_lora_rank, + kv_lora_rank=kv_lora_rank, + qk_rope_head_dim=qk_rope_head_dim, + qk_nope_head_dim=scenario.qk_nope_head_dim, + v_head_dim=scenario.v_head_dim, + rope_append=False, + predicted_tokens_per_seq=1, + hidden_size=scenario.hidden_size, + ) + pos_embd_params = _create_pos_embd_params(scenario) + mscale = 0.1 * pos_embd_params.rope.mscale_all_dim * math.log(pos_embd_params.rope.scale) + 1.0 + q_scaling = 1.0 / (mscale * mscale) + + trtllm_layers = { + layer_idx: DeepseekV4TrtllmAttention( + layer_idx=layer_idx, + num_heads=num_heads, + head_dim=head_dim, + num_kv_heads=1, + q_scaling=q_scaling, + pos_embd_params=pos_embd_params, + mla_params=mla_params, + sparse_attention_config=sparse_config, + skip_create_weights_in_init=True, + ) + for layer_idx in TEST_LAYERS + } + for layer in trtllm_layers.values(): + layer.update_quant_config(None) + + vanilla_layers = _create_vanilla_layers( + TEST_LAYERS, + num_heads, + head_dim, + q_scaling, + pos_embd_params, + mla_params, + sparse_config, + ) + + for layer_idx in TEST_LAYERS: + if scenario.compress_ratios[layer_idx] <= 1: + continue + _prefill_compress_buffer( + cache_manager, + layer_idx, + context_lengths, + request_ids, + head_dim, + device, + ) + + metadata = DeepseekV4TrtllmAttentionMetadata( + seq_lens=torch.tensor(context_lengths, dtype=torch.int), + request_ids=request_ids, + max_num_requests=len(request_ids), + num_contexts=len(request_ids), + prompt_lens=context_lengths, + max_num_tokens=total_tokens, + kv_cache_manager=cache_manager, + kv_cache_params=KVCacheParams( + use_cache=True, + num_cached_tokens_per_seq=[0] * len(request_ids), + ), + mapping=Mapping(world_size=1, tp_size=1, rank=0), + sparse_attention_config=sparse_config, + ) + + try: + metadata.prepare() + rope_cos_sin = _create_rope_cos_sin(scenario, device) + token_positions = [ + position for context_length in context_lengths for position in range(context_length) + ] + for layer_idx in TEST_LAYERS: + ratio = scenario.compress_ratios[layer_idx] + fused_q = torch.randn( + total_tokens, + num_heads * head_dim, + dtype=dtype, + device=device, + ) + q_pe = fused_q.view(total_tokens, num_heads, head_dim)[..., -qk_rope_head_dim:].clone() + compressed_kv = torch.randn( + total_tokens, + kv_lora_rank, + dtype=dtype, + device=device, + ) + k_pe = torch.randn( + total_tokens, + qk_rope_head_dim, + dtype=dtype, + device=device, + ) + latent_cache = torch.cat([compressed_kv, k_pe], dim=-1) + attention_sink = torch.randn( + num_heads, + dtype=torch.float32, + device=device, + ) + topk_indices = ( + _build_compressed_topk_indices( + token_positions, + ratio, + scenario.index_topk, + device, + ) + if ratio == 4 + else None + ) + trtllm_output = torch.empty( + total_tokens, + num_heads * head_dim, + dtype=dtype, + device=device, + ) + + result = trtllm_layers[layer_idx].forward( + fused_q.clone(), + None, + None, + metadata, + forward_args=AttentionForwardArgs( + output=trtllm_output, + latent_cache=latent_cache.clone(), + q_pe=q_pe, + attention_sinks=attention_sink, + attention_input_type=AttentionInputType.context_only, + sparse_backend_args=SparseBackendForwardArgs(topk_indices=topk_indices), + ), + ) + fused_q_rot = _rotate_fused_q_for_ctx( + fused_q, + rope_cos_sin, + context_lengths, + num_heads, + kv_lora_rank, + qk_rope_head_dim, + ) + k_pe_rot = _rotate_k_pe_for_ctx(k_pe, rope_cos_sin, context_lengths) + latent_cache_rot = torch.cat([compressed_kv, k_pe_rot], dim=-1) + vanilla_output = torch.empty_like(trtllm_output) + golden = vanilla_layers[layer_idx].forward( + fused_q_rot, + None, + None, + metadata, + forward_args=AttentionForwardArgs( + output=vanilla_output, + latent_cache=latent_cache_rot, + attention_sinks=attention_sink, + attention_input_type=AttentionInputType.context_only, + sparse_backend_args=SparseBackendForwardArgs(topk_indices=topk_indices), + ), + ) + + assert result.data_ptr() == trtllm_output.data_ptr() + assert golden.data_ptr() == vanilla_output.data_ptr() + torch.testing.assert_close(result, golden, atol=0.2, rtol=2e-2) + finally: + cache_manager.shutdown() + + @skip_pre_blackwell @pytest.mark.skip_less_device_memory(80000) @pytest.mark.parametrize("context_lengths", [[4399], [14, 508, 3947], [2, 1406, 3327]]) @@ -908,6 +1060,15 @@ def yarn_get_mscale(scale=1, mscale=1): ) layer.update_quant_config(None) layers[layer_idx] = layer + vanilla_layers = _create_vanilla_layers( + TEST_LAYERS, + num_heads, + head_dim, + q_scaling, + pos_embd_params, + mla_params, + sparse_config, + ) # Install a per-layer attention sink attn_sinks: Dict[int, torch.Tensor] = {} @@ -1088,7 +1249,7 @@ def yarn_get_mscale(scale=1, mscale=1): sparse_backend_args=SparseBackendForwardArgs(topk_indices=topk_indices), ) - # Reference computation + # Vanilla golden computation k_pe_ref = _rotate_k_pe_for_ctx(k_pe, rope_cos_sin, context_lengths) latent_cache_ref = torch.cat([compressed_kv, k_pe_ref], dim=-1) fused_q_rot = _rotate_fused_q_for_ctx( @@ -1099,21 +1260,15 @@ def yarn_get_mscale(scale=1, mscale=1): kv_lora_rank, qk_rope_head_dim, ) - ref_result = calculate_deepseek_v4_ref_ctx_sparse( + ref_result = vanilla_layers[layer_idx].forward( fused_q_rot, - latent_cache_ref, - compress_ref_data.get(layer_idx), - scenario.window_size, - topk_indices, - context_lengths, - num_heads, - kv_lora_rank, - v_head_dim, - qk_nope_head_dim, - qk_rope_head_dim, - q_scaling, - ratio, - attn_sink=attn_sinks[layer_idx], + None, + None, + attn_metadata, + attention_input_type=AttentionInputType.context_only, + latent_cache=latent_cache_ref, + attention_sinks=attn_sinks[layer_idx], + sparse_backend_args=SparseBackendForwardArgs(topk_indices=topk_indices), ) latent_cache_ref_all[layer_idx] = latent_cache_ref @@ -1373,6 +1528,15 @@ def test_deepseek_v4_sparse_mla_mixed_batch(context_lengths: List[int]): ) layer.update_quant_config(None) layers[li] = layer + vanilla_layers = _create_vanilla_layers( + TEST_LAYERS, + num_heads, + head_dim, + q_scaling, + pos_embd_params, + mla_params, + sparse_config, + ) attn_sinks: Dict[int, torch.Tensor] = {} for li in TEST_LAYERS: @@ -1558,6 +1722,27 @@ def test_deepseek_v4_sparse_mla_mixed_batch(context_lengths: List[int]): sparse_backend_args=SparseBackendForwardArgs(topk_indices=ctx_topk), ) + ctx_k_pe_ref = _rotate_k_pe_for_ctx(ctx_k_pe, rope_cos_sin, [context_lengths[0]]) + ctx_latent_ref = torch.cat([ctx_compressed_kv, ctx_k_pe_ref], dim=-1) + ctx_q_rot = _rotate_fused_q_for_ctx( + ctx_fused_q, + rope_cos_sin, + [context_lengths[0]], + num_heads, + kv_lora_rank, + qk_rope_head_dim, + ) + ctx_ref = vanilla_layers[li].forward( + ctx_q_rot, + None, + None, + mixed_metadata, + attention_input_type=AttentionInputType.context_only, + latent_cache=ctx_latent_ref, + attention_sinks=attn_sinks[li], + sparse_backend_args=SparseBackendForwardArgs(topk_indices=ctx_topk), + ) + # Generation forward → output[total_ctx_tokens:] num_seqs = mixed_metadata.kv_lens_cuda_runtime.size(0) cu_q = torch.empty(num_seqs + 1, dtype=torch.int32, device=device) @@ -1590,34 +1775,6 @@ def test_deepseek_v4_sparse_mla_mixed_batch(context_lengths: List[int]): sparse_backend_args=SparseBackendForwardArgs(topk_indices=gen_topk), ) - # Context reference - ctx_k_pe_ref = _rotate_k_pe_for_ctx(ctx_k_pe, rope_cos_sin, [context_lengths[0]]) - ctx_latent_ref = torch.cat([ctx_compressed_kv, ctx_k_pe_ref], dim=-1) - ctx_q_rot = _rotate_fused_q_for_ctx( - ctx_fused_q, - rope_cos_sin, - [context_lengths[0]], - num_heads, - kv_lora_rank, - qk_rope_head_dim, - ) - ctx_ref = calculate_deepseek_v4_ref_ctx_sparse( - ctx_q_rot, - ctx_latent_ref, - None, - scenario.window_size, - ctx_topk, - [context_lengths[0]], - num_heads, - kv_lora_rank, - v_head_dim, - qk_nope_head_dim, - qk_rope_head_dim, - q_scaling, - ratio, - attn_sink=attn_sinks[li], - ) - # Generation reference prefill_k_pe_ref = _rotate_k_pe_for_ctx(prefill_k_pe, rope_cos_sin, gen_ctx_lengths) gen_latent_ref = torch.cat([prefill_compressed_kv, prefill_k_pe_ref], dim=-1) diff --git a/tests/unittest/_torch/attention/sparse/deepseek_v4/test_deepseek_v4_vanilla.py b/tests/unittest/_torch/attention/sparse/deepseek_v4/test_deepseek_v4_vanilla.py new file mode 100644 index 000000000000..ece27df48328 --- /dev/null +++ b/tests/unittest/_torch/attention/sparse/deepseek_v4/test_deepseek_v4_vanilla.py @@ -0,0 +1,309 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Tests for the DeepSeek-V4 Vanilla selected-attention golden.""" + +import math +from types import SimpleNamespace + +import pytest +import torch + +from tensorrt_llm._torch.attention_backend.interface import AttentionForwardArgs, AttentionInputType +from tensorrt_llm._torch.attention_backend.sparse.deepseek_v4 import ( + DeepseekV4AttentionType, + DeepSeekV4Params, + DeepseekV4VanillaAttention, +) +from tensorrt_llm._torch.attention_backend.sparse.params import SparseBackendForwardArgs +from tensorrt_llm._torch.attention_backend.sparse.registry import ( + get_vanilla_sparse_attn_attention_backend, +) + + +class _FakeDeepseekV4CacheManager: + def __init__( + self, + swa_cache: torch.Tensor, + compressed_cache: torch.Tensor | None, + tokens_per_block: int, + compressed_tokens_per_block: int, + ) -> None: + self.swa_cache = swa_cache + self.compressed_cache = compressed_cache + self.tokens_per_block = tokens_per_block + self.compressed_block_sizes = [compressed_tokens_per_block] + self.layer_offsets = [0] + + def get_buffers( + self, + layer_idx: int, + attention_type: DeepseekV4AttentionType, + ) -> torch.Tensor: + assert layer_idx == 0 + if attention_type == DeepseekV4AttentionType.SWA: + return self.swa_cache + if attention_type == DeepseekV4AttentionType.COMPRESS: + assert self.compressed_cache is not None + return self.compressed_cache + raise ValueError(f"Unexpected attention type: {attention_type}") + + +def _create_backend(compress_ratio: int, window_size: int) -> DeepseekV4VanillaAttention: + backend = object.__new__(DeepseekV4VanillaAttention) + backend.layer_idx = 0 + backend.num_heads = 2 + backend.head_dim = 4 + backend.num_kv_heads = 1 + backend.quant_config = None + backend.q_scaling = 1.0 + backend.qk_nope_head_dim = 2 + backend.qk_rope_head_dim = 2 + backend.v_head_dim = 4 + backend.compress_ratio = compress_ratio + backend.window_size = window_size + return backend + + +def _write_paged_rows( + cache: torch.Tensor, + block_table: torch.Tensor, + rows: torch.Tensor, + tokens_per_block: int, +) -> None: + positions = torch.arange(rows.shape[0], dtype=torch.long) + pages = block_table[torch.div(positions, tokens_per_block, rounding_mode="floor")] + offsets = torch.remainder(positions, tokens_per_block) + cache[pages, offsets] = rows + + +def _reference_attention( + q: torch.Tensor, + full_latent: torch.Tensor, + compressed: torch.Tensor | None, + topk_indices: torch.Tensor | None, + *, + past: int, + compress_ratio: int, + window_size: int, + attention_sink: torch.Tensor, +) -> torch.Tensor: + outputs = [] + scale = 1.0 / math.sqrt(4) + for token_idx in range(q.shape[0]): + position = past + token_idx + swa_start = max(0, position - window_size + 1) + selected = [full_latent[swa_start : position + 1]] + if compress_ratio == 4: + assert compressed is not None and topk_indices is not None + row = topk_indices[token_idx] + valid = row[row >= 0].long() + if valid.numel() > 0: + selected.append(compressed.index_select(0, valid)) + elif compress_ratio == 128: + assert compressed is not None + num_compressed = (position + 1) // compress_ratio + if num_compressed > 0: + selected.append(compressed[:num_compressed]) + + selected_latent = torch.cat(selected) + scores = q[token_idx] @ selected_latent.transpose(0, 1) * scale + scores = scores.float() + sink = attention_sink.float().unsqueeze(1) + max_score = torch.maximum(scores.amax(dim=-1, keepdim=True), sink) + numerator = torch.exp(scores - max_score) + denominator = numerator.sum(dim=-1, keepdim=True) + denominator += torch.exp(sink - max_score) + probabilities = (numerator / denominator).to(q.dtype) + outputs.append(probabilities @ selected_latent) + return torch.stack(outputs).flatten(1) + + +@pytest.mark.parametrize( + ("compress_ratio", "attention_input_type", "past", "q_len"), + [ + (1, AttentionInputType.context_only, 0, 7), + (4, AttentionInputType.context_only, 0, 9), + (4, AttentionInputType.generation_only, 8, 2), + (128, AttentionInputType.generation_only, 128, 2), + ], +) +def test_deepseek_v4_vanilla_selected_attention( + compress_ratio: int, + attention_input_type: AttentionInputType, + past: int, + q_len: int, +) -> None: + torch.manual_seed(41 + compress_ratio + past) + window_size = 4 + tokens_per_block = 4 + compressed_tokens_per_block = 2 + total_length = past + q_len + num_swa_pages = math.ceil(total_length / tokens_per_block) + swa_cache = torch.zeros(num_swa_pages, tokens_per_block, 4) + swa_block_table = torch.arange(num_swa_pages, dtype=torch.int32) + + past_latent = torch.randn(past, 4) + new_latent = torch.randn(q_len, 4) + _write_paged_rows(swa_cache, swa_block_table, past_latent, tokens_per_block) + + num_compressed = total_length // compress_ratio if compress_ratio > 1 else 0 + compressed = torch.randn(num_compressed, 4) if num_compressed > 0 else None + compressed_cache = None + compressed_block_table = None + if compress_ratio > 1: + num_compressed_pages = max(1, math.ceil(num_compressed / compressed_tokens_per_block)) + compressed_cache = torch.zeros(num_compressed_pages, compressed_tokens_per_block, 4) + compressed_block_table = torch.arange(num_compressed_pages, dtype=torch.int32) + if compressed is not None: + _write_paged_rows( + compressed_cache, + compressed_block_table, + compressed, + compressed_tokens_per_block, + ) + + cache_manager = _FakeDeepseekV4CacheManager( + swa_cache, + compressed_cache, + tokens_per_block, + compressed_tokens_per_block, + ) + sliding_block_tables = torch.full( + (1, len(DeepseekV4AttentionType), 1, num_swa_pages), + -1, + dtype=torch.int32, + ) + sliding_block_tables[0, DeepseekV4AttentionType.SWA.value, 0] = swa_block_table + compress_block_tables = ( + {compress_ratio: compressed_block_table.unsqueeze(0)} + if compressed_block_table is not None + else {} + ) + metadata = SimpleNamespace( + seq_lens=torch.tensor([q_len], dtype=torch.int32), + num_contexts=1 if attention_input_type == AttentionInputType.context_only else 0, + num_seqs=1, + kv_cache_params=SimpleNamespace(num_cached_tokens_per_seq=[past]), + kv_cache_manager=cache_manager, + sliding_block_tables=sliding_block_tables, + compress_block_tables=compress_block_tables, + multi_item_part_lens=None, + ) + + topk_indices = None + if compress_ratio == 4: + topk_indices = torch.full((q_len, 3), -1, dtype=torch.int32) + for token_idx in range(q_len): + available = (past + token_idx + 1) // compress_ratio + if available > 0: + topk_indices[token_idx, 0] = 0 + topk_indices[token_idx, 1] = available - 1 + + q = torch.randn(q_len, 2, 4) + attention_sink = torch.tensor([0.25, -0.5]) + output = torch.empty(q_len, 8) + backend = _create_backend(compress_ratio, window_size) + result = backend.forward( + q.flatten(1), + None, + None, + metadata, + forward_args=AttentionForwardArgs( + output=output, + latent_cache=new_latent, + attention_sinks=attention_sink, + attention_input_type=attention_input_type, + sparse_backend_args=SparseBackendForwardArgs(topk_indices=topk_indices), + ), + ) + + expected = _reference_attention( + q, + torch.cat([past_latent, new_latent]), + compressed, + topk_indices, + past=past, + compress_ratio=compress_ratio, + window_size=window_size, + attention_sink=attention_sink, + ) + assert result.data_ptr() == output.data_ptr() + torch.testing.assert_close(result, expected) + + first_stored = max(past, total_length - window_size) + stored_positions = torch.arange(first_stored, total_length, dtype=torch.long) + stored_pages = swa_block_table[ + torch.div(stored_positions, tokens_per_block, rounding_mode="floor") + ] + stored_offsets = torch.remainder(stored_positions, tokens_per_block) + actual_stored = swa_cache[stored_pages, stored_offsets] + torch.testing.assert_close(actual_stored, new_latent[first_stored - past :]) + + +def test_deepseek_v4_vanilla_rejects_future_compressed_index() -> None: + backend = _create_backend(compress_ratio=4, window_size=4) + swa_cache = torch.zeros(1, 4, 4) + compressed_cache = torch.zeros(1, 2, 4) + cache_manager = _FakeDeepseekV4CacheManager(swa_cache, compressed_cache, 4, 2) + sliding_block_tables = torch.full( + (1, len(DeepseekV4AttentionType), 1, 1), + -1, + dtype=torch.int32, + ) + sliding_block_tables[0, DeepseekV4AttentionType.SWA.value, 0, 0] = 0 + metadata = SimpleNamespace( + seq_lens=torch.tensor([1], dtype=torch.int32), + num_contexts=1, + num_seqs=1, + kv_cache_params=SimpleNamespace(num_cached_tokens_per_seq=[0]), + kv_cache_manager=cache_manager, + sliding_block_tables=sliding_block_tables, + compress_block_tables={4: torch.zeros(1, 1, dtype=torch.int32)}, + multi_item_part_lens=None, + ) + + with pytest.raises(ValueError, match="future compressed entry"): + backend.forward( + torch.randn(1, 8), + None, + None, + metadata, + forward_args=AttentionForwardArgs( + latent_cache=torch.randn(1, 4), + attention_input_type=AttentionInputType.context_only, + sparse_backend_args=SparseBackendForwardArgs( + topk_indices=torch.tensor([[0]], dtype=torch.int32) + ), + ), + ) + + +def test_deepseek_v4_vanilla_backend_registry() -> None: + sparse_params = DeepSeekV4Params(compress_ratios=[1]) + assert get_vanilla_sparse_attn_attention_backend(sparse_params) is DeepseekV4VanillaAttention + + +def test_deepseek_v4_vanilla_mla_rope_generation_accepts_full_contract() -> None: + backend = _create_backend(compress_ratio=1, window_size=4) + tensor = torch.empty(0) + + backend.mla_rope_generation( + None, + None, + tensor, + SimpleNamespace(), + tensor, + tensor, + tensor, + None, + None, + None, + out_scale=tensor, + kv_norm_weight=tensor, + kv_norm_eps=1e-5, + precomputed_cu_seqlens=True, + precomputed_fmha_scheduler=True, + kv_only=True, + kv_done_elsewhere=False, + quant_scale_qkv=tensor, + ) diff --git a/tests/unittest/_torch/attention/sparse/test_sparse_mla_forward.py b/tests/unittest/_torch/attention/sparse/test_sparse_mla_forward.py index 120771e3a884..af5400243200 100644 --- a/tests/unittest/_torch/attention/sparse/test_sparse_mla_forward.py +++ b/tests/unittest/_torch/attention/sparse/test_sparse_mla_forward.py @@ -22,20 +22,18 @@ import pytest import torch -import torch.nn.functional as F import tensorrt_llm import tensorrt_llm.bindings from tensorrt_llm._torch.attention_backend.interface import ( AttentionBackend, PositionalEmbeddingParams, RopeParams) -from tensorrt_llm._torch.attention_backend.sparse.deepseek_v4 import \ - DeepseekV4CacheManager +from tensorrt_llm._torch.attention_backend.sparse.deepseek_v4 import ( + DeepseekV4CacheManager, DeepseekV4VanillaIndexer) from tensorrt_llm._torch.attention_backend.sparse.deepseek_v4.module import \ forward_context_sparse_attn as forward_context_deepseek_v4_mla from tensorrt_llm._torch.attention_backend.sparse.deepseek_v4.module import \ forward_generation_sparse_attn as forward_generation_deepseek_v4_mla -from tensorrt_llm._torch.attention_backend.sparse.dsa import (HAS_FAST_HADAMARD, - DSACacheManager) +from tensorrt_llm._torch.attention_backend.sparse.dsa import DSACacheManager from tensorrt_llm._torch.attention_backend.sparse.dsa.module import \ _forward_dsa_attn from tensorrt_llm._torch.attention_backend.utils import get_attention_backend @@ -53,8 +51,6 @@ from tensorrt_llm.quantization.mode import QuantAlgo from .deepseek_v4.test_compressor_module import RefCompressor -from .deepseek_v4.test_compressor_module import \ - rotate_activation as ref_rotate_activation # DSACacheManager creates background ThreadPoolExecutor threads. pytestmark = pytest.mark.threadleak(enabled=False) @@ -277,98 +273,6 @@ def _copy_ref_compressor_weights(ref_compressor: RefCompressor, ref_compressor.norm.weight.data.copy_(compressor.norm.weight.data) -def _topk_from_scores(scores: torch.Tensor, topk_tokens: int) -> torch.Tensor: - row = torch.full((topk_tokens, ), - -1, - dtype=torch.int32, - device=scores.device) - if scores.numel() == 0: - return row - k = min(topk_tokens, scores.numel()) - row[:k] = torch.topk(scores.float(), k, dim=-1).indices.to(torch.int32) - return row - - -def _ceil_pow2_scale(amax: torch.Tensor, max_value_inv: float, - min_amax: float) -> torch.Tensor: - scaled = torch.clamp(amax.float(), min=min_amax) * max_value_inv - bits = scaled.contiguous().view(torch.int32) - exp_part = ((bits >> 23) & 0xFF) - 127 - has_mantissa = (bits & 0x7FFFFF).ne(0).to(torch.int32) - log2_ceil = exp_part + has_mantissa - pow2_bits = (log2_ceil + 127) << 23 - return pow2_bits.view(torch.float32) - - -def _e2m1_nibbles_reference(x: torch.Tensor, - round_ties_to_even: bool) -> torch.Tensor: - thresholds = torch.tensor([0.0, 0.25, 0.75, 1.25, 1.75, 2.5, 3.5, 5.0], - device=x.device) - abs_x = x.abs().clamp(max=6.0) - idx = torch.zeros_like(abs_x, dtype=torch.uint8) - for level_idx in range(1, 8): - take_upper = abs_x > thresholds[level_idx] - if round_ties_to_even: - take_upper = take_upper | ((abs_x == thresholds[level_idx]) & - (level_idx % 2 == 0)) - idx = torch.where(take_upper, torch.full_like(idx, level_idx), idx) - sign = torch.signbit(x).to(torch.uint8) << 3 - return idx | sign - - -def _mxfp4_quant_dequant_reference( - x: torch.Tensor, - *, - round_ties_to_even: bool, - min_amax: float, - block_size: int = 32, -) -> torch.Tensor: - assert x.shape[-1] % block_size == 0 - fp4_values = torch.tensor([0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0], - device=x.device, - dtype=torch.float32) - orig_shape = x.shape - blocks = x.float().reshape(-1, x.shape[-1] // block_size, block_size) - scale = _ceil_pow2_scale(blocks.abs().amax(dim=-1, keepdim=True), 1.0 / 6.0, - min_amax) - nibbles = _e2m1_nibbles_reference(blocks / scale, round_ties_to_even) - value_idx = (nibbles & 0x07).to(torch.int64) - values = fp4_values[value_idx] - values = torch.where((nibbles & 0x08) != 0, -values, values) - return (values * scale).reshape(orig_shape) - - -def _maybe_rotate_for_fp4_indexer(x: torch.Tensor) -> torch.Tensor: - if not HAS_FAST_HADAMARD: - return x - return ref_rotate_activation(x) - - -def _prepare_fp4_indexer_q_reference(q: torch.Tensor) -> torch.Tensor: - q = _maybe_rotate_for_fp4_indexer(q) - # DSv4 Q uses fused_cat_fp4, whose E2M1 thresholds round ties down. - return _mxfp4_quant_dequant_reference(q, - round_ties_to_even=False, - min_amax=1.0e-12) - - -def _prepare_fp4_indexer_k_reference(k: torch.Tensor) -> torch.Tensor: - k = _maybe_rotate_for_fp4_indexer(k) - fp4_min_amax = 6.0 * torch.finfo(torch.float32).tiny - # The compressor kernel uses the hardware E2M1 cast, which is RNE. - return _mxfp4_quant_dequant_reference(k, - round_ties_to_even=True, - min_amax=fp4_min_amax) - - -def _reference_indexer_scores(q: torch.Tensor, k: torch.Tensor, - weights: torch.Tensor, - softmax_scale: float) -> torch.Tensor: - head_scores = torch.einsum("hd,kd->hk", q.float(), k.float()) - head_scores = F.relu(head_scores) * softmax_scale - return (head_scores * weights.float().unsqueeze(-1)).sum(dim=0) - - def calculate_reference_deepseek_v4_topk_indices( mla, ref_indexer_compressor: RefCompressor, @@ -385,25 +289,15 @@ def calculate_reference_deepseek_v4_topk_indices( device: torch.device, ) -> torch.Tensor: """Independent PyTorch reference for HF DS4 ratio-4 indexer top-k indices.""" - indexer = mla.mqa.indexer + vanilla_indexer = DeepseekV4VanillaIndexer(mla.mqa.indexer) num_tokens = hidden_states.shape[0] topk_indices = torch.full((num_tokens, topk_tokens), -1, dtype=torch.int32, device=device) - q_ref = F.linear(qr, indexer.wq_b.weight) - q_ref = q_ref.view(num_tokens, indexer.n_heads, indexer.head_dim) - q_ref = q_ref.unsqueeze(0) - apply_rotary_emb(q_ref[..., -indexer.rope_dim:], - freqs_cis[position_ids.long()]) - q_ref = q_ref.squeeze(0) - use_fp4_indexer = indexer.indexer_k_dtype == "fp4" - if use_fp4_indexer: - q_ref = _prepare_fp4_indexer_q_reference(q_ref) - - weights = F.linear(hidden_states, indexer.weights_proj.weight) - weights = weights.float() * (indexer.n_heads**-0.5) + q_ref = vanilla_indexer.project_query(qr, position_ids, freqs_cis) + weights = vanilla_indexer.token_weights(hidden_states) offset = 0 for req_idx in ctx_indices: @@ -421,17 +315,15 @@ def calculate_reference_deepseek_v4_topk_indices( if compressed_kv is not None: compressed_kv = compressed_kv.squeeze(0) - compressed_kv_for_scores = ( - _prepare_fp4_indexer_k_reference(compressed_kv) - if use_fp4_indexer else compressed_kv) + compressed_kv_for_scores = vanilla_indexer.prepare_key( + compressed_kv) for token_idx in range(seq_len): valid_len = (token_idx + 1) // compress_ratio k_for_scores = compressed_kv_for_scores[:valid_len] - scores = _reference_indexer_scores(q_ref[offset + token_idx], - k_for_scores, - weights[offset + token_idx], - indexer.softmax_scale) - topk_indices[offset + token_idx] = _topk_from_scores( + scores = vanilla_indexer.scores(q_ref[offset + token_idx], + k_for_scores, + weights[offset + token_idx]) + topk_indices[offset + token_idx] = vanilla_indexer.topk( scores, topk_tokens) offset += seq_len @@ -470,14 +362,11 @@ def calculate_reference_deepseek_v4_topk_indices( if compressed_kv is None or valid_len == 0: continue - compressed_kv_for_scores = ( - _prepare_fp4_indexer_k_reference(compressed_kv) - if use_fp4_indexer else compressed_kv) - scores = _reference_indexer_scores(q_ref[token_idx], - compressed_kv_for_scores[:valid_len], - weights[token_idx], - indexer.softmax_scale) - topk_indices[token_idx] = _topk_from_scores(scores, topk_tokens) + compressed_kv_for_scores = vanilla_indexer.prepare_key(compressed_kv) + scores = vanilla_indexer.scores(q_ref[token_idx], + compressed_kv_for_scores[:valid_len], + weights[token_idx]) + topk_indices[token_idx] = vanilla_indexer.topk(scores, topk_tokens) return topk_indices diff --git a/tests/unittest/_torch/attention/test_attention_backends.py b/tests/unittest/_torch/attention/test_attention_backends.py index ca00dcd21125..b7f5f0b01509 100644 --- a/tests/unittest/_torch/attention/test_attention_backends.py +++ b/tests/unittest/_torch/attention/test_attention_backends.py @@ -50,8 +50,8 @@ def get_long_seq_len(window: int) -> int: return window + 17 -def _phases_from_window(window: int) -> dict: - long_len = get_long_seq_len(window) +def _phases(long_len: int, gen_len: int) -> dict: + """The ctx/gen/mix batch shapes, parameterized by context and generation len.""" return { "ctx": dict( seq_lens=[long_len, 73, 41], @@ -59,25 +59,38 @@ def _phases_from_window(window: int) -> dict: num_contexts=3, ), "gen": dict( - seq_lens=[1, 1, 1], + seq_lens=[gen_len] * 3, num_cached_tokens=[long_len, 73, 41], num_contexts=0, ), "mix": dict( - seq_lens=[long_len, 1, 1], + seq_lens=[long_len, gen_len, gen_len], num_cached_tokens=[0, long_len, 73], num_contexts=1, ), } +def _phases_from_window(window: int) -> dict: + return _phases(get_long_seq_len(window), gen_len=1) + + # Standard self-attention batch phases. Non-sliding cases use a nominal window # only to choose non-tiny, non-power-of-two lengths; the backend still receives # sliding_window=None. _PHASES = _phases_from_window(_NON_SLIDING_PHASE_WINDOW) +# Model-agnostic sparse sweep dimensions, shared by every sparse config. +_SPARSE_COMPUTE_DTYPE = "bfloat16" +_SPARSE_KV_LAYOUT = "HND" +_SPARSE_PAGE_SIZE = 64 +_SPARSE_USE_KVM_V2 = False + + def _phases_for(cfg: ModelAttnConfig) -> dict: + if cfg.sparse_attention_config is not None: + return _phases(cfg.sparse_topk + 32, gen_len=1) if cfg.mask != "sliding": return _PHASES @@ -117,7 +130,10 @@ def _common(cfg: ModelAttnConfig) -> dict: qk_nope_head_dim=cfg.qk_nope_head_dim, qk_rope_head_dim=cfg.qk_rope_head_dim, v_head_dim=cfg.v_head_dim, + hidden_size=cfg.hidden_size, ) + if cfg.sparse_attention_config is not None: + common.update(sparse_attention_config=cfg.sparse_attention_config) return common @@ -140,6 +156,29 @@ def _expand(cfg: ModelAttnConfig, precisions, kv_layouts, page_sizes): common = _common(cfg) phases = _phases_for(cfg) + # Sparse cases use one model-agnostic sweep (bf16 latent cache, fixed layout/ + # page/manager). Each phase mixes singleton and random selection rows in a + # single execution; the injected selections isolate backend execution from + # the separately-tested algorithm/indexer. + if cfg.sparse_attention_config is not None: + is_deepseek_v4 = cfg.sparse_attention_config.algorithm == "deepseek_v4" + page_size = 128 if is_deepseek_v4 else _SPARSE_PAGE_SIZE + manager = "v2" if is_deepseek_v4 or _SPARSE_USE_KVM_V2 else "v1" + tag = f"{_prec_tag(_SPARSE_COMPUTE_DTYPE, None)}-{_SPARSE_KV_LAYOUT}-p{page_size}-{manager}" + for phase_name in _phases_to_run(cfg, phases): + yield ( + f"{cfg.id}-{phase_name}-{tag}", + BackendCase( + page_size=page_size, + kv_layout=_SPARSE_KV_LAYOUT, + dtype=_SPARSE_COMPUTE_DTYPE, + use_kv_cache_manager_v2=_SPARSE_USE_KVM_V2, + **phases[phase_name], + **common, + ), + ) + return + # Bidirectional, KV-cache-free DiT / encoder workloads: only compute dtype. if cfg.no_cache: for dtype, kvd in precisions: