From 7bb247a83d26bff1a3a53cb99e5510d1f4cfd2aa Mon Sep 17 00:00:00 2001 From: Po-Han Huang Date: Mon, 3 Aug 2026 01:20:03 -0700 Subject: [PATCH] Make paged MQA metadata Torch ABI independent --- csrc/jit_kernels/impls/sm100_mqa_logits.hpp | 27 +++++--- .../impls/sm120_paged_mqa_logits.hpp | 30 ++++++--- .../jit_kernels/impls/sm90_fp8_mqa_logits.hpp | 27 +++++--- csrc/tvm_ffi_api.cpp | 63 ++++++++++++++++--- sgl_deep_gemm/tests/test_attention.py | 25 ++++++++ 5 files changed, 140 insertions(+), 32 deletions(-) diff --git a/csrc/jit_kernels/impls/sm100_mqa_logits.hpp b/csrc/jit_kernels/impls/sm100_mqa_logits.hpp index e8c5b2fac3..75a919d8e6 100644 --- a/csrc/jit_kernels/impls/sm100_mqa_logits.hpp +++ b/csrc/jit_kernels/impls/sm100_mqa_logits.hpp @@ -52,12 +52,12 @@ static void __instantiate_kernel() {{ } }; -static void sm100_paged_mqa_logits_metadata(const torch::Tensor& context_lens, - const torch::Tensor& schedule_meta, - const int& num_requests, const int& num_q_tokens_total, - const int& next_n, const int& num_sms, - const bool& is_context_lens_2d, const bool& is_varlen, - const int* indices_ptr) { +static void sm100_paged_mqa_logits_metadata_raw(int* context_lens_ptr, + int* schedule_meta_ptr, + const int& num_requests, const int& num_q_tokens_total, + const int& next_n, const int& num_sms, + const bool& is_context_lens_2d, const bool& is_varlen, + const int* indices_ptr) { constexpr int split_kv = 256; const int num_threads = 256; // smem: prefix_work[num_requests] + request_q_token_start[num_requests] @@ -72,9 +72,9 @@ static void sm100_paged_mqa_logits_metadata(const torch::Tensor& context_lens, .num_sms = num_sms, .num_requests = num_requests, .num_q_tokens_total = num_q_tokens_total, - .context_lens = context_lens.data_ptr(), + .context_lens = context_lens_ptr, .indices = const_cast(indices_ptr), - .schedule_meta = schedule_meta.data_ptr(), + .schedule_meta = schedule_meta_ptr, .launch_args = LaunchArgs(1, num_threads, smem_size) }; const auto code = SM100PagedMQALogitsMetadataRuntime::generate(args); @@ -82,6 +82,17 @@ static void sm100_paged_mqa_logits_metadata(const torch::Tensor& context_lens, SM100PagedMQALogitsMetadataRuntime::launch(runtime, args); } +static void sm100_paged_mqa_logits_metadata(const torch::Tensor& context_lens, + const torch::Tensor& schedule_meta, + const int& num_requests, const int& num_q_tokens_total, + const int& next_n, const int& num_sms, + const bool& is_context_lens_2d, const bool& is_varlen, + const int* indices_ptr) { + sm100_paged_mqa_logits_metadata_raw( + context_lens.data_ptr(), schedule_meta.data_ptr(), num_requests, + num_q_tokens_total, next_n, num_sms, is_context_lens_2d, is_varlen, indices_ptr); +} + // Unified contiguous-KV runtime for FP4/FP8; FP8 reuses the unused `sf_q` descriptor slot class SM100MQALogitsRuntime final: public LaunchRuntime { public: diff --git a/csrc/jit_kernels/impls/sm120_paged_mqa_logits.hpp b/csrc/jit_kernels/impls/sm120_paged_mqa_logits.hpp index b1ac3b358a..4f8e1d081c 100644 --- a/csrc/jit_kernels/impls/sm120_paged_mqa_logits.hpp +++ b/csrc/jit_kernels/impls/sm120_paged_mqa_logits.hpp @@ -53,13 +53,13 @@ static void __instantiate_kernel() {{ } }; -static void sm120_paged_mqa_logits_metadata(const torch::Tensor& context_lens, - const torch::Tensor& schedule_metadata, - const int& batch_size, const int& next_n, - const int& block_kv, const int& num_sms, - const bool& is_context_lens_2d, - const int& num_next_n_atoms, - const bool& is_varlen, const int* indices_ptr) { +static void sm120_paged_mqa_logits_metadata_raw(int* context_lens_ptr, + int* schedule_metadata_ptr, + const int& batch_size, const int& next_n, + const int& block_kv, const int& num_sms, + const bool& is_context_lens_2d, + const int& num_next_n_atoms, + const bool& is_varlen, const int* indices_ptr) { constexpr int split_kv = 128; constexpr int num_threads = 32; const int aligned_batch_size = align(batch_size, 32); @@ -78,9 +78,9 @@ static void sm120_paged_mqa_logits_metadata(const torch::Tensor& context_lens, .next_n = next_n, .num_next_n_atoms = num_next_n_atoms, .is_context_lens_2d = is_context_lens_2d, - .context_lens = context_lens.data_ptr(), + .context_lens = context_lens_ptr, .indices = const_cast(indices_ptr), - .schedule_metadata = schedule_metadata.data_ptr(), + .schedule_metadata = schedule_metadata_ptr, .launch_args = LaunchArgs(1, num_threads, smem_size) }; const auto code = SM120PagedMQALogitsMetadataRuntime::generate(args); @@ -88,6 +88,18 @@ static void sm120_paged_mqa_logits_metadata(const torch::Tensor& context_lens, SM120PagedMQALogitsMetadataRuntime::launch(runtime, args); } +static void sm120_paged_mqa_logits_metadata(const torch::Tensor& context_lens, + const torch::Tensor& schedule_metadata, + const int& batch_size, const int& next_n, + const int& block_kv, const int& num_sms, + const bool& is_context_lens_2d, + const int& num_next_n_atoms, + const bool& is_varlen, const int* indices_ptr) { + sm120_paged_mqa_logits_metadata_raw( + context_lens.data_ptr(), schedule_metadata.data_ptr(), batch_size, next_n, + block_kv, num_sms, is_context_lens_2d, num_next_n_atoms, is_varlen, indices_ptr); +} + // ---- FP8 paged ---- class SM120FP8PagedMQALogitsRuntime final: public LaunchRuntime { public: diff --git a/csrc/jit_kernels/impls/sm90_fp8_mqa_logits.hpp b/csrc/jit_kernels/impls/sm90_fp8_mqa_logits.hpp index 7e3cea5b22..371ffcd122 100644 --- a/csrc/jit_kernels/impls/sm90_fp8_mqa_logits.hpp +++ b/csrc/jit_kernels/impls/sm90_fp8_mqa_logits.hpp @@ -196,12 +196,12 @@ static void __instantiate_kernel() {{ } }; -static void sm90_paged_mqa_logits_metadata(const torch::Tensor& context_lens, - const torch::Tensor& schedule_metadata, - const int& batch_size, const int& next_n, - const int& block_kv, const int& num_sms, - const bool& is_context_lens_2d, - const bool& is_varlen, const int* indices_ptr) { +static void sm90_paged_mqa_logits_metadata_raw(int* context_lens_ptr, + int* schedule_metadata_ptr, + const int& batch_size, const int& next_n, + const int& block_kv, const int& num_sms, + const bool& is_context_lens_2d, + const bool& is_varlen, const int* indices_ptr) { constexpr int split_kv = 256; constexpr int num_threads = 32; const int aligned_batch_size = align(batch_size, 32); @@ -219,9 +219,9 @@ static void sm90_paged_mqa_logits_metadata(const torch::Tensor& context_lens, .batch_size = batch_size, .next_n = next_n, .is_context_lens_2d = is_context_lens_2d, - .context_lens = context_lens.data_ptr(), + .context_lens = context_lens_ptr, .indices = const_cast(indices_ptr), - .schedule_metadata = schedule_metadata.data_ptr(), + .schedule_metadata = schedule_metadata_ptr, .launch_args = LaunchArgs(1, num_threads, smem_size) }; const auto code = SM90PagedMQALogitsMetadataRuntime::generate(args); @@ -229,6 +229,17 @@ static void sm90_paged_mqa_logits_metadata(const torch::Tensor& context_lens, SM90PagedMQALogitsMetadataRuntime::launch(runtime, args); } +static void sm90_paged_mqa_logits_metadata(const torch::Tensor& context_lens, + const torch::Tensor& schedule_metadata, + const int& batch_size, const int& next_n, + const int& block_kv, const int& num_sms, + const bool& is_context_lens_2d, + const bool& is_varlen, const int* indices_ptr) { + sm90_paged_mqa_logits_metadata_raw( + context_lens.data_ptr(), schedule_metadata.data_ptr(), batch_size, next_n, + block_kv, num_sms, is_context_lens_2d, is_varlen, indices_ptr); +} + class SM90FP8PagedMQALogitsRuntime final: public LaunchRuntime { public: struct Args { diff --git a/csrc/tvm_ffi_api.cpp b/csrc/tvm_ffi_api.cpp index bee2310a7f..ffc9fc0f69 100644 --- a/csrc/tvm_ffi_api.cpp +++ b/csrc/tvm_ffi_api.cpp @@ -1,3 +1,4 @@ +#include #include #include #include @@ -573,13 +574,61 @@ Tensor dg_fp8_mqa_logits(TensorView q, TensorView kv_data, TensorView kv_sf, Tensor dg_get_paged_mqa_logits_metadata(TensorView context_lens, int64_t block_kv, int64_t num_sms, Optional indices) { - auto indices_val = indices.has_value()? - std::optional(convert_to_torch_tensor(indices.value())) - : std::nullopt; - auto result = attention::get_paged_mqa_logits_metadata( - convert_to_torch_tensor(context_lens), static_cast(block_kv), - static_cast(num_sms), indices_val); - return Tensor::FromDLPack(at::toDLPack(result)); + DG_HOST_ASSERT(context_lens.ndim() == 2); + DG_HOST_ASSERT(context_lens.dtype().code == kDLInt and context_lens.dtype().bits == 32 and + context_lens.dtype().lanes == 1); + DG_HOST_ASSERT(context_lens.IsContiguous()); + + const int batch_size = static_cast(context_lens.size(0)); + const int next_n = static_cast(context_lens.size(1)); + const bool is_varlen = indices.has_value(); + int* indices_ptr = nullptr; + if (is_varlen) { + const auto indices_view = indices.value(); + DG_HOST_ASSERT(indices_view.ndim() == 1 and indices_view.size(0) == batch_size); + DG_HOST_ASSERT(indices_view.dtype().code == kDLInt and indices_view.dtype().bits == 32 and + indices_view.dtype().lanes == 1); + DG_HOST_ASSERT(indices_view.IsContiguous()); + indices_ptr = reinterpret_cast( + static_cast(indices_view.data_ptr()) + indices_view.byte_offset()); + } + + std::array output_shape{num_sms + 1, 2}; + auto schedule_metadata = Tensor::FromEnvAlloc( + TVMFFIEnvTensorAlloc, ShapeView(output_shape.data(), output_shape.size()), + context_lens.dtype(), context_lens.device()); + auto* context_lens_ptr = reinterpret_cast( + static_cast(context_lens.data_ptr()) + context_lens.byte_offset()); + auto* schedule_metadata_ptr = static_cast(schedule_metadata.data_ptr()); + + const auto arch_major = device_runtime->get_arch_major(); + if (is_varlen) { + DG_HOST_ASSERT(arch_major == 10 and next_n == 1 and (block_kv == 64 or block_kv == 32)); + sm100_paged_mqa_logits_metadata_raw( + context_lens_ptr, schedule_metadata_ptr, batch_size, batch_size * next_n, next_n, + static_cast(num_sms), true, true, indices_ptr); + } else if (arch_major == 10) { + DG_HOST_ASSERT(block_kv == 64 or block_kv == 32); + sm100_paged_mqa_logits_metadata_raw( + context_lens_ptr, schedule_metadata_ptr, batch_size, batch_size * next_n, next_n, + static_cast(num_sms), true, false, nullptr); + } else if (arch_major == 12) { + DG_HOST_ASSERT(block_kv == 64); + const int next_n_atom = (next_n >= 2) ? 2 : 1; + const int num_next_n_atoms = (next_n + next_n_atom - 1) / next_n_atom; + sm120_paged_mqa_logits_metadata_raw( + context_lens_ptr, schedule_metadata_ptr, batch_size, next_n, + static_cast(block_kv), static_cast(num_sms), true, num_next_n_atoms, + false, nullptr); + } else if (arch_major == 9) { + DG_HOST_ASSERT(block_kv == 64); + sm90_paged_mqa_logits_metadata_raw( + context_lens_ptr, schedule_metadata_ptr, batch_size, next_n, + static_cast(block_kv), static_cast(num_sms), true, false, nullptr); + } else { + DG_HOST_UNREACHABLE("Unsupported architecture"); + } + return schedule_metadata; } Tensor dg_fp8_paged_mqa_logits(TensorView q, TensorView fused_kv_cache, diff --git a/sgl_deep_gemm/tests/test_attention.py b/sgl_deep_gemm/tests/test_attention.py index aa2e400c3e..40cc191907 100644 --- a/sgl_deep_gemm/tests/test_attention.py +++ b/sgl_deep_gemm/tests/test_attention.py @@ -77,6 +77,30 @@ def ref_diff_tol(has_bf16: bool) -> float: return 3e-5 if has_bf16 else 5e-6 +def test_paged_mqa_logits_metadata() -> None: + print('Testing paged MQA logits metadata:') + num_sms = deep_gemm.get_num_sms() + cases = ( + [[64]], + [[64, 128]], + [[1], [63], [64], [65], [1024]], + [[31, 32, 33], [63, 64, 65]], + ) + + for values in cases: + context_lens = torch.tensor(values, device='cuda', dtype=torch.int32) + schedule_meta = deep_gemm.get_paged_mqa_logits_metadata( + context_lens, 64, num_sms + ) + torch.cuda.synchronize() + + assert schedule_meta.shape == (num_sms + 1, 2) + assert schedule_meta.dtype == torch.int32 + assert schedule_meta.device == context_lens.device + assert schedule_meta.is_contiguous() + print() + + def dtype_tag(dtype: torch.dtype) -> str: return 'BF16' if dtype == torch.bfloat16 else 'FP32' @@ -493,5 +517,6 @@ def enumerate_paged_mqa_logits(): random.seed(0) test_gemm_skip_head_mid() + test_paged_mqa_logits_metadata() test_mqa_logits() test_paged_mqa_logits()