Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
36 changes: 33 additions & 3 deletions deep_gemm/include/deep_gemm/scheduler/sm90_paged_mqa_logits.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand All @@ -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;
Expand Down Expand Up @@ -207,8 +220,14 @@ struct SM90PagedMQALogitsScheduler : SM90IndicesStorage<kIsVarlen> {
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.
Expand All @@ -223,6 +242,17 @@ struct SM90PagedMQALogitsScheduler : SM90IndicesStorage<kIsVarlen> {
}

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;
Expand Down
46 changes: 46 additions & 0 deletions tests/test_attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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()
Loading