Skip to content
Closed
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
6 changes: 3 additions & 3 deletions csrc/apis/attention.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
}
Expand All @@ -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] =
Expand Down
62 changes: 62 additions & 0 deletions tests/test_attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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()
Loading