perf(jit_kernel/deepseek_v4): optimize paged_mqa_metadata (#25855)

Co-authored-by: SII-yangdian <yangdian@sii.edu.cn>
This commit is contained in:
SII-yangdian
2026-08-13 19:00:22 -07:00
committed by GitHub
co-authored by SII-yangdian
parent 704e512836
commit f2b2b567aa
4 changed files with 641 additions and 58 deletions
@@ -1,93 +1,314 @@
// paged_mqa_metadata: batch-size-adaptive dispatch.
//
// Replaces upstream's single-block kernel (grid=1, Phase-3 lane-serial
// advance, O(bs) dependent loads on the critical path) with three internal
// kernels dispatched by batch_size, all sharing the same Phase-1/2 prefix
// sum and a `num_sm + 1`-thread parallel upper_bound for Phase 3.
#include <sgl_kernel/tensor.h>
#include <sgl_kernel/utils.h>
#include <sgl_kernel/utils.cuh>
#include <sgl_kernel/warp.cuh>
#include <cub/block/block_scan.cuh>
#include <dlpack/dlpack.h>
#include <tvm/ffi/container/tensor.h>
#include <cstdint>
namespace sglang {
constexpr uint32_t kBlockSize = 1024;
constexpr uint32_t kSplitKV = 256; // const for both SM90 and SM100
constexpr uint32_t kTinyBlock = 256;
constexpr uint32_t kTinyMax = 64;
constexpr uint32_t kSmallBlock = 256;
constexpr uint32_t kSmallMax = 2048;
constexpr uint32_t kSmallItemsPerThread = 8;
static_assert(kSmallBlock * kSmallItemsPerThread == kSmallMax);
constexpr uint32_t kMBTileSize = 4096;
constexpr uint32_t kMBBlockSize = 1024;
constexpr uint32_t kMBItemsPerThread = 4;
constexpr uint32_t kKernelBThreads = 256;
static_assert(kMBBlockSize * kMBItemsPerThread == kMBTileSize);
struct MetadataParams {
/// NOTE: batch_size > 0
uint32_t batch_size;
uint32_t num_sm;
const uint32_t* __restrict__ context_lens;
uint32_t* __restrict__ schedule_metadata;
bool use_smem = true;
};
__global__ __launch_bounds__(kBlockSize, 1) //
void smxx_paged_mqa_logits_metadata(const MetadataParams params) {
using namespace device;
extern __shared__ uint32_t s_length[];
static constexpr auto kNumWarps = kBlockSize / kWarpThreads;
static_assert(kNumWarps == kWarpThreads);
// bs <= 64. Warp-0 inclusive scan, 256 B static smem.
__global__ __launch_bounds__(kTinyBlock, 1) //
void paged_mqa_metadata_tiny_kernel(const MetadataParams params) {
__shared__ uint32_t s_prefix[kTinyMax];
__shared__ uint32_t s_global_sum;
const auto tx = threadIdx.x;
const auto lane_id = tx % kWarpThreads;
const auto warp_id = tx / kWarpThreads;
const uint32_t tx = threadIdx.x;
const uint32_t bs = params.batch_size;
const uint32_t num_sm = params.num_sm;
__shared__ uint32_t s_warp_sum[kNumWarps];
uint32_t local_sum = 0;
for (uint32_t i = tx; i < params.batch_size; i += kBlockSize) {
const auto length = params.context_lens[i];
local_sum += (length + kSplitKV - 1) / kSplitKV;
if (params.use_smem) s_length[i] = length;
if (tx < 32) {
uint32_t running = 0;
#pragma unroll
for (uint32_t base = 0; base < kTinyMax; base += 32) {
const uint32_t idx = base + tx;
uint32_t v = 0;
if (idx < bs) {
const uint32_t length = params.context_lens[idx];
v = (length + kSplitKV - 1) >> 8;
}
#pragma unroll
for (int o = 1; o < 32; o <<= 1) {
uint32_t y = __shfl_up_sync(0xffffffff, v, o);
if (tx >= static_cast<uint32_t>(o)) v += y;
}
v += running;
if (idx < bs) s_prefix[idx] = v;
running = __shfl_sync(0xffffffff, v, 31);
}
if (tx == 0) s_global_sum = running;
}
s_warp_sum[warp_id] = warp::reduce_sum(local_sum);
__syncthreads();
const auto global_sum = warp::reduce_sum(s_warp_sum[lane_id]);
if (lane_id != 0) return;
const uint32_t global_sum = s_global_sum;
const uint32_t avg = global_sum / num_sm;
const uint32_t ret = global_sum % num_sm;
const uint32_t pivot = num_sm - ret;
const auto length_ptr = params.use_smem ? s_length : params.context_lens;
// Stride loop so num_sm > blockDim.x - 1 is fully written.
for (uint32_t i = tx; i <= num_sm; i += blockDim.x) {
// Match DeepGEMM's reversed remainder allocation: leading SMs get
// `avg` work and the final `ret` SMs get `avg + 1`. When global_sum is
// smaller than num_sm, empty SMs stay at the valid (q=0, offset=0)
// boundary instead of starting at q=batch_size.
const uint32_t target = i * avg + (i > pivot ? i - pivot : 0);
const auto avg = global_sum / params.num_sm;
const auto ret = global_sum % params.num_sm;
uint32_t q = 0;
uint32_t num_work = (length_ptr[0] + kSplitKV - 1) / kSplitKV;
uint32_t sum_work = num_work;
for (auto i = warp_id; i <= params.num_sm; i += kNumWarps) {
const auto target = i * avg + min(i, ret);
while (sum_work <= target) {
if (++q >= params.batch_size) break;
num_work = (length_ptr[q] + kSplitKV - 1) / kSplitKV;
sum_work += num_work;
uint32_t lo = 0;
uint32_t hi = bs;
while (lo < hi) {
const uint32_t mid = (lo + hi) >> 1;
if (s_prefix[mid] <= target)
lo = mid + 1;
else
hi = mid;
}
if (q >= params.batch_size) {
params.schedule_metadata[2 * i + 0] = params.batch_size;
const uint32_t q = lo;
if (q >= bs) {
params.schedule_metadata[2 * i + 0] = bs;
params.schedule_metadata[2 * i + 1] = 0;
} else {
// sum > target && (sum - length) <= target
const uint32_t prefix_prev = (q == 0) ? 0u : s_prefix[q - 1];
params.schedule_metadata[2 * i + 0] = q;
params.schedule_metadata[2 * i + 1] = target - (sum_work - num_work);
params.schedule_metadata[2 * i + 1] = target - prefix_prev;
}
}
}
template <auto* f, size_t kMaxDynamicSMEM>
void setup_kernel_smem_once(host::DebugInfo where = {}) {
[[maybe_unused]]
static const auto result = [] {
const auto fptr = std::bit_cast<const void*>(f);
return ::cudaFuncSetAttribute(fptr, ::cudaFuncAttributeMaxDynamicSharedMemorySize, kMaxDynamicSMEM);
}();
host::RuntimeDeviceCheck(result, where);
// 64 < bs <= 2048. CUB BlockScan, 8 KB static smem.
__global__ __launch_bounds__(kSmallBlock, 1) //
void paged_mqa_metadata_small_kernel(const MetadataParams params) {
using BlockScan = cub::BlockScan<uint32_t, kSmallBlock, cub::BLOCK_SCAN_WARP_SCANS>;
__shared__ uint32_t s_prefix[kSmallMax];
__shared__ typename BlockScan::TempStorage temp_storage;
__shared__ uint32_t s_global_sum;
const uint32_t tx = threadIdx.x;
const uint32_t bs = params.batch_size;
const uint32_t num_sm = params.num_sm;
uint32_t thread_items[kSmallItemsPerThread];
#pragma unroll
for (uint32_t k = 0; k < kSmallItemsPerThread; ++k) {
const uint32_t i = tx * kSmallItemsPerThread + k;
if (i < bs) {
const uint32_t length = params.context_lens[i];
thread_items[k] = (length + kSplitKV - 1) >> 8;
} else {
thread_items[k] = 0;
}
}
uint32_t block_aggregate;
BlockScan(temp_storage).InclusiveSum(thread_items, thread_items, block_aggregate);
if (tx == 0) s_global_sum = block_aggregate;
#pragma unroll
for (uint32_t k = 0; k < kSmallItemsPerThread; ++k) {
const uint32_t i = tx * kSmallItemsPerThread + k;
if (i < bs) s_prefix[i] = thread_items[k];
}
__syncthreads();
const uint32_t global_sum = s_global_sum;
const uint32_t avg = global_sum / num_sm;
const uint32_t ret = global_sum % num_sm;
const uint32_t pivot = num_sm - ret;
// Stride loop so num_sm > blockDim.x - 1 is fully written.
for (uint32_t i = tx; i <= num_sm; i += blockDim.x) {
const uint32_t target = i * avg + (i > pivot ? i - pivot : 0);
uint32_t lo = 0;
uint32_t hi = bs;
while (lo < hi) {
const uint32_t mid = (lo + hi) >> 1;
if (s_prefix[mid] <= target)
lo = mid + 1;
else
hi = mid;
}
const uint32_t q = lo;
if (q >= bs) {
params.schedule_metadata[2 * i + 0] = bs;
params.schedule_metadata[2 * i + 1] = 0;
} else {
const uint32_t prefix_prev = (q == 0) ? 0u : s_prefix[q - 1];
params.schedule_metadata[2 * i + 0] = q;
params.schedule_metadata[2 * i + 1] = target - prefix_prev;
}
}
}
// bs > 2048, Phase 1: ceil(bs / kMBTileSize) blocks each emit an in-tile
// inclusive prefix into scratch_prefix and a per-tile sum into tile_sums.
__global__ __launch_bounds__(kMBBlockSize, 1) //
void phase1_tile_scan_kernel(
const MetadataParams params, uint32_t* __restrict__ scratch_prefix, uint32_t* __restrict__ tile_sums) {
using TileBlockScan = cub::BlockScan<uint32_t, kMBBlockSize, cub::BLOCK_SCAN_WARP_SCANS>;
__shared__ typename TileBlockScan::TempStorage temp_storage;
const uint32_t bs = params.batch_size;
const uint32_t tile_idx = blockIdx.x;
const uint32_t tile_base = tile_idx * kMBTileSize;
const uint32_t tx = threadIdx.x;
uint32_t thread_items[kMBItemsPerThread];
#pragma unroll
for (uint32_t k = 0; k < kMBItemsPerThread; ++k) {
const uint32_t i = tile_base + tx * kMBItemsPerThread + k;
if (i < bs) {
const uint32_t length = params.context_lens[i];
thread_items[k] = (length + kSplitKV - 1) >> 8;
} else {
thread_items[k] = 0;
}
}
uint32_t block_aggregate;
TileBlockScan(temp_storage).InclusiveSum(thread_items, thread_items, block_aggregate);
#pragma unroll
for (uint32_t k = 0; k < kMBItemsPerThread; ++k) {
const uint32_t i = tile_base + tx * kMBItemsPerThread + k;
if (i < bs) scratch_prefix[i] = thread_items[k];
}
if (tx == 0) tile_sums[tile_idx] = block_aggregate;
}
// bs > 2048, Phase 2/3: one block, kKernelBThreads threads. Warp-0 scans
// tile_sums into s_tile_prefix; then num_sm+1 threads do tile-level
// upper_bound + within-tile upper_bound to recover (batch_idx, offset).
__global__ __launch_bounds__(kKernelBThreads, 1) //
void schedule_from_tiles_kernel(
const MetadataParams params,
const uint32_t* __restrict__ scratch_prefix,
const uint32_t* __restrict__ tile_sums,
uint32_t num_tiles) {
extern __shared__ uint32_t s_tile_prefix[];
__shared__ uint32_t s_global_sum;
const uint32_t tx = threadIdx.x;
const uint32_t bs = params.batch_size;
const uint32_t num_sm = params.num_sm;
if (tx < 32) {
uint32_t running = 0;
for (uint32_t base = 0; base < num_tiles; base += 32) {
const uint32_t idx = base + tx;
uint32_t v = (idx < num_tiles) ? tile_sums[idx] : 0;
#pragma unroll
for (int o = 1; o < 32; o <<= 1) {
uint32_t y = __shfl_up_sync(0xffffffff, v, o);
if (tx >= static_cast<uint32_t>(o)) v += y;
}
v += running;
if (idx < num_tiles) s_tile_prefix[idx] = v;
running = __shfl_sync(0xffffffff, v, 31);
}
if (tx == 0) s_global_sum = running;
}
__syncthreads();
const uint32_t global_sum = s_global_sum;
const uint32_t avg = global_sum / num_sm;
const uint32_t ret = global_sum % num_sm;
const uint32_t pivot = num_sm - ret;
// Stride loop so num_sm > blockDim.x - 1 is fully written. `continue`
// replaces the original early-out `return` so later strided targets
// still get processed.
for (uint32_t i = tx; i <= num_sm; i += blockDim.x) {
const uint32_t target = i * avg + (i > pivot ? i - pivot : 0);
uint32_t t_lo = 0;
uint32_t t_hi = num_tiles;
while (t_lo < t_hi) {
const uint32_t mid = (t_lo + t_hi) >> 1;
if (s_tile_prefix[mid] <= target)
t_lo = mid + 1;
else
t_hi = mid;
}
const uint32_t tile = t_lo;
if (tile >= num_tiles) {
params.schedule_metadata[2 * i + 0] = bs;
params.schedule_metadata[2 * i + 1] = 0;
continue;
}
const uint32_t tile_offset = (tile == 0) ? 0u : s_tile_prefix[tile - 1];
const uint32_t tile_start = tile * kMBTileSize;
uint32_t tile_end = tile_start + kMBTileSize;
if (tile_end > bs) tile_end = bs;
const uint32_t local_target = target - tile_offset;
uint32_t lo = tile_start;
uint32_t hi = tile_end;
while (lo < hi) {
const uint32_t mid = (lo + hi) >> 1;
if (scratch_prefix[mid] <= local_target)
lo = mid + 1;
else
hi = mid;
}
const uint32_t q = lo;
if (q >= bs) {
params.schedule_metadata[2 * i + 0] = bs;
params.schedule_metadata[2 * i + 1] = 0;
} else {
const uint32_t prefix_prev = (q == tile_start) ? tile_offset : (scratch_prefix[q - 1] + tile_offset);
params.schedule_metadata[2 * i + 0] = q;
params.schedule_metadata[2 * i + 1] = target - prefix_prev;
}
}
}
struct IndexerMetadataKernel {
static constexpr auto kMaxBatchSizeInSmem = 16384 * 2; // 128 KB smeme
static void run(tvm::ffi::TensorView seq_lens, tvm::ffi::TensorView metadata) {
static void run(tvm::ffi::TensorView seq_lens, tvm::ffi::TensorView metadata, tvm::ffi::TensorView workspace) {
using namespace host;
auto N = SymbolicSize{"batch_size"};
auto M = SymbolicSize{"num_sm"};
auto W = SymbolicSize{"workspace"};
auto device = SymbolicDevice{};
device.set_options<kDLCUDA>();
TensorMatcher({N}) //
@@ -98,21 +319,40 @@ struct IndexerMetadataKernel {
.with_dtype<int32_t>()
.with_device(device)
.verify(metadata);
TensorMatcher({W}) //
.with_dtype<int32_t>()
.with_device(device)
.verify(workspace);
const auto batch_size = static_cast<uint32_t>(N.unwrap());
const auto num_sm = static_cast<uint32_t>(M.unwrap()) - 1;
RuntimeCheck(num_sm <= 1024);
const auto use_smem = batch_size <= kMaxBatchSizeInSmem;
RuntimeCheck(num_sm >= 1 && num_sm <= 1024);
const auto params = MetadataParams{
.batch_size = batch_size,
.num_sm = num_sm,
.context_lens = static_cast<uint32_t*>(seq_lens.data_ptr()),
.schedule_metadata = static_cast<uint32_t*>(metadata.data_ptr()),
.use_smem = use_smem,
};
constexpr auto kernel = smxx_paged_mqa_logits_metadata;
setup_kernel_smem_once<kernel, (kMaxBatchSizeInSmem + 1) * sizeof(uint32_t)>();
const auto smem = use_smem ? (batch_size + 1) * sizeof(uint32_t) : 0;
LaunchKernel(1, kBlockSize, device.unwrap(), smem)(kernel, params);
const auto dl_device = device.unwrap();
if (batch_size <= kTinyMax) {
LaunchKernel(1, kTinyBlock, dl_device)(paged_mqa_metadata_tiny_kernel, params);
} else if (batch_size <= kSmallMax) {
LaunchKernel(1, kSmallBlock, dl_device)(paged_mqa_metadata_small_kernel, params);
} else {
const auto num_tiles = (batch_size + kMBTileSize - 1) / kMBTileSize;
const auto required = static_cast<int64_t>(batch_size) + num_tiles;
RuntimeCheck(static_cast<int64_t>(W.unwrap()) >= required, "workspace too small for multi-block path");
auto* scratch_prefix = static_cast<uint32_t*>(workspace.data_ptr());
auto* tile_sums = scratch_prefix + batch_size;
LaunchKernel(num_tiles, kMBBlockSize, dl_device)(phase1_tile_scan_kernel, params, scratch_prefix, tile_sums);
const auto kb_smem_bytes = static_cast<size_t>(num_tiles) * sizeof(uint32_t);
LaunchKernel(1, kKernelBThreads, dl_device, kb_smem_bytes)(
schedule_from_tiles_kernel, params, scratch_prefix, tile_sums, num_tiles);
}
}
};
@@ -45,9 +45,18 @@ def _jit_fused_store_module(
def get_paged_mqa_logits_metadata(seq_lens: torch.Tensor, page_size: int, num_sm: int):
assert page_size == 64
seq_lens = seq_lens.view(-1).to(torch.int32)
bs = int(seq_lens.shape[0])
metadata = seq_lens.new_empty(num_sm + 1, 2)
# Workspace for the multi-block path; kMBTileSize must match the .cuh.
if bs > 2048:
kMBTileSize = 4096
workspace = seq_lens.new_empty(
bs + (bs + kMBTileSize - 1) // kMBTileSize, dtype=torch.int32
)
else:
workspace = seq_lens.new_empty(0, dtype=torch.int32)
module = _jit_metadata_module()
module.run(seq_lens, metadata)
module.run(seq_lens, metadata, workspace)
return metadata