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()