diff --git a/deep_gemm/include/deep_gemm/scheduler/sm90_paged_mqa_logits.cuh b/deep_gemm/include/deep_gemm/scheduler/sm90_paged_mqa_logits.cuh index caa6ba770..2a3a83766 100644 --- a/deep_gemm/include/deep_gemm/scheduler/sm90_paged_mqa_logits.cuh +++ b/deep_gemm/include/deep_gemm/scheduler/sm90_paged_mqa_logits.cuh @@ -99,6 +99,16 @@ void sm90_paged_mqa_logits_metadata(const uint32_t batch_size, const uint32_t ne // The host supplies this because SM90 next_n=4 uses one scheduler item // backed by a two-CTA cluster, while other kernels may atomize next_n. const uint32_t total = sum * num_next_n_atoms; + if (total == 0) { + // Emit only one-past-the-end sentinels. Besides avoiding the + // prefix_sum[batch_size] OOB, this prevents a final SM from + // executing a synthetic zero-KV task. + for (uint32_t sm_idx = lane_idx; sm_idx <= kNumSMs; sm_idx += 32) { + schedule_metadata[sm_idx * 2] = batch_size * num_next_n_atoms; + schedule_metadata[sm_idx * 2 + 1] = 0; + } + return; + } const uint32_t q = total / kNumSMs, r = total % kNumSMs; const uint32_t pivot = kNumSMs - r; for (uint32_t sm_idx = lane_idx; sm_idx < kNumSMs; sm_idx += 32) { @@ -110,7 +120,10 @@ void sm90_paged_mqa_logits_metadata(const uint32_t batch_size, const uint32_t ne lo = pred ? mid + 1 : lo; hi = pred ? hi : mid; } - const uint32_t q_idx = lo; + // lo == batch_size when every context length is zero. Clamp before + // reading prefix_sum[q_idx]; the scheduler constructor separately + // guards empty ranges, so these entries produce no device work. + const uint32_t q_idx = min(lo, batch_size - 1); const uint32_t offset_in_q = (q_idx == 0 ? seg_starts : seg_starts - prefix_sum[q_idx - 1] * num_next_n_atoms); const uint32_t num_segs_q = (q_idx == 0 ? prefix_sum[0] : prefix_sum[q_idx] - prefix_sum[q_idx - 1]); const uint32_t atom_idx = num_segs_q > 0 ? offset_in_q / num_segs_q : 0; @@ -207,8 +220,14 @@ struct SM90PagedMQALogitsScheduler : SM90IndicesStorage { current_q_atom_idx = current_pack.x, current_kv_idx = current_pack.y * kNumBlocksPerSplit; end_q_atom_idx = end_pack.x, end_kv_idx = end_pack.y * kNumBlocksPerSplit; - // NOTES: unconditional call is safe — reversed metadata allocation ensures `current_q_atom_idx` is always in-bounds. - refresh_num_kv_and_advance(current_q_atom_idx); + // Empty metadata ranges may carry the one-past-the-end sentinel (notably + // all-zero context lengths). Do not dereference context_lens or + // indices until this SM actually owns a task. + current_advance = 1; + current_num_kv = 0; + last_advance = 1; + if (exist_q_atom_idx(current_q_atom_idx)) + refresh_num_kv_and_advance(current_q_atom_idx); } // Whether num_kv should be refreshed after advancing to q_atom_idx. @@ -223,6 +242,17 @@ struct SM90PagedMQALogitsScheduler : SM90IndicesStorage { } CUTLASS_DEVICE bool fetch_next_task(uint32_t &q_atom_idx, uint32_t &kv_idx, uint32_t &num_kv) { + // A zero-context request has no KV task. Metadata naturally assigns it + // zero work, but traversal can still cross it between two non-empty + // requests; skip all of its atoms before exposing a task to the kernel. + while (current_num_kv == 0 and + not (current_q_atom_idx == end_q_atom_idx and current_kv_idx == end_kv_idx)) { + current_kv_idx = 0; + current_q_atom_idx += current_advance; + if (should_refresh_num_kv(current_q_atom_idx) and exist_q_atom_idx(current_q_atom_idx)) + refresh_num_kv_and_advance(current_q_atom_idx); + } + q_atom_idx = current_q_atom_idx; kv_idx = current_kv_idx; num_kv = current_num_kv; diff --git a/tests/test_attention.py b/tests/test_attention.py index 019252d62..e8ccfc168 100644 --- a/tests/test_attention.py +++ b/tests/test_attention.py @@ -619,6 +619,51 @@ def make_sparse_kv_block_indices(context_lens: List[int], request_indices: List[ return torch.tensor(indices, device='cuda', dtype=torch.int32), num_blocks_per_q +@test_filter(lambda: get_arch_major() == 9) +def test_paged_mqa_logits_zero_context(): + # A zero context length gives a request no KV work at all. `test_paged_mqa_logits` + # never generates one (context lens are drawn around a positive average), so the + # scheduler's empty-range handling needs its own case: with every length zero the + # binary search in `sm90_paged_mqa_logits_metadata` runs off the end of the batch, + # and reading `prefix_sum[batch_size]` is out of bounds of a shared buffer sized + # to exactly `align(batch_size, 32)` ints. + print('Testing Paged MQA Logits (zero context lengths):') + num_sms = deep_gemm.get_num_sms() + + for block_kv in (32, 64): + for next_n in (1, 2, 4): + # SM90 next_n=4 schedules one item per two-CTA cluster, not per SM. + num_slots = num_sms // (2 if next_n == 4 else 1) + # batch_size == align(batch_size, 32) puts `prefix_sum[batch_size]` exactly + # one element past the end of the kernel's shared memory allocation. + for batch_size in (32, 1024): + # SM90 passes num_next_n_atoms=1, so the one-past-the-end q atom + # index the kernel writes for an empty range is just `batch_size`. + sentinel = batch_size + case = f'block_kv={block_kv}, next_n={next_n}, batch_size={batch_size}' + + # All requests empty: every slot must be the one-past-the-end sentinel. + context_lens = torch.zeros((batch_size, next_n), device='cuda', dtype=torch.int) + metadata = deep_gemm.get_paged_mqa_logits_metadata( + context_lens=context_lens, block_kv=block_kv, num_sms=num_slots) + torch.cuda.synchronize() + assert metadata.size(0) == num_slots + 1, case + assert (metadata[:, 0] == sentinel).all(), f'{case}: {metadata[:, 0].unique().tolist()}' + assert (metadata[:, 1] == 0).all(), f'{case}: {metadata[:, 1].unique().tolist()}' + + # Empty requests interleaved with non-empty ones: scheduled q atoms must + # stay addressable, and the trailing slot must still be the sentinel. + context_lens = torch.zeros((batch_size, next_n), device='cuda', dtype=torch.int) + context_lens[::2] = 512 + metadata = deep_gemm.get_paged_mqa_logits_metadata( + context_lens=context_lens, block_kv=block_kv, num_sms=num_slots) + torch.cuda.synchronize() + q_atom_idx = metadata[:, 0] + assert (q_atom_idx <= sentinel).all(), f'{case}: {q_atom_idx.max().item()} > {sentinel}' + assert q_atom_idx[-1].item() == sentinel, f'{case}: {q_atom_idx[-1].item()}' + print(' > Passed\n') + + @test_filter(lambda: get_arch_major() == 10) def test_sparse_mqa_logits() -> None: num_heads, head_dim = 32, 128 @@ -862,4 +907,5 @@ def enumerate_sparse_mqa_logits(): test_gemm_skip_head_mid() test_mqa_logits() test_paged_mqa_logits() + test_paged_mqa_logits_zero_context() test_sparse_mqa_logits()