From 2a95b271c7decb99a4773d17299f393167ab0a01 Mon Sep 17 00:00:00 2001 From: Woosuk Kwon Date: Thu, 17 Sep 2026 22:05:53 +0000 Subject: [PATCH] Fix sparse MQA merge partitions with repeated block indices Pair only the last Q0 duplicate with the first Q1 duplicate so stable merge partitions preserve sparse slot order and coverage. Co-authored-by: Codex Signed-off-by: Woosuk Kwon --- csrc/apis/attention.hpp | 6 +- .../sm100_sparse_mqa_logits_metadata.cuh | 11 +++- tests/test_attention.py | 62 +++++++++++++++++++ 3 files changed, 74 insertions(+), 5 deletions(-) diff --git a/csrc/apis/attention.hpp b/csrc/apis/attention.hpp index 1d0903f12..f730f1dd6 100644 --- a/csrc/apis/attention.hpp +++ b/csrc/apis/attention.hpp @@ -239,7 +239,7 @@ static const torch::Tensor& get_sparse_mqa_logits_workspace(const torch::TensorO } // Each sparse-index row starts with the valid blocks inferred from its KV length. This prefix must -// contain unique, strictly increasing absolute block indices within the corresponding KV range. +// contain nondecreasing absolute block indices within the corresponding KV range; repeats are allowed. // With unaligned ks, block i starts at i * sparse_block_kv + ks % sparse_block_kv. static torch::Tensor get_sparse_mqa_logits_metadata(const torch::Tensor& cu_seq_len_k_start, const torch::Tensor& cu_seq_len_k_end, @@ -269,8 +269,8 @@ static torch::Tensor get_sparse_mqa_logits_metadata(const torch::Tensor& cu_seq_ return metadata; } -// Queries belonging to one request must be consecutive. Each sparse-index row starts with unique, -// strictly increasing logical block indices within its context length. Paired queries must also +// Queries belonging to one request must be consecutive. Each sparse-index row starts with nondecreasing +// logical block indices within its context length; repeats are allowed. Paired queries must also // have identical block-table rows. static torch::Tensor get_paged_sparse_mqa_logits_metadata(const torch::Tensor& context_lens, const torch::Tensor& block_table, diff --git a/deep_gemm/include/deep_gemm/scheduler/sm100_sparse_mqa_logits_metadata.cuh b/deep_gemm/include/deep_gemm/scheduler/sm100_sparse_mqa_logits_metadata.cuh index ec5dbd188..3894e9ef4 100644 --- a/deep_gemm/include/deep_gemm/scheduler/sm100_sparse_mqa_logits_metadata.cuh +++ b/deep_gemm/include/deep_gemm/scheduler/sm100_sparse_mqa_logits_metadata.cuh @@ -292,7 +292,9 @@ void sm100_sparse_mqa_logits_metadata( // Drop a duplicate carried across merge partitions if (num_remaining_inputs > 0 and q0_slot_idx > 0 and q1_slot_idx < num_kv_blocks_in_q1 and - ptx::ld_shared(smem.logical_kv_block_indices[0] + q0_slot_idx - 1) == ptx::ld_shared(smem.logical_kv_block_indices[1] + q1_slot_idx)) { + ptx::ld_shared(smem.logical_kv_block_indices[0] + q0_slot_idx - 1) == ptx::ld_shared(smem.logical_kv_block_indices[1] + q1_slot_idx) and + (q0_slot_idx == num_kv_blocks_in_q0 or ptx::ld_shared(smem.logical_kv_block_indices[0] + q0_slot_idx) != ptx::ld_shared(smem.logical_kv_block_indices[1] + q1_slot_idx)) and + (q1_slot_idx == 0 or ptx::ld_shared(smem.logical_kv_block_indices[1] + q1_slot_idx - 1) != ptx::ld_shared(smem.logical_kv_block_indices[1] + q1_slot_idx))) { ++ q1_slot_idx; -- num_remaining_inputs; } @@ -313,7 +315,12 @@ void sm100_sparse_mqa_logits_metadata( const uint32_t q0_logical_kv_block_idx = q0_slot_idx < num_kv_blocks_in_q0 ? ptx::ld_shared(smem.logical_kv_block_indices[0] + q0_slot_idx) : ~0u; const uint32_t q1_logical_kv_block_idx = q1_slot_idx < num_kv_blocks_in_q1 ? ptx::ld_shared(smem.logical_kv_block_indices[1] + q1_slot_idx) : ~0u; const bool in_q0 = q0_logical_kv_block_idx <= q1_logical_kv_block_idx; - const bool in_q1 = q1_logical_kv_block_idx <= q0_logical_kv_block_idx; + // Stable merge-path puts all equal Q0 slots before Q1. Only pair + // the last equal Q0 slot with the first Q1 slot, including repeats. + const bool last_equal_q0 = q0_slot_idx + 1 >= num_kv_blocks_in_q0 or + ptx::ld_shared(smem.logical_kv_block_indices[0] + q0_slot_idx + 1) != q0_logical_kv_block_idx; + const bool in_q1 = q1_logical_kv_block_idx < q0_logical_kv_block_idx or + (q1_logical_kv_block_idx == q0_logical_kv_block_idx and last_equal_q0); // Carry a final duplicate into the next partition const bool consume_q1 = in_q1 and num_remaining_inputs > in_q0; packed_slots_in_thread[merged_kv_block_offset_in_thread] = diff --git a/tests/test_attention.py b/tests/test_attention.py index daa1b54c8..2a940eb8f 100644 --- a/tests/test_attention.py +++ b/tests/test_attention.py @@ -665,6 +665,67 @@ def test_paged_mqa_logits_zero_context(): print(' > Passed\n') +@test_filter(lambda: get_arch_major() == 10) +def test_sparse_mqa_logits_repeated_blocks() -> None: + """Repeated padding must preserve every Q slot across merge partitions.""" + torch.manual_seed(42) + num_heads, head_dim, sparse_block_kv, max_blocks = 32, 128, 8, 2048 + # 389 repeated blocks make adjacent merge partitions consume different + # numbers of inputs. The old merge can move Q1's slot backwards here. + lengths = torch.tensor([3105, 3106], dtype=torch.int32, device='cuda') + for fmt in ('mxfp4', 'mxfp8'): + cast_fwd = per_token_cast_to_fp4 if fmt == 'mxfp4' else per_token_cast_to_fp8 + elem_dim = head_dim // 2 if fmt == 'mxfp4' else head_dim + q_fp, q_sf = cast_fwd(torch.randn((2 * num_heads, head_dim), device='cuda', dtype=torch.bfloat16), + use_ue8m0=True, gran_k=32, use_packed_ue8m0=True) + q = q_fp.view(2, num_heads, elem_dim), q_sf.view(2, num_heads) + weights = to_mqa_weights(torch.randn((2, num_heads), device='cuda', dtype=torch.bfloat16), torch.bfloat16) + for mode, start in (('aligned', 0), ('unaligned', 3), ('paged', 0)): + starts = torch.full_like(lengths, start) + ends = starts + lengths + num_kv_tokens = 3200 + kv_input = torch.randn((num_kv_tokens, head_dim), device='cuda', dtype=torch.bfloat16) + kv_fp, kv_sf = cast_fwd(kv_input, + use_ue8m0=True, gran_k=32, use_packed_ue8m0=True) + kv = kv_fp, kv_sf.view(num_kv_tokens) + full = deep_gemm.fp8_fp4_mqa_logits(q, kv, weights, starts, ends, + clean_logits=False, max_seqlen_k=num_kv_tokens, + logits_dtype=torch.bfloat16) + if mode == 'paged': + page_kv = 64 + cast_cache = kv_cache_cast_to_mxfp4 if fmt == 'mxfp4' else kv_cache_cast_to_mxfp8 + pages, _ = cast_cache(kv_input.view(-1, page_kv, 1, head_dim)) + page_stride = ceil_div(page_kv * (elem_dim + 4), 512) * 512 + storage = torch.empty((pages.shape[0], page_stride), device='cuda', dtype=torch.uint8) + cache = storage.as_strided(pages.shape, (page_stride, elem_dim + 4, elem_dim + 4, 1)) + cache.copy_(pages) + table = torch.arange(pages.shape[0], device='cuda', dtype=torch.int32).repeat(2, 1) + requests = torch.zeros_like(lengths) + paged_q = q[0].unsqueeze(1), q[1].unsqueeze(1) + for prefix in (0, 8): + sparse_indices = torch.full((2, max_blocks), 388, dtype=torch.int32, device='cuda') + sparse_indices[:, :prefix] = torch.arange(prefix, dtype=torch.int32, device='cuda') + if mode == 'paged': + metadata = deep_gemm.get_paged_sparse_mqa_logits_metadata( + lengths, table, requests, page_kv, sparse_indices, q_fp.dtype, sparse_block_kv) + actual = deep_gemm.fp8_fp4_paged_sparse_mqa_logits( + paged_q, cache, weights, metadata, max_blocks, sparse_block_kv) + else: + metadata = deep_gemm.get_sparse_mqa_logits_metadata( + starts, ends, num_kv_tokens, sparse_indices, q_fp.dtype, sparse_block_kv, + use_unaligned_ks=bool(start)) + actual = deep_gemm.fp8_fp4_sparse_mqa_logits( + q, kv, weights, metadata, max_blocks, sparse_block_kv, use_unaligned_ks=bool(start)) + token_indices = (sparse_indices[:, :, None] * sparse_block_kv + start + + torch.arange(sparse_block_kv, device='cuda')).flatten(1).long() + counts = ceil_div(lengths, sparse_block_kv) * sparse_block_kv + valid = ((torch.arange(actual.shape[1], device='cuda')[None, :] < counts[:, None]) + & (token_indices < ends[:, None])) + expected = full.gather(1, token_indices - start) + assert_bitwise_equal(actual[valid], expected[valid], f'repeated sparse blocks: {fmt}, {mode}, {prefix=}') + print(' > Repeated sparse blocks passed\n') + + @test_filter(lambda: get_arch_major() == 10) def test_sparse_mqa_logits() -> None: num_heads, head_dim = 32, 128 @@ -909,4 +970,5 @@ def enumerate_sparse_mqa_logits(): test_mqa_logits() test_paged_mqa_logits() test_paged_mqa_logits_zero_context() + test_sparse_mqa_logits_repeated_blocks() test_sparse_mqa_logits()