diff --git a/tests/v1/attention/test_dspark_noncausal_sparse_mla.py b/tests/v1/attention/test_dspark_noncausal_sparse_mla.py index d5c794131474..30c35a730efc 100644 --- a/tests/v1/attention/test_dspark_noncausal_sparse_mla.py +++ b/tests/v1/attention/test_dspark_noncausal_sparse_mla.py @@ -604,3 +604,153 @@ def test_dspark_noncausal_differs_from_causal( f"non-causal backend output matches the causal reference " f"(max abs diff={causal_err}); future-pointing indices are not attended to" ) + + +@pytest.mark.parametrize("context_len", [20, 128, 900]) +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float8_e4m3fn]) +@pytest.mark.parametrize("num_heads", [16, 64]) +def test_dsv41_flashinfer_dspark_window_matches_reference( + context_len, dtype, num_heads, monkeypatch +): + """Draft queries see the full block without treating padded slots as keys.""" + if not current_platform.is_device_capability_family(100): + pytest.skip("DSV4 TRTLLM sparse attention requires SM100") + from vllm.models.deepseek_v41.nvidia.flashinfer_sparse import ( + DeepseekSparseSWAFlashInferMetadataBuilder, + DeepseekV4FlashInferMLAAttention, + ) + from vllm.models.deepseek_v41.sparse_mla import ( + DeepseekV41SparseSWAMetadataBuilder, + ) + + torch.manual_seed(123) + device = "cuda" + query_lens = [5, 3] + context_lens = [context_len, context_len + 137] + num_real_tokens, num_tokens, head_dim = 8, 9, 512 + cache = torch.randn(32, 128, head_dim, device=device, dtype=torch.bfloat16).to( + dtype + ) + query = torch.randn( + num_tokens, num_heads, head_dim, device=device, dtype=torch.bfloat16 + ).to(dtype) + indices = torch.full((num_tokens, 256), -1, device=device, dtype=torch.int32) + visible_indices = [] + visible_lens = [] + for req, (context, query_len) in enumerate(zip(context_lens, query_lens)): + visible = ( + torch.arange(max(context - 128, 0), context + query_len, device=device) + + req * 16 * 128 + ) + visible_indices.extend([visible] * query_len) + visible_lens.extend([visible.numel()] * query_len) + for token, visible in enumerate(visible_indices): + indices[token, : visible.numel()] = visible.to(torch.int32) + query_start_loc = torch.tensor([0, 5, 8, 9], dtype=torch.int32) + metadata = SimpleNamespace( + num_decodes=3, + num_prefills=0, + num_decode_tokens=num_tokens, + num_prefill_tokens=0, + seq_lens=torch.tensor( + [context_lens[0] + 5, context_lens[1] + 3, 1], + device=device, + dtype=torch.int32, + ), + query_start_loc=query_start_loc.to(device), + query_start_loc_cpu=query_start_loc, + token_to_req_indices=torch.tensor( + [0] * 5 + [1] * 3 + [2], device=device, dtype=torch.int32 + ), + decode_swa_indices=indices, + decode_swa_width=256, + decode_swa_lens=torch.tensor( + visible_lens + [0], device=device, dtype=torch.int32 + ), + block_table=torch.arange(48, device=device, dtype=torch.int32).view(3, -1), + block_size=128, + flashinfer_sparse_index_cache={}, + max_decode_query_len=5, + ) + + # Exercise FlashInfer preparation without constructing a model/config. + def init_parent(builder): + builder._max_tokens = num_tokens + builder.device = device + builder.window_size = 128 + + monkeypatch.setattr(DeepseekV41SparseSWAMetadataBuilder, "__init__", init_parent) + monkeypatch.setattr( + DeepseekV41SparseSWAMetadataBuilder, + "build", + lambda *args: metadata, + ) + builder = DeepseekSparseSWAFlashInferMetadataBuilder() + common_metadata = SimpleNamespace(causal=False) + builder.build(0, common_metadata) + prepared = ( + metadata.flashinfer_decode_topk_lens, + metadata.flashinfer_decode_seq_lens, + ) + attention = SimpleNamespace( + kv_cache_torch_dtype=dtype, + window_size=128, + compress_ratio=0, + topk_indices_buffer=torch.empty( + num_tokens, 0, device=device, dtype=torch.int32 + ), + scale=1 / math.sqrt(head_dim), + _flashinfer_fp8_bmm1_scale=1 / math.sqrt(head_dim), + _flashinfer_fp8_bmm2_scale=1.0, + attn_sink=None, + ) + attention._build_sparse_index_metadata = MethodType( + DeepseekV4FlashInferMLAAttention._build_sparse_index_metadata, attention + ) + output = torch.empty_like(query, dtype=torch.bfloat16) + + def forward(): + metadata.flashinfer_sparse_index_cache.clear() + DeepseekV4FlashInferMLAAttention._forward( + attention, query, None, cache, metadata, None, True, output + ) + + def check_output(): + references = [] + for token, visible in enumerate(visible_indices): + keys = cache.flatten(0, 1)[visible].float() + weights = torch.softmax( + query[token].float() @ keys.T / math.sqrt(head_dim), -1 + ) + references.append(weights @ keys) + atol = 0.01 if dtype == torch.bfloat16 else 0.05 + torch.testing.assert_close( + output[:num_real_tokens].float(), + torch.stack(references), + atol=atol, + rtol=0.05, + ) + + forward() + check_output() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + forward() + # Replay must consume the current inputs, including KV written since capture. + query.copy_(torch.randn_like(query, dtype=torch.bfloat16).to(dtype)) + cache.copy_(torch.randn_like(cache, dtype=torch.bfloat16).to(dtype)) + # A new step changes visibility as well as Q/KV, without recapturing. + metadata.seq_lens.sub_(1) + metadata.decode_swa_lens[:num_real_tokens].sub_(1) + for token, visible in enumerate(visible_indices): + indices[token, visible.numel() - 1] = -1 + visible_indices[token] = visible[:-1] + builder.build(0, common_metadata) + assert metadata.flashinfer_decode_topk_lens.data_ptr() == prepared[0].data_ptr() + assert metadata.flashinfer_decode_seq_lens.data_ptr() == prepared[1].data_ptr() + torch.testing.assert_close(prepared[0], metadata.decode_swa_lens.clamp_min(128)) + torch.testing.assert_close( + prepared[1], metadata.seq_lens[metadata.token_to_req_indices.long()] + ) + graph.replay() + check_output() diff --git a/vllm/models/deepseek_v41/nvidia/flashinfer_sparse.py b/vllm/models/deepseek_v41/nvidia/flashinfer_sparse.py index 409725434a17..a247ae7d7409 100644 --- a/vllm/models/deepseek_v41/nvidia/flashinfer_sparse.py +++ b/vllm/models/deepseek_v41/nvidia/flashinfer_sparse.py @@ -26,7 +26,11 @@ ) from vllm.platforms.interface import DeviceCapability from vllm.utils.flashinfer import flashinfer_trtllm_batch_decode_sparse_mla_dsv4 -from vllm.v1.attention.backend import AttentionCGSupport, MultipleOf +from vllm.v1.attention.backend import ( + AttentionCGSupport, + CommonAttentionMetadata, + MultipleOf, +) from vllm.v1.attention.backends.mla.compressor_utils import ( get_dspark_swa_index_width, ) @@ -183,6 +187,39 @@ class DeepseekSparseSWAFlashInferMetadataBuilder(DeepseekV41SparseSWAMetadataBui _cudagraph_support: ClassVar[AttentionCGSupport] = AttentionCGSupport.ALWAYS + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + # Graphs retain these addresses while each build refreshes their contents. + self._decode_topk_lens = torch.empty( + self._max_tokens, dtype=torch.int32, device=self.device + ) + self._decode_seq_lens = torch.empty_like(self._decode_topk_lens) + + def build( + self, + common_prefix_len: int, + common_attn_metadata: CommonAttentionMetadata, + fast_build: bool = False, + ) -> "DeepseekSparseSWAMetadata": + metadata = super().build(common_prefix_len, common_attn_metadata, fast_build) + num_tokens = metadata.num_decode_tokens + if not common_attn_metadata.causal and num_tokens > 0: + assert metadata.decode_swa_lens is not None + assert metadata.seq_lens is not None + assert metadata.token_to_req_indices is not None + topk_lens = self._decode_topk_lens[:num_tokens] + seq_lens = self._decode_seq_lens[:num_tokens] + torch.clamp(metadata.decode_swa_lens, min=self.window_size, out=topk_lens) + torch.index_select( + metadata.seq_lens, + 0, + metadata.token_to_req_indices[:num_tokens], + out=seq_lens, + ) + metadata.flashinfer_decode_topk_lens = topk_lens + metadata.flashinfer_decode_seq_lens = seq_lens + return metadata + class DeepseekSparseSWAFlashInferBackend(DeepseekSparseSWABackend): @staticmethod @@ -493,21 +530,40 @@ def _forward( # Keep the TRTLLM-gen decode/prefill split: the launcher is tuned for # uniform-q batches, and this avoids flattening mixed batches into one call. if num_decode_tokens > 0: + decode_query = query[:num_decode_tokens] + decode_output = output[:num_decode_tokens] decode_cu = query_start_loc[: num_decodes + 1] + decode_seq_lens = seq_lens[:num_decodes] + decode_topk_lens = sparse_topk_lens[:num_decode_tokens] + max_decode_query_len = swa_metadata.max_decode_query_len + if swa_metadata.decode_swa_width > self.window_size: + # DSpark's non-causal window extends past the fixed 128 SWA + # columns into the aliased compressed pool. Exclude padding, + # and expose the full block to each query instead of letting + # TRTLLM derive a causal SWA length from its query position. + assert swa_only + assert swa_metadata.flashinfer_decode_topk_lens is not None + assert swa_metadata.flashinfer_decode_seq_lens is not None + decode_topk_lens = swa_metadata.flashinfer_decode_topk_lens + decode_seq_lens = swa_metadata.flashinfer_decode_seq_lens + decode_query = decode_query.unsqueeze(1) + decode_output = decode_output.unsqueeze(1) + decode_cu = None + max_decode_query_len = 1 flashinfer_trtllm_batch_decode_sparse_mla_dsv4( - query=query[:num_decode_tokens], + query=decode_query, swa_kv_cache=swa_k_cache, workspace_buffer=workspace, sparse_indices=sparse_indices[:num_decode_tokens], compressed_kv_cache=compressed_kv_cache, - sparse_topk_lens=sparse_topk_lens[:num_decode_tokens], - seq_lens=seq_lens[:num_decodes], - out=output[:num_decode_tokens], + sparse_topk_lens=decode_topk_lens, + seq_lens=decode_seq_lens, + out=decode_output, bmm1_scale=bmm1_scale, bmm2_scale=bmm2_scale, sinks=self.attn_sink, cum_seq_lens_q=decode_cu, - max_q_len=swa_metadata.max_decode_query_len, + max_q_len=max_decode_query_len, ) if num_prefill_tokens > 0: diff --git a/vllm/v1/attention/backends/mla/sparse_swa.py b/vllm/v1/attention/backends/mla/sparse_swa.py index e88631c32bb3..0fa5ae414ce5 100644 --- a/vllm/v1/attention/backends/mla/sparse_swa.py +++ b/vllm/v1/attention/backends/mla/sparse_swa.py @@ -205,6 +205,8 @@ class DeepseekSparseSWAMetadata: num_decode_tokens: int = 0 num_prefill_tokens: int = 0 max_decode_query_len: int = 1 + flashinfer_decode_topk_lens: torch.Tensor | None = None + flashinfer_decode_seq_lens: torch.Tensor | None = None # Pre-computed prefill metadata shared across all DeepseekV4 attention layers. prefill_seq_lens: torch.Tensor | None = None