diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/block_amax.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/block_amax.cuh new file mode 100644 index 000000000..330073800 --- /dev/null +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/block_amax.cuh @@ -0,0 +1,155 @@ +#pragma once + +#include +#include + +#include +#include + +#include +#include + +#include +#include + +namespace sglang { + +/// Level-one keys of the two-level indexer: the max of each kBlockTokens-score +/// block, the row's newest block forced to +inf. Contract: BlockAmaxKernel. +struct BlockAmaxConfig { + using DType = float; + static constexpr uint32_t kBlockTokens = 8; // scores per key + static constexpr uint32_t kBlockSize = 512; + static constexpr uint32_t kNumItems = 2; // keys per thread + static constexpr uint32_t kOccupancy = 4; + static constexpr uint32_t kKeysPerCTA = kBlockSize * kNumItems; + // One block is 32 B: a single load on Blackwell, two 16 B loads before it. + static constexpr uint32_t kVecSize = device::kMaxVecBytes / sizeof(DType); + static constexpr uint32_t kVecsPerBlock = kBlockTokens / kVecSize; + static_assert(kVecsPerBlock * kVecSize == kBlockTokens); + using vec_t = device::AlignedVector; +}; + +struct BlockAmaxParams { + const BlockAmaxConfig::DType* __restrict__ scores; + BlockAmaxConfig::DType* __restrict__ amax_scores; + const int32_t* __restrict__ seq_len; + int64_t stride_scores; // in elements + int64_t stride_amax_scores; // in elements + uint32_t topk; // rows with <= topk blocks are skipped, 0 = never skip +}; + +/// grid = (rows, ceil(max_keys / kKeysPerCTA)); a CTA owns kKeysPerCTA consecutive +/// keys of one row, a thread kNumItems keys kBlockSize apart (coalesced loads). +template +__global__ __launch_bounds__(BlockAmaxConfig::kBlockSize, BlockAmaxConfig::kOccupancy) // + void amax8_varlen_kernel(const __grid_constant__ BlockAmaxParams params) { + using namespace device; + using C = BlockAmaxConfig; + using T = typename C::DType; + using vec_t = typename C::vec_t; + const auto bx = blockIdx.x; + const auto by = blockIdx.y; + const auto tx = threadIdx.x; + + PDLWaitPrimary(); // seq_len and scores are the previous kernels' outputs + const auto seq_len = static_cast(params.seq_len[bx]); + const auto num_keys = (seq_len + C::kBlockTokens - 1) / C::kBlockTokens; + const auto first_key = by * C::kKeysPerCTA; + if (num_keys <= params.topk || first_key >= num_keys) { + return PDLTriggerSecondary(); + } + const auto* __restrict__ in = params.scores + bx * params.stride_scores; + auto* __restrict__ out = params.amax_scores + bx * params.stride_amax_scores; + + vec_t vec[C::kNumItems][C::kVecsPerBlock]; +#pragma unroll + for (uint32_t i = 0; i < C::kNumItems; ++i) { + const auto idx = first_key + tx + i * C::kBlockSize; + if (idx < num_keys) { +#pragma unroll + for (uint32_t v = 0; v < C::kVecsPerBlock; ++v) { + vec[i][v].load(in, idx * C::kVecsPerBlock + v); + } + } + } + // The dependent grid may start its prologue now; its griddepcontrol.wait still + // covers every store below (it waits for this grid to complete). + PDLTriggerSecondary(); + +#pragma unroll + for (uint32_t i = 0; i < C::kNumItems; ++i) { + const auto idx = first_key + tx + i * C::kBlockSize; + if (idx < num_keys) { + T key = vec[i][0][0]; +#pragma unroll + for (uint32_t v = 0; v < C::kVecsPerBlock; ++v) { +#pragma unroll + for (uint32_t j = 0; j < C::kVecSize; ++j) { + key = fmaxf(key, vec[i][v][j]); // a NaN score is ignored, torch.amax would propagate it + } + } + out[idx] = idx + 1 == num_keys ? std::numeric_limits::infinity() : key; + } + } +} + +/// Host entry: `amax_scores[b, i] = max(scores[b, 8 i : 8 i + 8])` for +/// `i < ceil(seq_len[b] / 8)`, the last of them +inf; rows with at most `topk` +/// blocks untouched. `scores` rows must stay 32 B aligned (stride % 8 == 0). +/// The grid covers `amax_scores`' width, so the caller sizes it for the longest +/// row: `seq_len[b] <= 8 * amax_scores.shape[1]` for every row (not checked). +template +struct BlockAmaxKernel { + static void amax8_varlen( + const tvm::ffi::TensorView scores, + const tvm::ffi::TensorView seq_lens, + const tvm::ffi::TensorView amax_scores, + const uint32_t topk) { + using namespace host; + using C = BlockAmaxConfig; + auto B = SymbolicSize{"batch_size"}; + auto L = SymbolicSize{"max_seq_len"}; + auto S = SymbolicSize{"stride_scores"}; + auto K = SymbolicSize{"max_keys"}; + auto O = SymbolicSize{"stride_amax_scores"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + TensorMatcher({B, L}) // scores + .with_strides({S, 1}) + .with_dtype() + .with_device(device_) + .verify(scores); + TensorMatcher({B}) // seq_lens + .with_dtype() + .with_device(device_) + .verify(seq_lens); + TensorMatcher({B, K}) // amax_scores + .with_strides({O, 1}) + .with_dtype() + .with_device(device_) + .verify(amax_scores); + RuntimeCheck(S.unwrap() % C::kBlockTokens == 0, "stride_scores must keep every block 32 B aligned"); + RuntimeCheck( + reinterpret_cast(scores.data_ptr()) % (C::kBlockTokens * sizeof(typename C::DType)) == 0, + "scores must be 32 B aligned"); + RuntimeCheck(K.unwrap() > 0, "amax_scores must hold at least one key per row"); + const auto max_keys = K.unwrap(); // ceil(longest row / 8), sized by the caller + const auto params = BlockAmaxParams{ + .scores = static_cast(scores.data_ptr()), + .amax_scores = static_cast(amax_scores.data_ptr()), + .seq_len = static_cast(seq_lens.data_ptr()), + .stride_scores = S.unwrap(), + .stride_amax_scores = O.unwrap(), + .topk = topk, + }; + const auto grid = dim3( + static_cast(B.unwrap()), + static_cast(div_ceil(max_keys, static_cast(C::kKeysPerCTA)))); + LaunchKernel(grid, C::kBlockSize, device_.unwrap()) + .config({.use_pdl = kPDL}) + .launch(amax8_varlen_kernel, params); + } +}; + +} // namespace sglang diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/candidate_block_table.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/candidate_block_table.cuh new file mode 100644 index 000000000..3cffe255c --- /dev/null +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/candidate_block_table.cuh @@ -0,0 +1,218 @@ +#pragma once + +#include +#include + +#include +#include +#include + +#include +#include + +#include +#include +#include + +namespace sglang { + +/// Finalises the sparse indexer's block table: the top-k block ids a row +/// selected (any order, -1 padded) become, in place, the same ids ascending with +/// INT32_MAX past the row's count, plus each block as a pool slot / 8 +/// (`page_table[b, id / bpp] * bpp + id % bpp`, `bpp` blocks per index page). A +/// row with at most `topk` blocks gets the identity table without reading its +/// input. +/// +/// Counting sort over a per-row bitmap (one bit per block, 16 KiB for a 1M-token +/// row): single-bit words are emitted by their owner, denser words go to a +/// block-wide queue the warps drain one lane per bit. +struct CandidateBlockTableConfig { + static constexpr uint32_t kBlockSize = 1024; + static constexpr uint32_t kOccupancy = 2; + static constexpr uint32_t kNumWarps = kBlockSize / device::kWarpThreads; + static constexpr uint32_t kBlockTokens = 8; + static constexpr uint32_t kMaxSeqLen = 128 * 1024; // blocks: 1M tokens / kBlockTokens + static constexpr uint32_t kMaxTopK = 2048; + static constexpr uint32_t kWordsPerThread = kMaxSeqLen / 32 / kBlockSize; + static_assert(kWordsPerThread == 4 && kNumWarps == device::kWarpThreads); + static constexpr int32_t kPad = std::numeric_limits::max(); + using word_vec_t = device::AlignedVector; + struct WriteItem { + uint32_t start; // rank of the word's first bit | word index << 16 + uint32_t bits; + }; + struct Smem { + uint32_t queue_size; + uint32_t warp_sum[kNumWarps]; + union { + alignas(16) uint32_t bitmap[kMaxSeqLen / 32]; + WriteItem write_queue[kMaxTopK]; // a queued word holds >= 2 of the topk bits + }; + }; +}; + +struct CandidateBlockTableParams { + const uint32_t* __restrict__ seq_len; // [rows] tokens + const int32_t* __restrict__ page_table; // [rows, pages] index-pool pages + int32_t* __restrict__ indices; // [rows, topk] blocks, -1 padded in, ascending + kPad out + int32_t* __restrict__ out_pages; // [rows, topk] the same blocks as pool slots / 8 + int64_t page_table_stride; + int64_t indices_stride; + int64_t out_pages_stride; + uint32_t topk; + uint32_t page_bits; // log2(page_size / kBlockTokens) +}; + +/// One CTA per row. +template +__global__ __launch_bounds__(CandidateBlockTableConfig::kBlockSize, CandidateBlockTableConfig::kOccupancy) // + void sort_128k_transform(const __grid_constant__ CandidateBlockTableParams params) { + using namespace device; + using C = CandidateBlockTableConfig; + __shared__ C::Smem smem; + const auto bx = blockIdx.x; + const auto tx = threadIdx.x; + const auto warp_id = tx / kWarpThreads; + const auto lane_id = tx % kWarpThreads; + const auto lanemask_lt = (1u << lane_id) - 1u; + + PDLWaitPrimary(); // indices is the block top-k's output + const auto seq_len = params.seq_len[bx]; + const auto nblocks = (seq_len + C::kBlockTokens - 1) / C::kBlockTokens; + const auto* __restrict__ table = params.page_table + bx * params.page_table_stride; + auto* __restrict__ indices = params.indices + bx * params.indices_stride; + auto* __restrict__ pages = params.out_pages + bx * params.out_pages_stride; + const auto bpp_mask = (1u << params.page_bits) - 1u; + const auto emit = [&](uint32_t rank, uint32_t id) { + indices[rank] = static_cast(id); + pages[rank] = (table[id >> params.page_bits] << params.page_bits) | static_cast(id & bpp_mask); + }; + const auto pad = [&](uint32_t rank) { + indices[rank] = C::kPad; + pages[rank] = C::kPad; + }; + + if (nblocks <= params.topk) { // every block is selected: the identity table + for (uint32_t t = tx; t < params.topk; t += C::kBlockSize) { + if (t < nblocks) { + emit(t, t); + } else { + pad(t); + } + } + return PDLTriggerSecondary(); + } + + // 1. the selected blocks as a bitmap + C::word_vec_t words; + words.fill(0u); + words.store(smem.bitmap, tx); + if (tx == 0) smem.queue_size = 0; + __syncthreads(); + for (uint32_t t = tx; t < params.topk; t += C::kBlockSize) { + const auto id = indices[t]; + if (id >= 0) atomicOr(&smem.bitmap[id >> 5], 1u << (id & 31)); + } + __syncthreads(); + + // 2. rank of every word's first bit: block-wide exclusive scan of the popcounts + words.load(smem.bitmap, tx); + uint32_t count[C::kWordsPerThread]; + uint32_t local = 0; +#pragma unroll + for (uint32_t j = 0; j < C::kWordsPerThread; ++j) { + count[j] = __popc(words[j]); + local += count[j]; + } + const auto warp_inc = warp::inclusive_sum(lane_id, local); + if (lane_id == kWarpThreads - 1) smem.warp_sum[warp_id] = warp_inc; + __syncthreads(); // also: every thread holds its words, the bitmap may become the queue + const auto peer_sum = smem.warp_sum[lane_id]; + const auto warp_prefix = warp::reduce_sum(lane_id < warp_id ? peer_sum : 0u); + const auto total = warp::reduce_sum(peer_sum); + uint32_t base = warp_prefix + warp_inc - local; + PDLTriggerSecondary(); + + // 3. single bits by their owner, denser words queued for the warps +#pragma unroll + for (uint32_t j = 0; j < C::kWordsPerThread; ++j) { + const auto word_idx = tx * C::kWordsPerThread + j; + if (count[j] == 1) { + emit(base, word_idx * 32 + __ffs(words[j]) - 1); + } else if (count[j] >= 2) { + const auto slot = atomicAdd(&smem.queue_size, 1u); + smem.write_queue[slot] = {base | (word_idx << 16), words[j]}; + } + base += count[j]; + } + for (uint32_t t = total + tx; t < params.topk; t += C::kBlockSize) { + pad(t); + } + __syncthreads(); + + // 4. drain the queue: one word per warp step, one lane per bit + const auto queue_size = smem.queue_size; + for (uint32_t q = warp_id; q < queue_size; q += C::kNumWarps) { + const auto item = smem.write_queue[q]; + if ((item.bits >> lane_id) & 1u) { + emit((item.start & 0xFFFFu) + __popc(item.bits & lanemask_lt), (item.start >> 16) * 32 + lane_id); + } + } +} + +/// Host entry: `indices` is rewritten in place; `page_size` is the index pool's, +/// a power of two >= 8, and the row's page table must cover its length. +template +struct CandidateBlockTableKernel { + static void transform( + const tvm::ffi::TensorView indices, + const tvm::ffi::TensorView seq_lens, + const tvm::ffi::TensorView page_table, + const tvm::ffi::TensorView out_pages, + const uint32_t page_size) { + launch>(indices, seq_lens, page_table, out_pages, page_size); + } + + private: + template + static void launch( + const tvm::ffi::TensorView indices, + const tvm::ffi::TensorView seq_lens, + const tvm::ffi::TensorView page_table, + const tvm::ffi::TensorView out_pages, + const uint32_t page_size) { + using namespace host; + using C = CandidateBlockTableConfig; + auto B = SymbolicSize{"batch_size"}; + auto K = SymbolicSize{"topk_blocks"}; + auto Si = SymbolicSize{"indices_stride"}; + auto Sp = SymbolicSize{"out_pages_stride"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + TensorMatcher({B, K}).with_strides({Si, 1}).with_dtype().with_device(device_).verify(indices); + TensorMatcher({B}).with_dtype().with_device(device_).verify(seq_lens); + TensorMatcher({B, -1}).with_strides({-1, 1}).with_dtype().with_device(device_).verify(page_table); + TensorMatcher({B, K}).with_strides({Sp, 1}).with_dtype().with_device(device_).verify(out_pages); + RuntimeCheck( + std::has_single_bit(page_size) && page_size >= C::kBlockTokens, + "page_size must be a power of two of at least 8"); + const auto topk = static_cast(K.unwrap()); + RuntimeCheck(topk > 0 && topk <= C::kMaxTopK, "topk_blocks must be in (0, kMaxTopK]"); + const auto params = CandidateBlockTableParams{ + .seq_len = static_cast(seq_lens.data_ptr()), + .page_table = static_cast(page_table.data_ptr()), + .indices = static_cast(indices.data_ptr()), + .out_pages = static_cast(out_pages.data_ptr()), + .page_table_stride = page_table.stride(0), + .indices_stride = Si.unwrap(), + .out_pages_stride = Sp.unwrap(), + .topk = topk, + .page_bits = static_cast(std::countr_zero(page_size / C::kBlockTokens)), + }; + LaunchKernel(static_cast(B.unwrap()), C::kBlockSize, device_.unwrap()) + .config({.use_pdl = kPDL}) + .launch(kKernel, params); + } +}; + +} // namespace sglang diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/topk_bf16_small.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/topk_bf16_small.cuh new file mode 100644 index 000000000..37daad49d --- /dev/null +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/topk_bf16_small.cuh @@ -0,0 +1,440 @@ +/** + * \brief DeepSeek-V4.1's bf16 top-k kernel for short rows (<= 16384 scores) + * Adapted from https://github.com/deepseek-ai/DeepSelect + * Plain SIMT (no tensor cores or clusters), tuned for 16384-wide rows with k = 512. + */ +#pragma once + +#include +#include + +#include +#include +#include + +#include +#include + +#include +#include + +namespace sglang { + +/** + * \brief bf16 top-k of one row that fits in registers (rows of at most 16384 scores: the + * DeepSeek-V4.1 sparse indexer's consumer rows), fused with a page-table transform. + * + * One CTA of 512 threads per row, up to 32 scores per thread held in registers. Two radix + * passes over the raw bf16 bytes locate the k-th largest value exactly, a census places every + * element relative to it, and the selected indices are staged in shared memory before one + * coalesced, page-transformed store. This is DeepSelect's init-window select. + * + * \note The value order used everywhere is the "distorted" order of the raw bf16 bits + * (`x ^ (x < 0 ? 0xFFFF : 0x8000)`, negatives below positives, -0 below +0). The + * histograms are indexed by the *raw* byte instead, and the pivot search undoes the + * permutation once per lane, so no element pays the distortion. + * \note NaN scores are not selected: the ordered compares never match them, so a row with n + * positive NaNs yields its top (k - n) real scores and -1 in the remaining slots (a + * negative NaN orders below -inf and is simply never picked). + */ +struct TopKBF16Config { + static constexpr uint32_t kBlockSize = 512; + static constexpr uint32_t kNumWarps = kBlockSize / device::kWarpThreads; + static constexpr uint32_t kOccupancy = 3; + static constexpr uint32_t kVecSize = 8; + static constexpr uint32_t kMaxVecs = 4; + static constexpr uint32_t kElemsPerThread = kVecSize * kMaxVecs; + static constexpr uint32_t kMaxSeqLen = kBlockSize * kElemsPerThread; + static constexpr uint32_t kMaxTopK = 2048; + static constexpr uint32_t kNumBins = 256; + static constexpr uint32_t kSinkBin = kNumBins; // LSB pass sends out-of-bucket elements here + /// NOTE: in the MSB row the negative half lives 16 words further up. Raw bytes 128 apart share + /// a bank, so without it +x and -x with the same exponent (the common case for centered data) + /// collide on every histogram update; measured as half of all atomic wavefronts. + static constexpr uint32_t kNegShift = 16; + static constexpr uint32_t kHistStride = kNumBins + kNegShift + 4; // keeps both rows 16 B aligned + /// NOTE: a negative NaN. In the distorted order it sits below -inf, and every ordered bf16 + /// comparison against it is false, so padding is never counted nor selected. + static constexpr uint32_t kPadElem = 0xFFFFu; + static constexpr uint32_t kNegZeroBits = 0x8000u; + using vec_t = device::AlignedVector; + static_assert(kMaxSeqLen == 16384 && kMaxSeqLen <= 0xFFFF); // the census counters pack in 16 bits + // one census bit per element of the slice, in a uint32_t + static_assert(kElemsPerThread == 32); + + struct Smem { + uint32_t count_gt_eq; // packed (gt << 16 | eq), the block-wide census prefix + uint32_t pivot_bin; + uint32_t pivot_remain; + union { + alignas(16) uint32_t histogram[2][kHistStride]; + alignas(16) uint32_t stage[kMaxTopK]; + }; + }; +}; + +struct TopKBF16Params { + const bf16_t* __restrict__ scores; + const int32_t* __restrict__ seq_lens; + const int32_t* __restrict__ page_table; + int32_t* __restrict__ page_indices; + int64_t score_stride; + int64_t page_table_stride; + int64_t page_indices_stride; + uint32_t topk; + uint32_t page_bits; +}; + +SGL_DEVICE uint32_t get_ptx_lane_id() { + uint32_t lane_id; + asm volatile("mov.u32 %0, %%laneid;" : "=r"(lane_id)); + return lane_id; +} + +/// \brief Exclusive suffix scan: lane `L` gets the sum over lanes `> L`. +SGL_DEVICE uint32_t warp_exclusive_suffix_sum(uint32_t x, uint32_t lane_id) { + uint32_t inc = x; +#pragma unroll + for (uint32_t offset = 1; offset < device::kWarpThreads; offset <<= 1) { + const auto t = __shfl_down_sync(device::kFullMask, inc, offset); + if (lane_id + offset < device::kWarpThreads) inc += t; + } + return inc - x; +} + +template +SGL_DEVICE To bitcast(const From& f) { + static_assert(sizeof(From) == sizeof(To)); + return reinterpret_cast(f); +} + +struct TopKBF16Pivot { + uint32_t bin; // in distorted (value-ascending) order, [0, 256) + uint32_t remain; // how many elements of `bin` still have to be taken +}; + +/// Locate the bin holding the k-th largest element in a 256-bin histogram indexed by a raw +/// byte; one warp, exactly one lane writes `smem.pivot_*`. `msb_mode`: lanes < 16 cover raw +/// 0xFF..0x80 (negatives, reversed), lanes >= 16 raw 0x00..0x7F. `negative` (LSB mode only): +/// the pivot bucket is negative, so the whole byte is reversed. +SGL_DEVICE void topk_bf16_find_pivot_warp( + const uint32_t* hist, uint32_t k, bool msb_mode, bool negative, uint32_t lane_id, TopKBF16Config::Smem& smem) { + using C = TopKBF16Config; + // lane L owns distorted bins [8L, 8L + 8) + const bool reverse = msb_mode ? lane_id < 16 : negative; + uint32_t raw_base = reverse ? 0xF8 - 8 * lane_id : 8 * lane_id - (msb_mode ? 0x80 : 0); + if (msb_mode && reverse) raw_base += C::kNegShift; + device::AlignedVector lo, hi; + lo.load(hist + raw_base); + hi.load(hist + raw_base + 4); + uint32_t count[8]; +#pragma unroll + for (uint32_t i = 0; i < 8; ++i) { + const auto fwd = i < 4 ? lo[i] : hi[i - 4]; + const auto rev = i < 4 ? hi[3 - i] : lo[7 - i]; + count[i] = reverse ? rev : fwd; + } + uint32_t local = 0; +#pragma unroll + for (uint32_t i = 0; i < 8; ++i) { + local += count[i]; + } + // suffix[j] = number of elements in bins >= 8L + j + uint32_t suffix[9]; + suffix[8] = warp_exclusive_suffix_sum(local, lane_id); +#pragma unroll + for (int32_t j = 7; j >= 0; --j) { + suffix[j] = suffix[j + 1] + count[j]; + } + // exactly one lane satisfies suffix[8] < k <= suffix[0]; inside it, the pivot is the largest + // offset j with suffix[j] >= k + const bool found = suffix[8] < k && k <= suffix[0]; + uint32_t offset = 0; + uint32_t next = suffix[1]; +#pragma unroll + for (uint32_t j = 1; j < 8; ++j) { + if (suffix[j] >= k) { + offset = j; + next = suffix[j + 1]; + } + } + if (found) { + smem.pivot_bin = 8 * lane_id + offset; + smem.pivot_remain = k - next; + } +} + +/// One hit bit per element of an 8-wide vector, from the per-pair 16-bit masks of `__hgt2_mask` and friends. +SGL_DEVICE uint32_t topk_bf16_pack_hits(const uint32_t (&m)[4]) { + // one flag byte per element (0xFF / 0x00), then signed dot products turn them into bits + const auto lo = __byte_perm(m[0], m[1], 0x7531); + const auto hi = __byte_perm(m[2], m[3], 0x7531); + const auto nib = __dp4a(static_cast(lo), static_cast(0xF8FCFEFFu), 0); // -1,-2,-4,-8 + return __dp4a(static_cast(hi), static_cast(0x80C0E0F0u), nib); // -16..-128 +} + +template +__global__ __launch_bounds__(TopKBF16Config::kBlockSize, TopKBF16Config::kOccupancy) // + void topk_bf16_small_kernel(const __grid_constant__ TopKBF16Params params) { + using namespace device; + using C = TopKBF16Config; + using vec_t = C::vec_t; + __shared__ C::Smem smem; + + const auto bx = blockIdx.x; + const auto tx = threadIdx.x; + const auto lane_id = get_ptx_lane_id(); + const auto warp_id = tx / kWarpThreads; + const auto topk = params.topk; + // a selected index i maps through this row's table to slot + // table[i >> page_bits] << page_bits | (i & mask); -1 past what the row has + const auto* __restrict__ table = params.page_table + bx * params.page_table_stride; + auto* __restrict__ out = params.page_indices + bx * params.page_indices_stride; + const auto page_bits = params.page_bits; + const auto page_mask = (1u << page_bits) - 1; + const auto transform = [&](uint32_t idx) -> int32_t { + return (table[idx >> page_bits] << page_bits) | static_cast(idx & page_mask); + }; + + { + using zero_vec_t = AlignedVector; + static_assert(sizeof(smem.histogram) % sizeof(zero_vec_t) == 0); + constexpr uint32_t kZeroVecs = sizeof(smem.histogram) / sizeof(zero_vec_t); + zero_vec_t zeros; + zeros.fill(0); +#pragma unroll + for (uint32_t idx = tx; idx < kZeroVecs; idx += C::kBlockSize) { + zeros.store(smem.histogram, idx); + } + if (tx == 0) smem.count_gt_eq = 0; + } + + // NOTE: we prefetch metadata like seq_len + const auto seq_len = static_cast(params.seq_lens[bx]); + const auto* __restrict__ scores_row = params.scores + bx * params.score_stride; + if (seq_len <= topk) { // every element is selected, -1 past the row + PDLWaitPrimary(); + for (uint32_t t = tx; t < topk; t += C::kBlockSize) { + out[t] = t < seq_len ? transform(t) : -1; + } + return PDLTriggerSecondary(); + } + PDLWaitPrimary(); + + // Contiguous slices of whole vectors, balanced so short rows still spread over the block. + // Only the last vector of a row can be partial; it is padded with NaNs (see kPadElem). + const uint32_t num_vecs = div_ceil(seq_len, C::kVecSize); + const uint32_t num_full = seq_len / C::kVecSize; + const uint32_t vecs_per_thread = num_vecs / C::kBlockSize; + const uint32_t vecs_rem = num_vecs % C::kBlockSize; + const uint32_t vec_start = tx * vecs_per_thread + min(tx, vecs_rem); + const uint32_t num_my = vecs_per_thread + (tx < vecs_rem ? 1 : 0); + vec_t vecs[C::kMaxVecs]; +#pragma unroll + for (uint32_t i = 0; i < C::kMaxVecs; ++i) { + if (i >= num_my) break; + const auto v = vec_start + i; + if (v < num_full) { + vecs[i].load(scores_row, v); + } else { + const auto* ptr = reinterpret_cast(scores_row) + v * C::kVecSize; + const auto n = seq_len - v * C::kVecSize; // in [1, kVecSize) +#pragma unroll + for (uint32_t j = 0; j < C::kVecSize / 2; ++j) { + vecs[i][j].x = bitcast(2 * j + 0 < n ? ptr[2 * j + 0] : static_cast(C::kPadElem)); + vecs[i][j].y = bitcast(2 * j + 1 < n ? ptr[2 * j + 1] : static_cast(C::kPadElem)); + } + } + } + + __syncthreads(); + + // Pass 1: histogram of the raw high byte (sign + 7 exponent bits) + const auto hist_msb = smem.histogram[0]; +#pragma unroll + for (uint32_t i = 0; i < C::kMaxVecs; ++i) { + if (i >= num_my) break; +#pragma unroll + for (uint32_t j = 0; j < C::kVecSize / 2; ++j) { + const auto raw = bitcast(vecs[i][j]); + /// NOTE: spelled as byte extraction so the address is one PRMT + one LEA per element + const auto b0 = __byte_perm(raw, 0, 0x4441); + const auto b1 = __byte_perm(raw, 0, 0x4443); + atomicAdd(hist_msb + b0 + (b0 >> 7) * C::kNegShift, 1); + atomicAdd(hist_msb + b1 + (b1 >> 7) * C::kNegShift, 1); + } + } + __syncthreads(); + + const auto pivot_of = [&](const uint32_t* hist, uint32_t k, bool msb_mode, bool neg) -> TopKBF16Pivot { + if (warp_id == 0) topk_bf16_find_pivot_warp(hist, k, msb_mode, neg, lane_id, smem); + __syncthreads(); + return {smem.pivot_bin, smem.pivot_remain}; + }; + const auto msb = pivot_of(hist_msb, topk, true, false); + const bool negative = msb.bin < 0x80; + const auto pivot_hi = negative ? 0xFF - msb.bin : msb.bin - 0x80; // raw high byte + + // Pass 2: among elements sharing the pivot's high byte, histogram the raw low byte. The high + // bytes are compared as tiny positive bf16 values (exact), the others land in the sink bin. + const auto hist_lsb = smem.histogram[1]; + const auto pivot_hi_x2 = bitcast(pivot_hi << 16 | pivot_hi); + constexpr uint32_t kSinkBinX2 = C::kSinkBin << 16 | C::kSinkBin; +#pragma unroll + for (uint32_t i = 0; i < C::kMaxVecs; ++i) { + if (i >= num_my) break; +#pragma unroll + for (uint32_t j = 0; j < C::kVecSize / 2; ++j) { + const auto raw = bitcast(vecs[i][j]); + const auto hi = __byte_perm(raw, 0, 0x5341); // {byte1, 0, byte3, 0} + const auto sel = __heq2_mask(bitcast(hi), pivot_hi_x2); + const auto lo = raw & 0x00FF00FFu; + const auto bins = (sel & lo) | (~sel & kSinkBinX2); // sel ? lo : kSinkBin + atomicAdd(hist_lsb + __byte_perm(bins, 0, 0x4410), 1); // bins & 0xFFFF + atomicAdd(hist_lsb + __byte_perm(bins, 0, 0x4432), 1); // bins >> 16 + } + } + __syncthreads(); + + const auto lsb = pivot_of(hist_lsb, msb.remain, false, negative); + const uint32_t pivot_lo = negative ? 0xFF - lsb.bin : lsb.bin; + const uint32_t pivot_bits = pivot_hi << 8 | pivot_lo; + const auto pivot_x2 = bitcast(pivot_bits << 16 | pivot_bits); + + // Census: one bit per element of the slice, in element order (vector i fills byte i) + uint32_t gt_mask = 0; + uint32_t eq_mask = 0; +#pragma unroll + for (uint32_t i = 0; i < C::kMaxVecs; ++i) { + if (i >= num_my) break; + uint32_t gt[4], eq[4]; +#pragma unroll + for (uint32_t j = 0; j < C::kVecSize / 2; ++j) { + gt[j] = __hgt2_mask(vecs[i][j], pivot_x2); + eq[j] = __heq2_mask(vecs[i][j], pivot_x2); + } + // drop the new byte into slot i, keeping the other three + constexpr uint32_t kInsert[4] = {0x3214, 0x3240, 0x3410, 0x4210}; + gt_mask = __byte_perm(gt_mask, topk_bf16_pack_hits(gt), kInsert[i]); + eq_mask = __byte_perm(eq_mask, topk_bf16_pack_hits(eq), kInsert[i]); + } + const uint32_t cnt_gt = __popc(gt_mask); + const uint32_t cnt_eq = __popc(eq_mask); + + // Block-wide exclusive prefix of (gt, eq), packed: one warp scan plus one shared atomic per + // warp. Warps land in arrival order, which is fine since the output is unordered. + const uint32_t local = cnt_gt << 16 | cnt_eq; + const uint32_t warp_inc = warp::inclusive_sum(lane_id, local); + uint32_t warp_base = 0; + if (lane_id == kWarpThreads - 1) warp_base = atomicAdd(&smem.count_gt_eq, warp_inc); + warp_base = __shfl_sync(kFullMask, warp_base, kWarpThreads - 1); + const uint32_t before = warp_base + warp_inc - local; + + // Everything above the pivot is taken, plus `remain` of the elements equal to it. + uint32_t eq_total = lsb.remain; + if (pivot_bits == C::kNegZeroBits) { + /// NOTE: the census compares as floats, so a -0 pivot also sees +0 as equal while the + /// histogram ranked +0 above it. Both are worth the same, so let the equal quota absorb + /// them: the quota then has to come from the census total (one extra barrier, rare). + __syncthreads(); + eq_total = topk - (smem.count_gt_eq >> 16); + } + const uint32_t gt_before = before >> 16; + const uint32_t eq_before = before & 0xFFFF; + const uint32_t eq_start = min(eq_before, eq_total); + const uint32_t eq_quota = min(eq_before + cnt_eq, eq_total) - eq_start; + // keep only `eq_quota` of the equal bits (which ones does not matter) + if (eq_quota == 0) { + eq_mask = 0; + } else { +#pragma unroll 1 + for (uint32_t n = cnt_eq; n > eq_quota; --n) { + eq_mask &= eq_mask - 1; + } + } + + uint32_t hits = gt_mask | eq_mask; + auto* dst = smem.stage + gt_before + eq_start; + const uint32_t elem_base = vec_start * C::kVecSize; + while (hits != 0) { + const auto e = __ffs(hits) - 1; + hits &= hits - 1; + *dst++ = elem_base + e; + } + + PDLTriggerSecondary(); + __syncthreads(); + + // Slots past the census total were never staged (only NaN scores cause that: counted by the + // histogram, never selected); write -1 there. + const uint32_t totals = smem.count_gt_eq; + const uint32_t num_staged = (totals >> 16) + min(totals & 0xFFFFu, eq_total); + // TODO(perf): unroll once k regularly exceeds 512. + for (uint32_t t = tx; t < topk; t += C::kBlockSize) { + out[t] = t < num_staged ? transform(smem.stage[t]) : -1; + } +} + +/// Host entry: bf16 top-k over rows of at most kMaxSeqLen, selected indices +/// written through a per-row table as `table[i >> log2(page_size)] << log2(page_size) +/// | (i & mask)`, -1 past min(topk, seq_len). +template +struct TopKBF16Kernel { + static void transform( + const tvm::ffi::TensorView scores, + const tvm::ffi::TensorView seq_lens, + const tvm::ffi::TensorView page_table, + const tvm::ffi::TensorView page_indices, + const uint32_t page_size) { + using namespace host; + using C = TopKBF16Config; + auto B = SymbolicSize{"batch_size"}; + auto L = SymbolicSize{"max_seq_len"}; + auto S = SymbolicSize{"score_stride"}; + auto K = SymbolicSize{"topk"}; + auto O = SymbolicSize{"page_indices_stride"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + TensorMatcher({B, L}) // scores + .with_strides({S, 1}) + .with_dtype() + .with_device(device_) + .verify(scores); + TensorMatcher({B}) // seq_lens + .with_dtype() + .with_device(device_) + .verify(seq_lens); + TensorMatcher({B, -1}) // page_table + .with_strides({-1, 1}) + .with_dtype() + .with_device(device_) + .verify(page_table); + TensorMatcher({B, K}) // page_indices + .with_strides({O, 1}) + .with_dtype() + .with_device(device_) + .verify(page_indices); + CHECK_HOST(std::has_single_bit(page_size)) << "page_size must be a power of 2"; + CHECK_HOST(L.unwrap() <= C::kMaxSeqLen) << "rows longer than kMaxSeqLen take the streaming top-k"; + /// NOTE: a row base must stay aligned to the vector width, not just the tensor base. + CHECK_HOST(S.unwrap() % C::kVecSize == 0) << "score_stride must keep every row vector-aligned"; + const auto topk = static_cast(K.unwrap()); + CHECK_HOST(topk > 0 && topk <= C::kMaxTopK) << "topk must be in (0, " << C::kMaxTopK << "]"; + const auto params = TopKBF16Params{ + .scores = static_cast(scores.data_ptr()), + .seq_lens = static_cast(seq_lens.data_ptr()), + .page_table = static_cast(page_table.data_ptr()), + .page_indices = static_cast(page_indices.data_ptr()), + .score_stride = S.unwrap(), + .page_table_stride = page_table.stride(0), + .page_indices_stride = O.unwrap(), + .topk = topk, + .page_bits = static_cast(std::countr_zero(page_size)), + }; + LaunchKernel(static_cast(B.unwrap()), C::kBlockSize, device_.unwrap()) + .config({.use_pdl = kPDL}) + .launch(topk_bf16_small_kernel, params); + } +}; + +} // namespace sglang diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/topk_v2.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/topk_v2.cuh index 426f04788..6e0033e20 100644 --- a/python/sglang/kernels/jit/csrc/deepseek_v4/topk_v2.cuh +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/topk_v2.cuh @@ -20,9 +20,9 @@ #include #include +#include #include #include -#include namespace sglang { @@ -38,27 +38,14 @@ enum class TopKMode { using Register2 = impl::TopKRegister<2>; // <= 8192, register-resident, 1 read using Register4 = impl::TopKRegister<4>; // <= 16384, register-resident, 1 read using Streaming = impl::TopKStreaming; -#ifndef USE_ROCM -using Cluster = impl::TopKCluster<8>; -#endif constexpr uint32_t kBlockSize = impl::TopKConfig::kBlockSize; constexpr uint32_t kOccupancy = impl::TopKConfig::kOccupancy; constexpr uint32_t kMaxTopK = impl::TopKConfig::kMaxTopK; -#ifndef USE_ROCM -constexpr uint32_t kClusterSize = Cluster::kClusterSize; -#endif constexpr uint32_t kReg2MaxSeqLen = Register2::kMaxSeqLen; // 8192 constexpr uint32_t kReg4MaxSeqLen = Register4::kMaxSeqLen; // 16384 #define TOPK_KERNEL __global__ __launch_bounds__(kBlockSize, kOccupancy) -#ifndef USE_ROCM -#define CLUSTER_TOPK_KERNEL TOPK_KERNEL __cluster_dims__(1, kClusterSize, 1) -#endif - -constexpr uint32_t kClusterFloor = 65536; -constexpr uint32_t kClusterMaxBatch = 512; -constexpr uint32_t kNumPersistentClusters = 15 * kOccupancy; /// Metadata tensor rows (each 8 B / 2 int32). Row 0 is the global plan result; /// rows 1..N are the (batch_id, seq_len) of items routed to the cluster pool. @@ -72,18 +59,30 @@ struct alignas(8) PlanItem { }; static_assert(sizeof(GlobalMetadata) == 2 * sizeof(int32_t) && sizeof(PlanItem) == sizeof(GlobalMetadata)); +struct PageTransform { + const int32_t* __restrict__ page_table; + uint32_t page_bits; + int32_t* __restrict__ raw_out; // the row's raw output, written in DUAL_OUTPUT only + + SGL_DEVICE int32_t page_to_indices(uint32_t i) const { + const uint32_t mask = (1u << page_bits) - 1u; + return (page_table[i >> page_bits] << page_bits) | (i & mask); + } +}; + struct TopKPagedParams { const float* __restrict__ scores; const int32_t* __restrict__ seq_lens; const int32_t* __restrict__ page_table; int32_t* __restrict__ page_indices; - int32_t* __restrict__ raw_indices; + int32_t* __restrict__ raw_indices; // DUAL_OUTPUT only, nullptr otherwise const PlanItem* __restrict__ metadata; // [0]=GlobalMetadata, [1+i]=PlanItem int64_t score_stride; int64_t page_table_stride; uint32_t topk; uint32_t page_bits; - uint32_t cluster_floor; // seq_len > this routes to the cluster path (batch-aware, host-set) + uint32_t static_cluster_floor; // only used in small batch variant + uint32_t batch_size; SGL_DEVICE const GlobalMetadata& global() const { return *reinterpret_cast(metadata); @@ -97,18 +96,19 @@ struct TopKPagedParams { SGL_DEVICE int32_t* get_output_ptr(uint32_t batch_id) const { return page_indices + batch_id * static_cast(topk); } - SGL_DEVICE int32_t* get_raw_output_ptr(uint32_t batch_id) const { - return raw_indices == nullptr ? nullptr : raw_indices + batch_id * static_cast(topk); + SGL_DEVICE PageTransform get_transform(uint32_t batch_id) const { + return { + page_table == nullptr ? nullptr : page_table + batch_id * page_table_stride, + page_bits, + raw_indices == nullptr ? nullptr : raw_indices + batch_id * static_cast(topk)}; } SGL_DEVICE TopKProblem problem(uint32_t batch_id, uint32_t seq_len) const { const auto k = static_cast(topk); return TopKProblem{ .in = scores + batch_id * score_stride, .out = page_indices + batch_id * k, - .page_table = page_table + batch_id * page_table_stride, .topk = topk, .seq_len = seq_len, - .page_bits = page_bits, }; } SGL_DEVICE TopKProblem problem(uint32_t batch_id) const { @@ -126,30 +126,9 @@ struct TopKRaggedParams { uint32_t topk; }; -#ifndef USE_ROCM -/** - * \brief Persistent cluster kernel for the long items. It will handle long inputs. - * The short items are handled by the separate topk_kernel. - */ -template -CLUSTER_TOPK_KERNEL void topk_persistent_cluster_kernel(const __grid_constant__ TopKPagedParams params) { - device::enable_smem_spilling(); - __shared__ impl::MaxSmem smem; - const uint32_t num_cluster_items = params.global().num_cluster_items; - device::PDLWaitPrimary(); - device::PDLTriggerSecondary(); -#pragma unroll 1 - for (uint32_t w = blockIdx.x; w < num_cluster_items; w += kNumPersistentClusters) { - const auto it = params.item(w); - const auto problem = params.problem(it.batch_id, it.seq_len); - Cluster::forward(problem, &smem); - __syncthreads(); - } -} -#endif // !USE_ROCM - template SGL_DEVICE void for_each_item(uint32_t topk, const F& f) { + static_assert(kMaxTopK % kBlockSize == 0); constexpr uint32_t kNumElems = kMaxTopK / kBlockSize; #pragma unroll for (uint32_t i = 0; i < kNumElems; ++i) { @@ -161,31 +140,35 @@ SGL_DEVICE void for_each_item(uint32_t topk, const F& f) { } template -SGL_DEVICE void trivial_transform(const TopKProblem& problem, int32_t* raw_output_ptr) { +SGL_DEVICE void trivial_transform(const TopKProblem& problem, const PageTransform& transform) { device::PDLWaitPrimary(); device::PDLTriggerSecondary(); for_each_item(problem.topk, [&](uint32_t tx, uint32_t) { - const auto idx = tx < problem.seq_len ? static_cast(tx) : -1; if constexpr (kMode == TopKMode::INDICES) { - problem.emit(tx, idx); + problem.out[tx] = tx < problem.seq_len ? static_cast(tx) : -1; } else { - problem.transform_output(tx, idx); - if constexpr (kMode == TopKMode::DUAL_OUTPUT) raw_output_ptr[tx] = idx; + problem.out[tx] = tx < problem.seq_len ? transform.page_to_indices(tx) : -1; + if constexpr (kMode == TopKMode::DUAL_OUTPUT) { + transform.raw_out[tx] = tx < problem.seq_len ? static_cast(tx) : -1; + } } }); } template -SGL_DEVICE void problem_transform(TopKProblem& problem, int32_t* output_ptr, int32_t* raw_output_ptr) { - static_assert(kMode != TopKMode::INDICES, "problem_transform requires page-table output"); +SGL_DEVICE void paged_transform(const TopKProblem& problem, int32_t* out, const PageTransform& transform) { + static_assert(kMode != TopKMode::INDICES, "paged_transform requires page-table output"); static_assert(kMaxTopK % kBlockSize == 0); constexpr uint32_t kNumElems = kMaxTopK / kBlockSize; - int32_t source_index[kNumElems]; - for_each_item(problem.topk, [&](uint32_t tx, uint32_t i) { source_index[i] = problem.out[tx]; }); - problem.out = output_ptr; + int32_t indices[kNumElems]; for_each_item(problem.topk, [&](uint32_t tx, uint32_t i) { - problem.transform_output(tx, source_index[i]); - if constexpr (kMode == TopKMode::DUAL_OUTPUT) raw_output_ptr[tx] = source_index[i]; + // load into register at once + indices[i] = problem.out[tx]; + }); + for_each_item(problem.topk, [&](uint32_t tx, uint32_t i) { + // safe write to output + out[tx] = indices[i] >= 0 ? transform.page_to_indices(indices[i]) : -1; + if constexpr (kMode == TopKMode::DUAL_OUTPUT) transform.raw_out[tx] = indices[i]; }); } @@ -244,18 +227,17 @@ TOPK_KERNEL void topk_ragged_kernel(const __grid_constant__ TopKRaggedParams par device::PDLWaitPrimary(); static_assert(kVecSize <= kBlockSize, "not enough threads "); if (const auto tx = threadIdx.x; tx < rem) { - score[row_start - rem + tx] = -std::numeric_limits::max(); + score[row_start - rem + tx] = impl::padding_value(); } } - + using device::topk::broadcast; const auto problem = TopKProblem{ .in = score + (row_start - rem), .out = out, - .page_table = nullptr, // unused .topk = topk, .seq_len = seq_len + rem, - .page_bits = 1, // unused - .bias = offset - static_cast(rem), + .bias = broadcast(offset - static_cast(rem)), + .input_start = broadcast(rem), }; __shared__ impl::MaxSmem smem; if (problem.seq_len <= Register2::kMaxSeqLen) { @@ -280,8 +262,7 @@ TOPK_KERNEL void topk_ragged_kernel(const __grid_constant__ TopKRaggedParams par template TOPK_KERNEL void topk_main_kernel(const __grid_constant__ TopKPagedParams params) { device::enable_smem_spilling(); - auto problem = params.problem(blockIdx.x); - constexpr uint32_t kU32Max = std::numeric_limits::max(); + constexpr bool kNeedStaging = kMode != TopKMode::INDICES; constexpr bool kHandleCluster = (kLevel == 3); // Only the cluster path consumes the cluster kernel's output, so only it waits // on that kernel (kPDLFinal). Every other path waits at most on the indexer @@ -290,15 +271,18 @@ TOPK_KERNEL void topk_main_kernel(const __grid_constant__ TopKPagedParams params constexpr bool kPDLEarly = kPDL && !kHandleCluster; constexpr bool kPDLFinal = kPDL && kHandleCluster; __shared__ impl::MaxSmem smem; - if (problem.seq_len <= problem.topk) - return trivial_transform(problem, params.get_raw_output_ptr(blockIdx.x)); + __shared__ int32_t s_topk_indices[kMaxTopK]; - constexpr bool kNeedStaging = kMode != TopKMode::INDICES; - __shared__ int32_t s_topk_indices[kNeedStaging ? kMaxTopK : 1]; - if constexpr (kNeedStaging) problem.out = s_topk_indices; + const auto bx = blockIdx.x; + auto problem = params.problem(bx); + if (problem.seq_len <= problem.topk) { + return trivial_transform(problem, params.get_transform(bx)); + } + if constexpr (kNeedStaging) { + problem.out = s_topk_indices; // write into stage buffer in smem first + } // non-trivial path: dispatch based on level and seq_len - const auto cluster_threshold = kHandleCluster ? params.cluster_threshold() : kU32Max; if constexpr (kLevel == 0) { __builtin_assume(problem.seq_len <= kReg2MaxSeqLen); Register2::forward(problem, &smem); @@ -306,83 +290,132 @@ TOPK_KERNEL void topk_main_kernel(const __grid_constant__ TopKPagedParams params __builtin_assume(problem.seq_len <= kReg4MaxSeqLen); Register4::forward(problem, &smem); // max_seq_len <= 16384 guarantees seq <= 16384 } else { + const auto cluster_threshold = kHandleCluster ? params.cluster_threshold() : UINT_MAX; static_assert(kLevel == 2 || kLevel == 3, "we only support level = 0,1,2,3 now"); if (problem.seq_len <= kReg4MaxSeqLen) { Register4::forward(problem, &smem); } else if (problem.seq_len <= cluster_threshold) { Streaming::forward(problem, &smem); - } else { + } else [[unlikely]] { // Cluster path: the pool already selected into our output row; the only // work left is the epilogue, so this is the one path that waits for it. - problem.out = params.get_output_ptr(blockIdx.x); - device::PDLWaitPrimary(); + if constexpr (kNeedStaging) { + device::PDLWaitPrimary(); + problem.out = params.get_output_ptr(bx); // in-place transform + device::PDLTriggerSecondary(); + return paged_transform(problem, problem.out, params.get_transform(bx)); + } else { + return device::PDLTriggerSecondary(); + } } } device::PDLTriggerSecondary(); if constexpr (kNeedStaging) { __syncthreads(); - problem_transform(problem, params.get_output_ptr(blockIdx.x), params.get_raw_output_ptr(blockIdx.x)); + paged_transform(problem, params.get_output_ptr(bx), params.get_transform(bx)); } } -#ifndef USE_ROCM -template -CLUSTER_TOPK_KERNEL void topk_small_batch_kernel(const __grid_constant__ TopKPagedParams params) { - device::enable_smem_spilling(); - auto problem = params.problem(blockIdx.x); - __shared__ impl::MaxSmem smem; - if (problem.seq_len <= problem.topk) - return trivial_transform(problem, params.get_raw_output_ptr(blockIdx.x)); +#if SUPPORT_CLUSTER +#ifndef SGL_TOPK_V2_MAX_C8_OCC2 +#if SGL_ARCH_BLACKWELL_OR_GREATER +#define SGL_TOPK_V2_MAX_C8_OCC2 33 // NOTE: B200 +#else +#define SGL_TOPK_V2_MAX_C8_OCC2 30 // NOTE: H200 +#endif +#endif + +#ifndef SGL_TOPK_V2_MAX_C16_OCC1 +#define SGL_TOPK_V2_MAX_C16_OCC1 7 +#endif + +constexpr uint32_t kNumPersistentClusters = SGL_TOPK_V2_MAX_C8_OCC2; +constexpr uint32_t kMaxCluster16BatchSize = SGL_TOPK_V2_MAX_C16_OCC1; +constexpr uint32_t kClusterMaxBatch = 512; +#define CLUSTER_TOPK_KERNEL TOPK_KERNEL __cluster_dims__(1, kClusterSize, 1) + +/// Persistent cluster kernel for the items the plan routed to the pool; topk_main_kernel handles the rest. +template +CLUSTER_TOPK_KERNEL void topk_persistent_cluster_kernel(const __grid_constant__ TopKPagedParams params) { + device::enable_smem_spilling(); + using ClusterN = impl::TopKCluster; + __shared__ impl::MaxSmem smem; + const auto bx = blockIdx.x; + const auto num_cluster_items = params.global().num_cluster_items; + device::PDLWaitPrimary(); + if (bx >= params.batch_size) return; + device::PDLTriggerSecondary(); + auto idx = static_cast(num_cluster_items - 1 - bx); +#pragma unroll 1 + while (idx >= 0) { + const auto it = params.item(idx); + const auto problem = params.problem(it.batch_id, it.seq_len); + ClusterN::template forward(problem, &smem); + idx -= kNumPersistentClusters; + if (idx >= 0) __syncthreads(); + } +} + +template +CLUSTER_TOPK_KERNEL void topk_small_batch_cluster_kernel(const __grid_constant__ TopKPagedParams params) { + device::enable_smem_spilling(); constexpr bool kNeedStaging = kMode != TopKMode::INDICES; - __shared__ int32_t s_topk_indices[kNeedStaging ? kMaxTopK : 1]; - if constexpr (kNeedStaging) problem.out = s_topk_indices; + const auto bx = blockIdx.x; + const auto by = blockIdx.y; + auto problem = params.problem(bx); + __shared__ int32_t s_topk_indices[kMaxTopK]; + using ClusterN = impl::TopKCluster; + __shared__ impl::MaxSmem smem; // randomly elect one worker rank to avoid workload imbalance - const auto worker_rank = blockIdx.x % kClusterSize; + const auto worker_rank = bx % kClusterSize; + if (problem.seq_len <= problem.topk) { + if (by != worker_rank) return; + return trivial_transform(problem, params.get_transform(bx)); + } + if constexpr (kNeedStaging) { + problem.out = s_topk_indices; // write into stage buffer in smem first + } // for small batch, we will fuse in the cluster case if (problem.seq_len <= kReg4MaxSeqLen) { - if (blockIdx.y != worker_rank) return; + if (by != worker_rank) return; Register4::forward(problem, &smem); - __syncthreads(); - } else if (problem.seq_len <= params.cluster_floor) { - if (blockIdx.y != worker_rank) return; + } else if (problem.seq_len <= params.static_cluster_floor) { + if (by != worker_rank) return; Streaming::forward(problem, &smem); - __syncthreads(); } else { auto cluster = cooperative_groups::this_cluster(); if constexpr (kNeedStaging) { - problem.out = cluster.map_shared_rank(s_topk_indices, worker_rank); + problem.out = cluster.map_shared_rank(s_topk_indices, 0); } - Cluster::forward(problem, &smem); + ClusterN::forward(problem, &smem); if constexpr (kNeedStaging) { + device::PDLTriggerSecondary(); cluster.sync(); - if (blockIdx.y != worker_rank) return; + if (by != 0) return; + problem.out = s_topk_indices; + return paged_transform(problem, params.get_output_ptr(bx), params.get_transform(bx)); + } else { + return device::PDLTriggerSecondary(); } } device::PDLTriggerSecondary(); if constexpr (kNeedStaging) { - // Only the elected worker reaches here, and it mapped `topk_indices` to - // itself, so `problem.out` is this block's own buffer. Stating that keeps the - // shared::cluster address out of the load problem_transform issues -- which is - // load-bearing, not an optimization: without it cicc segfaults on CUDA 13.1+ - // for sm_90a (issue #32830, previously worked around by copying `problem` in - // #32910). Verified: dropping this line reproduces the crash on 13.1/13.2/13.3. - __builtin_assume(problem.out == s_topk_indices); - problem_transform(problem, params.get_output_ptr(blockIdx.x), params.get_raw_output_ptr(blockIdx.x)); + __syncthreads(); + paged_transform(problem, params.get_output_ptr(bx), params.get_transform(bx)); } } -#endif // !USE_ROCM // --- Plan: choose cluster_threshold from the seq_len distribution ----------- -__global__ __launch_bounds__(kBlockSize, 1) void topk_plan( +__global__ __launch_bounds__(kBlockSize, 1) void topk_plan_cluster( const uint32_t* __restrict__ seq_lens, PlanItem* __restrict__ metadata, // [0]=GlobalMetadata, [1+i]=PlanItem const uint32_t batch_size, - const uint32_t static_cluster_threshold) { + const int32_t static_cluster_threshold) { // Candidate (threshold T_j, cap_j) pairs, T strictly increasing. The plan lowers // cluster_threshold to T_j while #(items with seq_len > T_j) <= cap_j, so cap_j // bounds how many long items go to the persistent pool. The pool runs N items in @@ -395,16 +428,24 @@ __global__ __launch_bounds__(kBlockSize, 1) void topk_plan( uint32_t max_batch_size; }; constexpr Pair kCandidates[] = { - {65536, 30}, // (65536,98304]: ~1 pool wave, streams beyond 30 - {98304, 48}, // (98304,131072] - {131072, 60}, // (131072,196608] - {196608, 80}, // (196608,262144] - {262144, 112}, // (262144,393216] - {393216, 128}, // (393216,inf): longest -- worth many pool waves; a top - // threshold here lets overloaded ~280-393K batches still stream +#if SGL_ARCH_BLACKWELL_OR_GREATER // tuned on B200 + {32768, 48}, + {131072, 66}, + {163840, 99}, + {196608, 132}, + {262144, 198}, + {393216, 231}, + {524288, 264}, +#else // tuned on H200 + {65536, 30}, + {98304, 45}, + {131072, 60}, + {196608, 80}, + {262144, 112}, + {393216, 128}, +#endif }; constexpr uint32_t kNumCandidates = std::size(kCandidates); - static_assert(kCandidates[0].threshold == kClusterFloor); __shared__ uint32_t s_counts[kNumCandidates]; __shared__ uint32_t s_threshold; @@ -415,15 +456,15 @@ __global__ __launch_bounds__(kBlockSize, 1) void topk_plan( if (tx == 0) s_count = 0; __syncthreads(); - if (static_cluster_threshold > 0) { + if (static_cluster_threshold >= 0) { if (tx == 0) s_threshold = static_cluster_threshold; } else { for (uint32_t i = tx; i < batch_size; i += kBlockSize) { - const uint32_t sl = seq_lens[i]; + const uint32_t seq_len = seq_lens[i]; uint32_t count = 0; #pragma unroll for (uint32_t j = 0; j < kNumCandidates; ++j) { - count += (sl > kCandidates[j].threshold ? 1 : 0); + count += (seq_len > kCandidates[j].threshold ? 1 : 0); } if (count > 0) atomicAdd(&s_counts[count - 1], 1); } @@ -442,15 +483,18 @@ __global__ __launch_bounds__(kBlockSize, 1) void topk_plan( } } __syncthreads(); + + constexpr uint32_t kClusterFloor = 32768; // a very loose lower bound on threshold const auto cluster_threshold = max(s_threshold, kClusterFloor); // Compact items with seq_len > threshold into metadata[1..N]: their batch ids // are the work list the persistent cluster pool fetches. for (uint32_t i = tx; i < batch_size; i += kBlockSize) { - const uint32_t sl = seq_lens[i]; - if (sl > cluster_threshold) { + const uint32_t seq_len = seq_lens[i]; + assert(static_cast(seq_len) >= 0 && "negative seq_len detected"); + if (seq_len > cluster_threshold) { const auto pos = atomicAdd(&s_count, 1); - metadata[1 + pos] = {i, sl}; + metadata[1 + pos] = {i, seq_len}; } } __syncthreads(); @@ -460,11 +504,14 @@ __global__ __launch_bounds__(kBlockSize, 1) void topk_plan( } } +#endif // SUPPORT_CLUSTER + +template struct TopKKernel { static void plan( // const tvm::ffi::TensorView seq_lens, const tvm::ffi::TensorView metadata, - const uint32_t static_cluster_threshold) { + const int32_t static_cluster_threshold) { using namespace host; auto B = SymbolicSize{"batch_size"}; auto Bp1 = SymbolicSize{"batch_size_plus_1"}; @@ -475,25 +522,27 @@ struct TopKKernel { .with_dtype() .with_device(device_) .verify(seq_lens); - TensorMatcher({Bp1, 2}) // metadata: [0]=GlobalMetadata, [1..N]=PlanItem(batch_id, seq_len) + TensorMatcher({-1, 2}) // metadata: [0]=GlobalMetadata, [1..N]=PlanItem(batch_id, seq_len) .with_dtype() .with_device(device_) .verify(metadata); - RuntimeCheck(Bp1.unwrap() == B.unwrap() + 1, "invalid metadata shape"); -#ifdef USE_ROCM - // ROCm compiles out the cluster path, the only consumer of this plan. - (void)static_cluster_threshold; - return; -#else + RuntimeCheck(metadata.size(0) == B.unwrap() + 1, "invalid metadata shape"); +#if SUPPORT_CLUSTER const auto batch_size = static_cast(B.unwrap()); + // persistent cluster not supported + if (kNumPersistentClusters == 0) return; + // will not route to persistent cluster + if (batch_size <= kNumPersistentClusters || batch_size > kClusterMaxBatch) return; const auto device = device_.unwrap(); LaunchKernel(1, kBlockSize, device)( // - topk_plan, + topk_plan_cluster, static_cast(seq_lens.data_ptr()), static_cast(metadata.data_ptr()), batch_size, static_cluster_threshold); +#else + static_cast(static_cluster_threshold); #endif } @@ -507,10 +556,8 @@ struct TopKKernel { const tvm::ffi::Optional raw_indices) { using namespace host; auto B = SymbolicSize{"batch_size"}; - auto Bp1 = SymbolicSize{"batch_size_plus_1"}; auto L = SymbolicSize{"max_seq_len"}; auto S = SymbolicSize{"score_stride"}; - auto P = SymbolicSize{"page_table_stride"}; auto K = SymbolicSize{"topk"}; auto device_ = SymbolicDevice{}; device_.set_options(); @@ -530,32 +577,36 @@ struct TopKKernel { int64_t page_table_stride = 0; if (page_table.has_value()) { TensorMatcher({B, -1}) // page_table - .with_strides({P, 1}) + .with_strides({-1, 1}) .with_dtype() .with_device(device_) .verify(page_table.value()); page_table_ptr = static_cast(page_table.value().data_ptr()); - page_table_stride = P.unwrap(); + page_table_stride = (page_table.value()).stride(0); } TensorMatcher({B, K}) // page_indices .with_dtype() .with_device(device_) .verify(page_indices); - TensorMatcher({Bp1, 2}) // metadata: [0]=GlobalMetadata, [1..N]=PlanItem(batch_id, seq_len) + TensorMatcher({-1, 2}) // metadata: [0]=GlobalMetadata, [1..N]=PlanItem(batch_id, seq_len) .with_dtype() .with_device(device_) .verify(metadata); - + // Present means "both outputs": `page_indices` receives the page-table + // transform and `raw_indices` the selected raw indices, same -1 padding. int32_t* raw_indices_ptr = nullptr; if (raw_indices.has_value()) { RuntimeCheck(page_table.has_value(), "raw_indices requires a page table"); - TensorMatcher({B, K}).with_dtype().with_device(device_).verify(raw_indices.value()); + TensorMatcher({B, K}) // raw_indices + .with_dtype() + .with_device(device_) + .verify(raw_indices.value()); raw_indices_ptr = static_cast(raw_indices.value().data_ptr()); } RuntimeCheck(std::has_single_bit(page_size), "page_size must be power of 2"); RuntimeCheck(S.unwrap() % 4 == 0, "score_stride must be a multiple of 4 (16-byte vectorized load)"); - RuntimeCheck(Bp1.unwrap() == B.unwrap() + 1, "invalid metadata shape"); + RuntimeCheck(metadata.size(0) == B.unwrap() + 1, "invalid metadata shape"); const auto topk = static_cast(K.unwrap()); RuntimeCheck(topk > 0 && topk <= kMaxTopK, "topk must be in (0, 2048]"); @@ -564,13 +615,17 @@ struct TopKKernel { const auto max_seq_len = static_cast(L.unwrap()); const auto device = device_.unwrap(); - // The fused kernel runs one 8-block cluster per batch element, and B200 fits one - // wave of exactly 15 such clusters (occ2). For batch <= 15 it stays latency-bound, - // so the 8-way split beats streaming from a much lower seq (measured crossover - // ~36-40K); batch 16 spills into a 2nd wave (+25%) and keeps the 64K floor. - // The floor is chosen on the host per launch. - constexpr uint32_t kClusterFloorSmall = 32768; - constexpr uint32_t kSmallBatchLowFloor = 15; + constexpr auto get_static_cluster_floor = [](uint32_t batch_size) -> uint32_t { + // NOTE: 15 is exactly 0.5 wave which saturate all cluster-8 SMs on Hopper/Blackwell + if constexpr (SGL_ARCH_BLACKWELL_OR_GREATER) { + return batch_size <= 15 ? 24576 : 30720; + } else if constexpr (SGL_ARCH_HOPPER_OR_GREATER) { + return batch_size <= 15 ? 32768 : 65536; + } else { + return UINT_MAX; + } + }; + const auto params = TopKPagedParams{ .scores = static_cast(scores.data_ptr()), .seq_lens = static_cast(seq_lens.data_ptr()), @@ -582,43 +637,66 @@ struct TopKKernel { .page_table_stride = page_table_stride, .topk = topk, .page_bits = page_bits, - .cluster_floor = (batch_size <= kSmallBatchLowFloor) ? kClusterFloorSmall : kClusterFloor, + // only used in small batch variant + .static_cluster_floor = get_static_cluster_floor(batch_size), + // used for persistent cluster kernel and main kernel + .batch_size = batch_size, }; -#ifndef USE_ROCM - const bool use_cluster = (max_seq_len > params.cluster_floor) && (batch_size <= kClusterMaxBatch); -#endif - constexpr bool kUsePDL = true; - const auto mode = raw_indices.has_value() ? TopKMode::DUAL_OUTPUT - : page_table.has_value() ? TopKMode::PAGE_TABLE - : TopKMode::INDICES; const auto dispatch = [&](F&& f) { + const auto mode = raw_indices.has_value() ? TopKMode::DUAL_OUTPUT + : page_table.has_value() ? TopKMode::PAGE_TABLE + : TopKMode::INDICES; switch (mode) { case TopKMode::INDICES: return f.template operator()(); + case TopKMode::PAGE_TABLE: + return f.template operator()(); case TopKMode::DUAL_OUTPUT: return f.template operator()(); default: - return f.template operator()(); + Panic("Invalid mode, this path should be unreachable"); } }; dispatch([&]() { -#ifndef USE_ROCM +#if SUPPORT_CLUSTER + const bool use_cluster = (max_seq_len > params.static_cluster_floor) && (batch_size <= kClusterMaxBatch); if (use_cluster) { - if (batch_size <= kNumPersistentClusters) { - LaunchKernel({batch_size, kClusterSize}, kBlockSize, device) - .config({.use_pdl = kUsePDL, .cluster_dim = dim3{1, kClusterSize}}) - .launch(topk_small_batch_kernel, params); - } else { - const uint32_t num_clusters = std::min(batch_size, kNumPersistentClusters); - LaunchKernel({num_clusters, kClusterSize}, kBlockSize, device) - .config({.use_pdl = kUsePDL, .cluster_dim = dim3{1, kClusterSize}}) - .launch(topk_persistent_cluster_kernel, params); - LaunchKernel(batch_size, kBlockSize, device) - .config({.use_pdl = kUsePDL}) - .launch(topk_main_kernel, params); + if constexpr (kMaxCluster16BatchSize > 0) { + if (batch_size <= kMaxCluster16BatchSize) { + constexpr uint32_t kClusterSize = 16; + // Widths above 8 are non-portable; the launch is rejected without this. + const auto kernel = topk_small_batch_cluster_kernel; + [[maybe_unused]] + static const bool _ = [&kernel] { + const auto kernel_ptr = reinterpret_cast(kernel); + CHECK_CUDA(::cudaFuncSetAttribute(kernel_ptr, ::cudaFuncAttributeNonPortableClusterSizeAllowed, 1)); + return true; + }(); + return LaunchKernel({batch_size, kClusterSize}, kBlockSize, device) + .config({.use_pdl = kUsePDL, .cluster_dim = dim3{1, kClusterSize}}) + .launch(kernel, params); + } + } + + if constexpr (kNumPersistentClusters > 0) { + if (batch_size <= kNumPersistentClusters) { + constexpr uint32_t kClusterSize = 8; + return LaunchKernel({batch_size, kClusterSize}, kBlockSize, device) + .config({.use_pdl = kUsePDL, .cluster_dim = dim3{1, kClusterSize}}) + .launch(topk_small_batch_cluster_kernel, params); + } else { + constexpr uint32_t kClusterSize = 8; + const uint32_t num_clusters = std::min(batch_size, kNumPersistentClusters); + LaunchKernel({num_clusters, kClusterSize}, kBlockSize, device) + .config({.use_pdl = kUsePDL, .cluster_dim = dim3{1, kClusterSize}}) + .launch(topk_persistent_cluster_kernel, params); + LaunchKernel(batch_size, kBlockSize, device) + .config({.use_pdl = kUsePDL}) + .launch(topk_main_kernel, params); + return void(); + } } - return; } #endif if (max_seq_len <= kReg2MaxSeqLen) { @@ -694,7 +772,6 @@ struct TopKKernel { const auto topk = static_cast(K.unwrap()); RuntimeCheck(topk > 0 && topk <= kMaxTopK, "topk must be in (0, 2048]"); - constexpr bool kUsePDL = true; const auto params = TopKRaggedParams{ .scores = static_cast(scores.data_ptr()), .seq_lens = static_cast(seq_lens.data_ptr()), diff --git a/python/sglang/kernels/jit/csrc/occupancy/cluster_probe.cuh b/python/sglang/kernels/jit/csrc/occupancy/cluster_probe.cuh new file mode 100644 index 000000000..1a990c715 --- /dev/null +++ b/python/sglang/kernels/jit/csrc/occupancy/cluster_probe.cuh @@ -0,0 +1,55 @@ +#pragma once + +#include + +#include + +namespace sglang { + +__global__ void dummy_probe_kernel() {} + +uint32_t get_max_active_clusters(uint32_t cluster_size, uint32_t num_waves) { +#if !SGL_ARCH_HOPPER_OR_GREATER + host::Panic("cluster is not supported on arch before CUDA sm90"); +#else + int device; + int max_threads_per_sm; + int smem_per_sm; + int smem_per_block; + int num_clusters; + CHECK_CUDA(cudaGetDevice(&device)); + CHECK_CUDA(cudaDeviceGetAttribute(&max_threads_per_sm, cudaDevAttrMaxThreadsPerMultiProcessor, device)); + CHECK_CUDA(cudaDeviceGetAttribute(&smem_per_sm, cudaDevAttrMaxSharedMemoryPerMultiprocessor, device)); + CHECK_CUDA(cudaDeviceGetAttribute(&smem_per_block, cudaDevAttrMaxSharedMemoryPerBlockOptin, device)); + + // A block caps at 1024 threads, so threads alone cannot pin `num_waves` blocks + // per SM; spend the shared budget instead, less the per-block driver reserve, + // floored to the 1 KiB granularity so the driver cannot round it back up. + const auto reserved = static_cast(smem_per_sm - smem_per_block); + const auto budget = static_cast(smem_per_sm) / num_waves; + const auto smem = (std::min(budget - std::min(budget, reserved), static_cast(smem_per_block))) & ~1023u; + const auto num_warps = std::max(1u, std::min(1024u, static_cast(max_threads_per_sm) / num_waves) / 32); + + // Widths above 8 are non-portable and the query rejects them without this. + CHECK_CUDA(cudaFuncSetAttribute( + reinterpret_cast(dummy_probe_kernel), cudaFuncAttributeNonPortableClusterSizeAllowed, 1)); + CHECK_CUDA(cudaFuncSetAttribute( + reinterpret_cast(dummy_probe_kernel), + cudaFuncAttributeMaxDynamicSharedMemorySize, + static_cast(smem))); + + cudaLaunchConfig_t config = {}; // stream/dynamicSmemBytes must not be garbage + config.gridDim = dim3{cluster_size, 1024u}; + config.blockDim = dim3{32, num_warps}; + config.dynamicSmemBytes = smem; + config.numAttrs = 1; + cudaLaunchAttribute attr = {}; + attr.id = cudaLaunchAttributeClusterDimension; + attr.val.clusterDim = {cluster_size, 1, 1}; + config.attrs = &attr; + CHECK_CUDA(cudaOccupancyMaxActiveClusters(&num_clusters, dummy_probe_kernel, &config)); + return num_clusters; +#endif +} + +} // namespace sglang diff --git a/python/sglang/kernels/jit/include/sgl_kernel/deepseek_v4/topk_impl.cuh b/python/sglang/kernels/jit/include/sgl_kernel/deepseek_v4/topk_impl.cuh index 7dedd89a5..8836389c8 100644 --- a/python/sglang/kernels/jit/include/sgl_kernel/deepseek_v4/topk_impl.cuh +++ b/python/sglang/kernels/jit/include/sgl_kernel/deepseek_v4/topk_impl.cuh @@ -7,11 +7,11 @@ /// Design notes: /// - top-k (`topk`) is a *runtime* value (<= kMaxTopK = 2048), never a /// compile-time constant. -/// - the output is the page-table transform of the selected raw indices -/// (`TopKProblem::emit` then `transform_output`). +/// - the dispatcher optionally transforms selected raw indices through a page +/// table after the device implementation writes them. /// - each block reads its own `seq_len` (per-batch ragged lengths) -- the host /// launches one universal kernel and dispatches per block. -/// - the cluster size is fixed at 8 (dynamic persistent clusters are hard). +/// - the dispatcher selects cluster size 8 or 16 from the probed occupancy. /// /// Algorithm: fp16 coarse histogram -> threshold bin -> fp32-boundary collect -> /// exact radix tie-break. @@ -23,11 +23,20 @@ #include #include +#include #include #include #include -#ifndef USE_ROCM +#if !defined(USE_ROCM) +// currently only apply cluster for SM90 & SM100, SM120 has poor cluster performance +#define SUPPORT_CLUSTER (SGL_CUDA_ARCH >= 900 && SGL_CUDA_ARCH < 1100) +#else +// AMD doesn't support cluster +#define SUPPORT_CLUSTER false +#endif + +#if SUPPORT_CLUSTER #include #endif @@ -35,40 +44,27 @@ namespace sglang { namespace device::topk { -#ifndef USE_ROCM -namespace cg = cooperative_groups; +/// Hints that `value` is warp-uniform so it can live in a uniform register. The +/// caller must already guarantee that: on ROCm this is the identity, since the +/// 32-bit mask below covers only half of a 64-lane wavefront and there is no +/// uniform register file to hint at. +template +SGL_DEVICE T broadcast(T value, uint32_t src = 0) { +#if defined(USE_ROCM) + static_cast(src); + return value; +#else + return __shfl_sync(0xFFFFFFFF, value, src); #endif +} /// sgl_kernel names the warp size `kWarpThreads`; alias it locally as `kWarpSize`. inline constexpr uint32_t kWarpSize = kWarpThreads; -// --------------------------------------------------------------------------- -// Shared-memory storage sized/aligned for several impl `Smem` types -// --------------------------------------------------------------------------- - -/// Compile-time max over a non-empty pack (avoids an dependency). -template -constexpr T ct_max(T a) { - return a; -} -template -constexpr T ct_max(T a, Ts... rest) { - const T m = ct_max(rest...); - return a > m ? a : m; -} - -/// Static shared-memory buffer sized + aligned to hold any one of the given -/// impl `Smem` types. A kernel that dispatches across several paths (e.g. the -/// fused small-batch kernel runs either Streaming or Cluster; the main kernel -/// runs any of Register2/Register4/Streaming) declares one -/// `__shared__ MaxSmem<...> smem` and hands `&smem` to whichever forward() it -/// calls -- instead of hand-picking "the largest" type and relying on it -/// staying the largest. `&smem` converts to the `void*` the forwards expect; -/// the buffer is aligned to the strictest member, so the cast is well-aligned. template struct MaxSmem { - static constexpr size_t kSize = ct_max(sizeof(Smems)...); - static constexpr size_t kAlign = ct_max(alignof(Smems)...); + static constexpr size_t kSize = std::max({sizeof(Smems)...}); + static constexpr size_t kAlign = std::max({alignof(Smems)...}); alignas(kAlign) uint8_t storage[kSize]; }; @@ -81,64 +77,75 @@ SGL_DEVICE uint32_t extract_exact_bin(float x) { return (bits & 0x80000000u) ? ~bits : (bits | 0x80000000u); } +constexpr float padding_value() { + return std::numeric_limits::quiet_NaN(); +} + +constexpr float infinity_value() { + return std::numeric_limits::infinity(); +} + template SGL_DEVICE uint32_t extract_coarse_bin(float x) { static_assert(0 < kBits && kBits < 15); - const auto hx = cast(x); - const uint16_t bits = *reinterpret_cast(&hx); - const uint16_t key = (bits & 0x8000) ? ~bits : bits | 0x8000; - return key >> (16 - kBits); + uint32_t b = (uint32_t)__half_as_ushort(__float2half_rn(x)) << 16; + uint32_t s = (uint32_t)((int32_t)b >> 31); + return (b ^ (s | 0x80000000u)) >> (32 - kBits); } -// Smallest fp32 value `v` for which `extract_coarse_bin(v) >= bin`, i.e. the -// lower fp32 boundary of coarse bin `bin`. Because `extract_coarse_bin` is monotonic -// non-decreasing in its argument, the collect pass can classify an element with two -// fp32 comparisons against these boundaries instead of recomputing the fp16 bin -- -// removing the F2F conversion and bit-twiddle from the (compute-bound) second pass. -// Returns -inf for bin 0 (everything qualifies) and +inf for bins past the top. +SGL_DEVICE uint16_t coarse_bin_to_bits_finite(uint32_t bin) { + const uint16_t ob = static_cast(bin); + return (ob & 0x8000) ? static_cast(ob ^ 0x8000) : static_cast(~ob); +} + +// Smallest fp32 `v` for which `extract_coarse_bin(v) >= bin`, i.e. the +// lower fp32 boundary of coarse bin `bin`. The collect pass classifies with two +// comparisons against these instead of recomputing the fp16 bin per element, so +// this must agree with `extract_coarse_bin` on every value -- a score sitting +// exactly on a boundary included. Two pairs no fp32 threshold can separate are +// left: -0.0 at the zero bin, and +inf at a NaN-key bin. template SGL_DEVICE float coarse_bin_lower_bound(uint32_t bin) { constexpr uint32_t kShift = 16 - kBits; - const uint32_t key = bin << kShift; // ordered16 key at the low edge of `bin` + constexpr uint32_t kInfBin = 0xFC00u >> kShift; // bin holding the +inf key + const uint32_t key = bin << kShift; // ordered16 key at the low edge // ordered16 -> fp16 value (inverse of the transform in extract_coarse_bin); - // finite keys only. - const auto to_finite_val = [](uint32_t okey) -> float { - const uint16_t ob = static_cast(okey); - const uint16_t hb = (ob & 0x8000) ? static_cast(ob ^ 0x8000) : static_cast(~ob); + constexpr auto to_finite_val = [](uint32_t okey) -> float { + const uint16_t hb = coarse_bin_to_bits_finite(okey); return cast(*reinterpret_cast(&hb)); }; + constexpr auto step_up = [](float v) -> float { + const int32_t b = __float_as_int(v); + return __int_as_float(b >= 0 ? b + 1 : b - 1); + }; // Fast path, hoisted above the per-key special cases so both keys are // range-checked at once: `key` and `key - 1` both land in the finite band - // [0x0401, 0xFBFF] -- every boundary a finite-score threshold produces. - // fp16 rounds to nearest, so the fp32 boundary is the midpoint between the - // fp16 values at `key` and `key - 1`. (Verified bit-exact against the slow - // path for every bin of kBits 10 and 12, and measured faster than either - // per-key dispatch or an ordered-bit decrement trick -- the two conversions - // are independent and issue in parallel.) - if (key - 0x0401u <= 0xFBFFu - 0x0401u && bin < (1u << kBits)) { - return 0.5f * (to_finite_val(key) + to_finite_val(key - 1)); + // [0x0401, 0xFBFF] -- every boundary a finite-score threshold produces. fp16 + // rounds to nearest, so the boundary is the midpoint between the fp16 values + // at `key` and `key - 1`. + if (key - 0x0401u <= 0xFBFFu - 0x0401u) { + const float mid = 0.5f * (to_finite_val(key) + to_finite_val(key - 1)); + // fp32 -> fp16 rounds to nearest EVEN, so on the ~half of bins whose fp16 + // value has an odd significand the midpoint still bins as `bin - 1`. + return (coarse_bin_to_bits_finite(key) & 1u) ? step_up(mid) : mid; } - // Slow path: an edge of `bin` touches the +/-inf keys or NaN key space. - // The ordered-key line is: [0, 0x03FF) negative-NaN space, 0x03FF = -inf, + // Slow path: an edge of `bin` touches the +/-inf keys or NaN key space. The + // ordered-key line is: [0, 0x03FF) negative-NaN space, 0x03FF = -inf, // [0x0400, 0xFC00) finite, 0xFC00 = +inf, (0xFC00, 0xFFFF] positive-NaN - // space. Treat the +/-inf keys as +/-65536 (one ideal step past fp16 max, - // so the midpoint lands exactly on +/-65520 -- the fp32->fp16 - // round-to-nearest overflow threshold) and saturate NaN-space keys, keeping - // the returned boundaries finite-or-inf and monotone. Otherwise a threshold - // bin at/next to the inf bin gets NaN boundaries, the collect pass matches - // nothing, and rows whose scores contain >= topk (+/-)inf or >65504 values - // come back short -- the padded slots then illegal-address downstream. - if (bin == 0) return -FLT_MAX; - if (bin >= (1u << kBits)) return FLT_MAX; + // space. The +/-inf keys stand in as +/-65536, one ideal step past fp16 max, + // so the midpoint lands on the +/-65520 fp32 -> fp16 overflow threshold. + if (bin == 0) return -infinity_value(); // every value bins at >= 0 + if (bin > kInfBin) return infinity_value(); // NaN key space: nothing bins that high const auto to_val = [&](uint32_t okey) -> float { - constexpr float k_Inf = std::numeric_limits::infinity(); - if (okey < 0x03FFu) return -k_Inf; + if (okey < 0x03FFu) return -infinity_value(); if (okey == 0x03FFu) return -65536.0f; if (okey == 0xFC00u) return 65536.0f; - if (okey > 0xFC00u) return FLT_MAX; return to_finite_val(okey); }; - return 0.5f * (to_val(key) + to_val(key - 1)); + // The +/-65536 stand-ins are not real fp16 neighbours, so the parity rule + // does not apply here; test the property directly instead. + const float mid = 0.5f * (to_val(key) + to_val(key - 1)); + return extract_coarse_bin(mid) < bin ? step_up(mid) : mid; } SGL_DEVICE uint32_t warp_inclusive_sum(uint32_t lane_id, uint32_t val) { @@ -171,7 +178,7 @@ struct alignas(8) TieValue { float value; uint32_t idx; inline static constexpr TieValue invalid() { - return TieValue{-FLT_MAX, 0xFFFFFFFFu}; + return TieValue{padding_value(), 0xFFFFFFFFu}; } }; @@ -179,29 +186,17 @@ struct alignas(8) TieValue { // Per-batch problem description + page-table transform sink // --------------------------------------------------------------------------- -SGL_DEVICE int32_t page_to_indices(const int32_t* __restrict__ page_table, uint32_t i, uint32_t page_bits) { - const uint32_t mask = (1u << page_bits) - 1u; - return (page_table[i >> page_bits] << page_bits) | (i & mask); -} - -/// One batch element's worth of work. `emit(pos, raw_idx)` writes the selected raw -/// index to output slot `pos`; `transform_output` then applies the page-table -/// transform in a separate pass. struct TopKProblem { const float* __restrict__ in; int32_t* __restrict__ out; // page_indices [topk] - const int32_t* __restrict__ page_table; uint32_t topk; uint32_t seq_len; - uint32_t page_bits; - int32_t bias = 0; // needed by ragged mode + int32_t bias = 0; + uint32_t input_start = 0; // needed by ragged mode SGL_DEVICE void emit(uint32_t pos, uint32_t raw_idx) const { out[pos] = static_cast(raw_idx) + bias; } - SGL_DEVICE void transform_output(uint32_t t, int32_t raw) const { - out[t] = raw < 0 ? -1 : page_to_indices(page_table, raw, page_bits); - } }; // --------------------------------------------------------------------------- @@ -227,14 +222,13 @@ struct TopKConfig { static_assert(kMaxNumTie >= kMaxTopK && kMaxNumTie % kBlockSize == 0 && kBlockSize % kNumWarps == 0); struct TieHandleSmem { - struct alignas(16) MatchBin { + struct MatchBin { uint32_t bin; uint32_t above_count; uint32_t equal_count; - uint32_t _pad = 0; }; - alignas(128) uint32_t counter; - alignas(128) uint32_t counter_final; + uint32_t counter; + uint32_t counter_final; MatchBin match; uint32_t warp_sum[kNumWarps]; uint32_t histogram[2][kRadixSize]; @@ -255,7 +249,7 @@ struct TopKConfig { }; const auto tx = threadIdx.x; const auto lane_id = tx % kWarpSize; - const auto warp_id = tx / kWarpSize; + const auto warp_id = broadcast(tx / kWarpSize); static_assert(kNumWarps == kWarpSize); if (num_ties <= topk) { @@ -322,11 +316,11 @@ struct TopKConfig { } } else if (num_ties <= kBlockSize) { // Common case: one candidate per thread. - radix_tie_select<1>(tie_buffer, problem, base, num_ties, topk, smem); + return radix_tie_select<1>(tie_buffer, problem, base, num_ties, topk, smem); } else { - // Rare overflow case (kBlockSize < num_ties <= kMaxNumTie), kept out of - // the common path so it alone pays the multi-item register cost. - radix_tie_select(tie_buffer, problem, base, num_ties, topk, smem); + // Rare overflow case. + static_assert(kTieItems == 2); + return radix_tie_select<2>(tie_buffer, problem, base, num_ties, topk, smem); } } @@ -343,7 +337,7 @@ struct TopKConfig { TieHandleSmem* smem) { const auto tx = threadIdx.x; const auto lane_id = tx % kWarpSize; - const auto warp_id = tx / kWarpSize; + const auto warp_id = broadcast(tx / kWarpSize); bool active[kItems]; uint32_t key[kItems]; @@ -398,7 +392,7 @@ struct TopKConfig { } __syncthreads(); - const auto [threshold_bin, above_count, equal_count, __] = smem->match; + const auto [threshold_bin, above_count, equal_count] = smem->match; if (round < 3) total_active = equal_count; topk_remain -= above_count; @@ -434,96 +428,116 @@ struct TopKConfig { template struct TopKRadixBase : TopKConfig { + public: static constexpr uint32_t kVecSize = 4; static constexpr uint32_t kHistBits = kHistBits_; static constexpr uint32_t kHistSize = 1 << kHistBits; using vec_t = AlignedVector; struct Smem { - using kHistVec = AlignedVector; - alignas(128) uint32_t count_eq; - alignas(128) uint32_t count_gt; - uint32_t threshold_bin; + uint32_t count_eq; + uint32_t count_gt; + float v_hi; + float v_lo; uint32_t warp_sum[kNumWarps]; - // The coarse histogram is dead once find_threshold() has published - // threshold_bin, and the tie machinery only comes alive after that: the - // collect pass fills tie.values, then handle_tie works over them with - // tie.handle as scratch. Overlaying the two phases keeps the - // kMaxNumTie-candidate buffer from growing the block's shared-memory - // footprint. tie.handle and tie.values are live TOGETHER, so they sit + // The coarse histogram is dead once find_threshold() has published its + // boundaries, and the tie machinery only comes alive after that, so the two + // phases overlay. tie_handle and tie_values are live TOGETHER, so they sit // side by side inside the overlay, not in a union with each other. union { - uint32_t histogram[kHistSize]; - kHistVec hist_vecs[kBlockSize]; + alignas(16) uint32_t histogram[kHistSize]; struct { - TieHandleSmem handle; - TieValue values[kMaxNumTie]; - } tie; + TieValue tie_values[kMaxNumTie]; + TieHandleSmem tie_handle; + }; }; }; protected: - template + template SGL_DEVICE static void for_each_input(const float* __restrict__ in, uint32_t seq_len, F&& fn) { + constexpr auto kStride = N * kBlockSize; const auto tx = threadIdx.x; - const uint32_t num_full = seq_len / kVecSize; // fully-in-bounds vectors - - vec_t next_vec; - uint32_t vi = tx; - if (vi < num_full) next_vec.load(in, vi); - while (vi < num_full) { - const auto cur = next_vec; - const auto base = vi * kVecSize; - vi += kBlockSize; - if (vi < num_full) next_vec.load(in, vi); + const auto num_full = seq_len / kVecSize; // fully-in-bounds vectors + const auto kChunk = 128u; + // lane | rank | warp + auto vi = N == 1 ? tx : (tx % kChunk) + blockIdx.y * kChunk + (tx / kChunk) * (N * kChunk); + if (vi < num_full) { + vec_t next_vec; + next_vec.load(in, vi); +#pragma unroll 1 + do { + const auto cur = next_vec; + vi += kStride; + if (vi < num_full) next_vec.load(in, vi); + const auto base = (vi - kStride) * kVecSize; #pragma unroll - for (uint32_t j = 0; j < kVecSize; ++j) { - fn(cur[j], base + j); - } + for (uint32_t j = 0; j < kVecSize; ++j) { + fn(cur[j], base + j); + } + } while (vi < num_full); } - // Tail: at most one partial vector, `rem` in [0, kVecSize). - static_assert(kVecSize <= kBlockSize); // ensure tail correctness - const uint32_t tail_start = num_full * kVecSize; - if (tx < seq_len - tail_start) { - const auto idx = tail_start + tx; - fn(in[idx], idx); + if (vi == num_full) { + const auto base = vi * kVecSize; + if (base == seq_len) return; + vec_t cur; + cur.load(in, vi); +#pragma unroll + for (uint32_t j = 0; j < kVecSize; ++j) { + if (base + j < seq_len) fn(cur[j], base + j); + } } } - SGL_DEVICE static void find_threshold(const uint32_t topk, const uint32_t seq_len, Smem* smem) { + SGL_DEVICE static void init_histogram(uint32_t (&histogram)[kHistSize], uint32_t tx) { + constexpr uint32_t kItems = kHistSize / kBlockSize; + AlignedVector vec; + vec.fill(0); + vec.store(histogram, tx); + } + + /// Same, but scanning a histogram that need not be `smem`'s own -- the cluster + /// path merges into one rank's copy and scans it there. + template + SGL_DEVICE static void find_threshold(const uint32_t topk, const uint32_t seq_len, Smem* smem, Fn fn) { const auto tx = threadIdx.x; constexpr uint32_t kItems = kHistSize / kBlockSize; - uint32_t orig[kItems]; - const auto hist_vec = smem->hist_vecs[tx]; - uint32_t tmp_local_sum = 0; + uint32_t local_exc_sum[kItems + 1]; + AlignedVector hist_vec; + hist_vec.load(smem->histogram, tx); + local_exc_sum[0] = 0; #pragma unroll for (uint32_t i = 0; i < kItems; ++i) { - orig[i] = hist_vec[i]; - tmp_local_sum += orig[i]; + local_exc_sum[i + 1] = hist_vec[i] + local_exc_sum[i]; } + const auto local_sum = local_exc_sum[kItems]; const auto lane_id = tx % kWarpSize; - const auto warp_id = tx / kWarpSize; - const auto warp_inc = warp_inclusive_sum(lane_id, tmp_local_sum); - const auto warp_exc = warp_inc - tmp_local_sum; - if (lane_id == kWarpSize - 1) smem->warp_sum[warp_id] = warp_inc; + const auto warp_id = broadcast(tx / kWarpSize); + const auto warp_inc_sum = warp_inclusive_sum(lane_id, local_sum); + const auto warp_exc_sum = warp_inc_sum - local_sum; + if (lane_id == kWarpSize - 1) smem->warp_sum[warp_id] = warp_inc_sum; __syncthreads(); const auto tmp = smem->warp_sum[lane_id]; - // Exactly one bin satisfies: above < K && above + count >= K - uint32_t prefix_sum = warp::reduce_sum(lane_id < warp_id ? tmp : 0); - prefix_sum += warp_exc; + const auto warp_prefix_sum = warp::reduce_sum(lane_id < warp_id ? tmp : 0); + const auto exc_sum = static_cast(warp_prefix_sum + warp_exc_sum); + const auto remained = static_cast(seq_len - topk - exc_sum); + // only 1 lane will execute this + if (remained >= 0 && remained < static_cast(local_sum)) [[unlikely]] { + uint32_t target = 0; #pragma unroll - for (uint32_t i = 0; i < kItems; ++i) { - prefix_sum += orig[i]; - const auto above = seq_len - prefix_sum; - if (above < topk && above + orig[i] >= topk) { - smem->threshold_bin = tx * kItems + i; + for (uint32_t i = 0; i < kItems; ++i) { + const auto prev = static_cast(local_exc_sum[i + 0]); + const auto next = static_cast(local_exc_sum[i + 1]); + if (remained >= prev && remained < next) target = tx * kItems + i; } + fn(target); } + __syncthreads(); } }; @@ -542,15 +556,11 @@ struct TopKRegister : TopKRadixBase<12> { using Smem = typename TopKRadixBase<12>::Smem; template - SGL_DEVICE static void forward(const TopKProblem problem, void* _smem) { + SGL_DEVICE static void forward(const TopKProblem& problem, void* _smem) { const auto tx = threadIdx.x; const auto smem = static_cast(_smem); - { - Smem::kHistVec hist_vec; - hist_vec.fill(0); - smem->hist_vecs[tx] = hist_vec; - } + init_histogram(smem->histogram, tx); if (tx == 0) { smem->count_eq = 0; smem->count_gt = 0; @@ -558,78 +568,83 @@ struct TopKRegister : TopKRadixBase<12> { __syncthreads(); PDLWaitPrimary(); - - // A vector `vi` is fully in bounds iff vi < num_full; only full vectors are - // vector-loaded (16B aligned, never straddling seq_len). The = num_full) break; - local_vecs[i].load(problem.in, vi); + if (vi < num_full) local_vecs[i].load(problem.in, vi); } + + const auto tail_start = (problem.seq_len - 1) % kVecSize + 1; #pragma unroll for (uint32_t i = 0; i < kLocalVecs; ++i) { const auto vi = tx + kBlockSize * i; if (vi >= num_full) break; + if (vi == num_full - 1) { #pragma unroll - for (uint32_t j = 0; j < kVecSize; ++j) + for (uint32_t j = 0; j < kVecSize; ++j) { + if (j >= tail_start) local_vecs[i][j] = padding_value(); + } + } +#pragma unroll + for (uint32_t j = 0; j < kVecSize; ++j) { atomicAdd(&smem->histogram[extract_coarse_bin(local_vecs[i][j])], 1); + } } - if (tx >= kBlockSize - tail) { - const uint32_t idx = tail_start + tx - (kBlockSize - tail); - atomicAdd(&smem->histogram[extract_coarse_bin(problem.in[idx])], 1); + const auto num_padding = kVecSize - tail_start + problem.input_start; + if (tx == 0 && num_padding > 0) { + // Ask the histogram's own binning where the platform's NaN landed. + atomicSub(&smem->histogram[extract_coarse_bin(padding_value())], num_padding); + atomicAdd(&smem->histogram[0], num_padding); } __syncthreads(); // Phase 2: Find the threshold bin - find_threshold(problem.topk, problem.seq_len, smem); + find_threshold(problem.topk, num_full * kVecSize, smem, [&](uint32_t threshold_bin) { + const auto v_hi = coarse_bin_lower_bound(threshold_bin + 1); + const auto v_lo = coarse_bin_lower_bound(threshold_bin + 0); + smem->v_hi = v_hi; + smem->v_lo = v_lo; + }); - // Phase 3: collect by two fp32 boundaries (raw indices; transform applied later) + // Phase 3: collect by two fp32 boundaries const auto topk = problem.topk; - const auto threshold_bin = smem->threshold_bin; - const auto v_hi = coarse_bin_lower_bound(threshold_bin + 1); - const auto v_lo = coarse_bin_lower_bound(threshold_bin); - const auto collect = [&](float val, uint32_t idx) { - if (val >= v_hi) { - const auto pos = atomicAdd(&smem->count_gt, 1); - if (pos < topk) [[likely]] - problem.emit(pos, idx); - } else if (val >= v_lo) { - const auto count_eq = atomicAdd(&smem->count_eq, 1); - if (count_eq < kMaxNumTie) [[likely]] - smem->tie.values[count_eq] = {val, idx}; - } - }; + const auto v_hi = smem->v_hi; + const auto v_lo = smem->v_lo; + #pragma unroll for (uint32_t i = 0; i < kLocalVecs; ++i) { const auto vi = tx + kBlockSize * i; const auto base = vi * kVecSize; if (vi >= num_full) break; #pragma unroll - for (uint32_t j = 0; j < kVecSize; ++j) - collect(local_vecs[i][j], base + j); - } - if (tx >= kBlockSize - tail) { - const uint32_t idx = tail_start + tx - (kBlockSize - tail); - collect(problem.in[idx], idx); + for (uint32_t j = 0; j < kVecSize; ++j) { + const auto idx = base + j; + const auto val = local_vecs[i][j]; + if (val >= v_hi) { + const auto pos = atomicAdd(&smem->count_gt, 1); + if (pos < topk) [[likely]] { + problem.emit(pos, idx); + } + } else if (val >= v_lo) { + const auto pos = atomicAdd(&smem->count_eq, 1); + if (pos < kMaxNumTie) [[likely]] { + smem->tie_values[pos] = {val, idx}; + } + } + } } // Phase 4: Handle ties. __syncthreads(); - const auto above_count = smem->count_gt; - const auto equal_count = smem->count_eq; - const auto remain_topk = above_count < topk ? topk - above_count : 0; - const auto tie_count = min(equal_count, kMaxNumTie); - handle_tie(smem->tie.values, problem, above_count, tie_count, remain_topk, &smem->tie.handle); + const auto count_gt = smem->count_gt; + const auto count_eq = smem->count_eq; + const auto remain_topk = count_gt < topk ? topk - count_gt : 0; + const auto tie_count = min(count_eq, kMaxNumTie); + handle_tie(smem->tie_values, problem, count_gt, tie_count, remain_topk, &smem->tie_handle); } }; @@ -637,20 +652,16 @@ struct TopKRegister : TopKRadixBase<12> { // Streaming path: seq_len > 8192 -- two vectorized passes over global memory // --------------------------------------------------------------------------- -struct TopKStreaming : TopKRegister<2> { +struct TopKStreaming : TopKRadixBase<12> { public: static constexpr uint32_t kMaxSeqLen = std::numeric_limits::max(); template - SGL_DEVICE static void forward(const TopKProblem problem, void* _smem) { + SGL_DEVICE static void forward(TopKProblem problem, void* _smem) { const auto tx = threadIdx.x; const auto smem = static_cast(_smem); - { - Smem::kHistVec hist_vec; - hist_vec.fill(0); - smem->hist_vecs[tx] = hist_vec; - } + init_histogram(smem->histogram, tx); if (tx == 0) { smem->count_eq = 0; smem->count_gt = 0; @@ -663,19 +674,28 @@ struct TopKStreaming : TopKRegister<2> { const auto bin = extract_coarse_bin(val); atomicAdd(&smem->histogram[bin], 1); }); + const auto num_padding = problem.input_start; + if (tx == 0 && num_padding != 0) { + atomicSub(&smem->histogram[extract_coarse_bin(padding_value())], num_padding); + atomicAdd(&smem->histogram[0], num_padding); + } __syncthreads(); // Phase 2: Find the threshold bin - find_threshold(problem.topk, problem.seq_len, smem); + find_threshold(problem.topk, problem.seq_len, smem, [&](uint32_t threshold_bin) { + const auto v_hi = coarse_bin_lower_bound(threshold_bin + 1); + const auto v_lo = coarse_bin_lower_bound(threshold_bin + 0); + smem->v_hi = v_hi; + smem->v_lo = v_lo; + }); // Phase 3: Collect candidates and sort. Classify by two fp32 boundaries derived // from the threshold bin instead of recomputing the fp16 bin per element: an // element is "above" iff val >= v_hi (bin > threshold) and a "tie" iff // v_lo <= val < v_hi (bin == threshold). This drops the F2F + bit-twiddle from // the second full pass over the input. - const auto threshold_bin = smem->threshold_bin; - const float v_hi = coarse_bin_lower_bound(threshold_bin + 1); - const float v_lo = coarse_bin_lower_bound(threshold_bin); + const auto v_hi = smem->v_hi; + const auto v_lo = smem->v_lo; const auto topk = problem.topk; for_each_input(problem.in, problem.seq_len, [&](float val, uint32_t idx) { if (val >= v_hi) { @@ -684,9 +704,9 @@ struct TopKStreaming : TopKRegister<2> { problem.emit(pos, idx); } } else if (val >= v_lo) { - const auto count_eq = atomicAdd(&smem->count_eq, 1); - if (count_eq < kMaxNumTie) [[likely]] { - smem->tie.values[count_eq] = {val, idx}; + const auto pos = atomicAdd(&smem->count_eq, 1); + if (pos < kMaxNumTie) [[likely]] { + smem->tie_values[pos] = {val, idx}; } } }); @@ -697,11 +717,11 @@ struct TopKStreaming : TopKRegister<2> { // "above" and "tie" sets. above_count is < topk by the threshold-bin invariant, // so the count_gt guard above effectively never triggers. __syncthreads(); - const auto above_count = smem->count_gt; - const auto equal_count = smem->count_eq; - const auto remain_topk = above_count < topk ? topk - above_count : 0; - const auto tie_count = min(equal_count, kMaxNumTie); - handle_tie(smem->tie.values, problem, above_count, tie_count, remain_topk, &smem->tie.handle); + const auto count_gt = smem->count_gt; + const auto count_eq = smem->count_eq; + const auto remain_topk = count_gt < topk ? topk - count_gt : 0; + const auto tie_count = min(count_eq, kMaxNumTie); + handle_tie(smem->tie_values, problem, count_gt, tie_count, remain_topk, &smem->tie_handle); } }; @@ -713,171 +733,175 @@ struct TopKStreaming : TopKRegister<2> { // equivalent. // --------------------------------------------------------------------------- -#ifndef USE_ROCM +#if SUPPORT_CLUSTER -template +template struct TopKCluster : TopKRadixBase<10> { public: - static constexpr uint32_t kClusterSize = kClusterSize_; + static constexpr uint32_t kClusterSize = N; static constexpr uint32_t kMaxSeqLen = std::numeric_limits::max(); - using Base = TopKRadixBase<10>; - struct Smem : Base::Smem { - using kHistVec = Base::Smem::kHistVec; - uint32_t start_eq_local, start_gt_local; - int32_t tmp_out[kMaxTopK]; + struct Smem { + uint32_t count_eq; + uint32_t count_gt; + uint32_t local_start_eq; + uint32_t local_start_gt; + float v_lo; + float v_hi; + uint32_t warp_sum[kNumWarps]; + union { + alignas(16) uint32_t histogram[kHistSize]; + TieHandleSmem tie_handle; + int32_t stage_out_idxs[kMaxTopK]; + }; + TieValue tie_values[kMaxNumTie]; }; - // Process ONE batch element (one cluster). NO PDL and NO trailing barrier -- - // the persistent kernel does PDLWaitPrimary once before its item loop and a - // cluster.sync() after each forward(). Writes raw indices to out; the kernel's - // transform pass applies the page-table transform. + SGL_DEVICE static void barrier_cluster_arrive_relaxed() { + asm volatile("barrier.cluster.arrive.relaxed.aligned;" ::: "memory"); + } + + SGL_DEVICE static void barrier_cluster_arrive_release() { + asm volatile("barrier.cluster.arrive.release.aligned;" ::: "memory"); + } + + SGL_DEVICE static void barrier_cluster_wait() { + asm volatile("barrier.cluster.wait.acquire.aligned;" ::: "memory"); + } + template SGL_DEVICE static void forward(TopKProblem problem, void* _smem) { const auto tx = threadIdx.x; const auto smem = static_cast(_smem); - const auto cluster = cg::this_cluster(); + const auto cluster = cooperative_groups::this_cluster(); const auto this_rank = blockIdx.y; - const bool is_primary = (this_rank == 0); - constexpr uint32_t kAlignElems = kWarpSize * kVecSize; - const uint32_t chunk_size = div_ceil(problem.seq_len, kClusterSize * kAlignElems) * kAlignElems; - const uint32_t chunk_start = min(this_rank * chunk_size, problem.seq_len); - const uint32_t chunk_finish = min(chunk_start + chunk_size, problem.seq_len); - const uint32_t local_seq_len = chunk_finish - chunk_start; - problem.in += chunk_start; - - { - typename Smem::kHistVec hist_vec; - hist_vec.fill(0); - smem->hist_vecs[tx] = hist_vec; - } + init_histogram(smem->histogram, tx); if (tx == 0) { smem->count_eq = 0; smem->count_gt = 0; } __syncthreads(); + // Rank 0's shared memory is read by its peers: the zeroed histogram they fold + // into after bar-0, v_hi / v_lo after bar-2. Those arrives release so the + // peers' acquire wait orders the reads after the writes at cluster scope; + // __syncthreads() alone is CTA-scoped. The peers publish nothing at these two + // barriers and keep the cheaper relaxed arrive. + if (this_rank == 0) { + barrier_cluster_arrive_release(); // bar-0 arrive + } else { + barrier_cluster_arrive_relaxed(); // bar-0 arrive + } PDLWaitPrimary(); // Phase 1: Load and build histogram over this rank's contiguous chunk. - for_each_input(problem.in, local_seq_len, [&](float val, uint32_t) { + for_each_input(problem.in, problem.seq_len, [&](float val, uint32_t) { const auto bin = extract_coarse_bin(val); atomicAdd(&smem->histogram[bin], 1); }); + + barrier_cluster_wait(); // bar-0 wait __syncthreads(); + if (this_rank != 0) { + const auto smem_0 = cluster.map_shared_rank(smem, 0); + // Phase 2. atomic flush all histogram into rank 0 + static_assert(kHistSize == kBlockSize); // one bin per thread - // Phase 1.5: reduce the histogram across the cluster - { - // 1-shot all-reduce: each rank owns kPartition consecutive bins; - // for each owned bin, gather the kClusterSize peer values (one per - // consecutive lane) via DSMEM, sum across the lanes, then scatter back. - cluster.sync(); - static_assert(kHistSize == kBlockSize); // we optimize on top of this - constexpr uint32_t kPartition = kHistSize / kClusterSize; - const auto start = this_rank * kPartition; - const auto which = start + tx / kClusterSize; - const auto peer_rank = tx % kClusterSize; - const auto addr = cluster.map_shared_rank(&smem->histogram[which], peer_rank); - const auto value = *addr; - *addr = warp::reduce_sum(value); - cluster.sync(); - } + if (const auto count = smem->histogram[tx]; count != 0) { + atomicAdd(&smem_0->histogram[tx], count); + } - // Phase 2: Find the threshold bin (uses global seq_len) - find_threshold(problem.topk, problem.seq_len, smem); + barrier_cluster_arrive_release(); // bar-1 arrive + barrier_cluster_wait(); // bar-1 wait - // Phase 3: Collect candidates over this rank's chunk; convert local indices - // back to global by adding chunk_start. Classify by two fp32 boundaries derived - // from the (global) threshold bin instead of recomputing the fp16 bin per - // element -- see TopKStreaming for the rationale. threshold_bin is identical - // across ranks, so v_hi/v_lo are too. - const auto topk = problem.topk; - const auto threshold_bin = smem->threshold_bin; - const float v_hi = coarse_bin_lower_bound(threshold_bin + 1); - const float v_lo = coarse_bin_lower_bound(threshold_bin); + barrier_cluster_arrive_relaxed(); // bar-2 arrive + barrier_cluster_wait(); // bar-2 wait - // Phase 3: collect candidates. The primary scatters straight into - // `problem.out`, the others stage into block-local `smem->tmp_out`. - // - // DO NOT merge these two loops back into one by selecting the destination - // first (`cur_out = is_primary ? problem.out : smem->tmp_out`). `problem.out` - // can be a shared::cluster (DSMEM) alias of the elected rank's buffer while - // `tmp_out` is shared::cta; merging them into a single pointer variable makes - // cicc 13.1+ mis-lower the block-local arm on sm_90a and *silently drop every - // non-primary rank's staged output* -- `tmp_out` stays zero, and phase 3.5 - // then faithfully copies zeros to correct DSMEM addresses. The result is a - // top-k output where only the primary's slots and the handle_tie tail are - // valid, which downstream sparse attention dereferences as garbage KV indices. - if (!is_primary) { - // stage to tmp_out first before writing to global/DSMEM - for_each_input(problem.in, local_seq_len, [&](float val, uint32_t local_idx) { - const auto idx = chunk_start + local_idx; + // Phase 4. non-0 rank stage to local smem, then write to rank-0 via DSMEM + const auto topk = problem.topk; + const auto v_hi = smem_0->v_hi; + const auto v_lo = smem_0->v_lo; + for_each_input(problem.in, problem.seq_len, [&](float val, uint32_t idx) { if (val >= v_hi) { const auto pos = atomicAdd(&smem->count_gt, 1); if (pos < topk) [[likely]] { - smem->tmp_out[pos] = idx; + smem->stage_out_idxs[pos] = idx; } } else if (val >= v_lo) { - const auto count_eq = atomicAdd(&smem->count_eq, 1); - if (count_eq < kMaxNumTie) [[likely]] { - smem->tie.values[count_eq] = {val, idx}; + const auto pos = atomicAdd(&smem->count_eq, 1); + if (pos < kMaxNumTie) [[likely]] { + smem->tie_values[pos] = {val, idx}; } } }); __syncthreads(); - const auto local_above_count = smem->count_gt; - const auto local_equal_count = min(smem->count_eq, kMaxNumTie); - const auto smem_0 = cluster.map_shared_rank(smem, 0); + const auto local_count_gt = smem->count_gt; + const auto local_count_eq = min(smem->count_eq, kMaxNumTie); if (tx == 0) { - const auto gt = atomicAdd(&smem_0->count_gt, local_above_count); - const auto eq = atomicAdd(&smem_0->count_eq, local_equal_count); - smem->start_gt_local = gt; - smem->start_eq_local = eq; + const auto gt = atomicAdd(&smem_0->count_gt, local_count_gt); + const auto eq = atomicAdd(&smem_0->count_eq, local_count_eq); + smem->local_start_gt = gt; + smem->local_start_eq = eq; } __syncthreads(); - const auto start_gt_local = smem->start_gt_local; - const auto start_eq_local = smem->start_eq_local; + const auto local_start_gt = smem->local_start_gt; + const auto local_start_eq = smem->local_start_eq; #pragma unroll for (uint32_t i = 0; i < kTieItems; ++i) { const auto t = tx + i * kBlockSize; - if (t < local_equal_count && start_eq_local + t < kMaxNumTie) { - smem_0->tie.values[start_eq_local + t] = smem->tie.values[t]; + if (t < local_count_eq && local_start_eq + t < kMaxNumTie) { + smem_0->tie_values[local_start_eq + t] = smem->tie_values[t]; } } - cluster.sync(); + cluster.sync(); // bar 3 - const auto start_write = start_gt_local; - const auto num_write = local_above_count; + const auto start_write = local_start_gt; + const auto num_write = local_count_gt; #pragma unroll for (uint32_t i = 0; i < kTopKItems; ++i) { if (const auto t = tx + i * kBlockSize; t < num_write && start_write + t < topk) { - problem.emit(start_write + t, smem->tmp_out[t]); + problem.emit(start_write + t, smem->stage_out_idxs[t]); } } } else { - for_each_input(problem.in, local_seq_len, [&](float val, uint32_t local_idx) { - const auto idx = chunk_start + local_idx; + barrier_cluster_arrive_relaxed(); // bar-1 arrive + barrier_cluster_wait(); // bar-1 wait + + // Phase 3. rank-0 find threshold and write to local smem for other ranks to read + find_threshold(problem.topk, problem.seq_len, smem, [&](uint32_t threshold_bin) { + smem->v_hi = coarse_bin_lower_bound(threshold_bin + 1); + smem->v_lo = coarse_bin_lower_bound(threshold_bin + 0); + }); + + barrier_cluster_arrive_release(); // bar-2 arrive: publishes v_hi / v_lo + barrier_cluster_wait(); // bar-2 wait + + // Phase 4. rank-0 directly write to output + const auto topk = problem.topk; + const auto v_hi = smem->v_hi; + const auto v_lo = smem->v_lo; + for_each_input(problem.in, problem.seq_len, [&](float val, uint32_t idx) { if (val >= v_hi) { const auto pos = atomicAdd(&smem->count_gt, 1); if (pos < topk) [[likely]] { problem.emit(pos, idx); } } else if (val >= v_lo) { - const auto count_eq = atomicAdd(&smem->count_eq, 1); - if (count_eq < kMaxNumTie) [[likely]] { - smem->tie.values[count_eq] = {val, idx}; + const auto pos = atomicAdd(&smem->count_eq, 1); + if (pos < kMaxNumTie) [[likely]] { + smem->tie_values[pos] = {val, idx}; } } }); - cluster.sync(); + cluster.sync(); // bar-3 // Phase 4: Handle ties. - const auto above_count = smem->count_gt; - const auto equal_count = smem->count_eq; - const auto remain_topk = above_count < topk ? topk - above_count : 0; - const auto tie_count = min(equal_count, kMaxNumTie); - handle_tie(smem->tie.values, problem, above_count, tie_count, remain_topk, &smem->tie.handle); + const auto count_gt = smem->count_gt; + const auto count_eq = smem->count_eq; + const auto remain_topk = count_gt < topk ? topk - count_gt : 0; + const auto tie_count = min(count_eq, kMaxNumTie); + handle_tie(smem->tie_values, problem, count_gt, tie_count, remain_topk, &smem->tie_handle); } } }; diff --git a/python/sglang/kernels/jit/utils/occupancy.py b/python/sglang/kernels/jit/utils/occupancy.py new file mode 100644 index 000000000..bdae2169b --- /dev/null +++ b/python/sglang/kernels/jit/utils/occupancy.py @@ -0,0 +1,47 @@ +"""Occupancy probes a host-side dispatch needs before it can size a grid.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +from sglang.kernels.jit.utils.common import cache_once +from sglang.kernels.jit.utils.compile import load_jit + +if TYPE_CHECKING: + from tvm_ffi.module import Module + +__all__ = ["get_max_active_clusters"] + + +@cache_once +def _jit_probe_module() -> Module: + return load_jit( + "occupancy_cluster_probe", + cuda_files=["occupancy/cluster_probe.cuh"], + cuda_wrappers=[("get_max_active_clusters", "get_max_active_clusters")], + ) + + +@cache_once +def _get_max_active_clusters(cluster_size: int, occupancy: int) -> int: + return int(_jit_probe_module().get_max_active_clusters(cluster_size, occupancy)) + + +def get_max_active_clusters(cluster_size: int, occupancy: int) -> int: + """Clusters of ``cluster_size`` blocks that can be resident at once. + + Asks the driver (``cudaOccupancyMaxActiveClusters``) rather than dividing SM + count by cluster size: a cluster's blocks must be co-scheduled within one + GPC, so the answer falls short of the division once the cluster stops + dividing a GPC evenly. The probe kernel is pinned to ``occupancy`` blocks per + SM, so pass the occupancy the real kernel reaches (its second + ``__launch_bounds__`` argument). Raises ``RuntimeError`` before sm90, which + has no clusters, and ``ValueError`` when nothing is schedulable. + """ + result = _get_max_active_clusters(cluster_size, occupancy) + if result == 0: + raise ValueError( + f"no cluster of {cluster_size} fits at occupancy {occupancy}; " + "the cluster width is likely beyond what this device supports" + ) + return result diff --git a/python/sglang/kernels/ops/attention/dsv4/candidate_table.py b/python/sglang/kernels/ops/attention/dsv4/candidate_table.py new file mode 100644 index 000000000..e956d2e69 --- /dev/null +++ b/python/sglang/kernels/ops/attention/dsv4/candidate_table.py @@ -0,0 +1,93 @@ +"""Candidate-table kernels of the two-level indexer: level-one block keys and +the sorted, page-transformed block table.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Optional + +import torch + +from sglang.kernels.jit.utils import ( + cache_once, + is_arch_support_pdl, + load_jit, + make_cpp_args, +) + +from .utils import make_name + +if TYPE_CHECKING: + pass + + +@cache_once +def _jit_block_amax_module(): + args = make_cpp_args(is_arch_support_pdl()) + return load_jit( + make_name("block_amax"), + *args, + cuda_files=["deepseek_v4/block_amax.cuh"], + cuda_wrappers=[("amax8_varlen", f"BlockAmaxKernel<{args}>::amax8_varlen")], + ) + + +def amax8_varlen( + scores: torch.Tensor, + seq_lens: torch.Tensor, + topk: int = 0, + *, + max_seqlen: int = 0, + out: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """Level-one keys of the two-level indexer: ``out[b, i]`` is the max of + ``scores[b, 8 i : 8 i + 8]`` for ``i < ceil(seq_lens[b] / 8)``, the last of + them ``+inf`` (the newest block is always selected), nothing written past + that count. Rows with at most ``topk`` blocks are skipped (every block is + selected anyway); ``topk=0`` never skips. ``out`` is allocated as + ``[rows, ceil(max_seqlen / 8)]`` when not given, ``max_seqlen`` defaulting to + the width of ``scores``; every ``seq_lens[b]`` must fit in ``8 * out.shape[1]``. + fp32 only for now; ``scores`` rows must be 32-byte aligned (stride a multiple + of 8). Returns ``out``. + """ + if out is None: + num_tokens, max_len = scores.shape + if max_seqlen == 0: + max_seqlen = max_len + out = scores.new_empty(num_tokens, (max_seqlen + 7) // 8) + _jit_block_amax_module().amax8_varlen(scores, seq_lens, out, topk) + return out + + +@cache_once +def _jit_candidate_block_table_module(): + args = make_cpp_args(is_arch_support_pdl()) + return load_jit( + make_name("candidate_block_table"), + *args, + cuda_files=["deepseek_v4/candidate_block_table.cuh"], + cuda_wrappers=[("transform", f"CandidateBlockTableKernel<{args}>::transform")], + ) + + +def sort_candidate_blocks( + blocks: torch.Tensor, + seq_lens: torch.Tensor, + page_table: torch.Tensor, + page_size: int, + *, + out_pages: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """The block table of the two-level indexer from a row's selected blocks, + in place: ``blocks`` ``[rows, k]`` int32 block ids in any order, ``-1`` + padded, become the same ids ascending with ``INT32_MAX`` past ``min(k, + ceil(seq_lens[b] / 8))``; the matching pool slots / 8 (``page_table[b, id // + bpp] * bpp + id % bpp``, ``bpp = page_size // 8``, same padding) go to + ``out_pages``. A row with at most ``k`` blocks gets the identity table + regardless of its input. Returns ``out_pages``. + """ + if out_pages is None: + out_pages = torch.empty_like(blocks) + _jit_candidate_block_table_module().transform( + blocks, seq_lens, page_table, out_pages, page_size + ) + return out_pages diff --git a/python/sglang/kernels/ops/attention/dsv4/topk.py b/python/sglang/kernels/ops/attention/dsv4/topk.py index e5f0d0fd1..1ed069eff 100644 --- a/python/sglang/kernels/ops/attention/dsv4/topk.py +++ b/python/sglang/kernels/ops/attention/dsv4/topk.py @@ -18,11 +18,6 @@ from .utils import make_name @cache_once def _jit_topk_v1_module(): - # topk (<= 1024) is a runtime argument, not a compile-time constant, so a - # single module serves every k. Baking it in via -DSGL_TOPK used to build one - # module per k, and since the macro fed a `constexpr` rather than a template - # parameter every module exported identically mangled symbols -- see the - # comment in topk_v1.cuh for how that broke the second module's launch. args = make_cpp_args(is_arch_support_pdl()) return load_jit( make_name("topk_v1"), @@ -34,19 +29,75 @@ def _jit_topk_v1_module(): @cache_once def _jit_topk_v2_module(): - # v2 is universal: topk (<= 2048) is a runtime argument, not a compile-time - # constant, so a single module serves every k. + from sglang.kernels.jit.utils.occupancy import get_max_active_clusters + + args = make_cpp_args(is_arch_support_pdl()) + # Leave these undefined if the probe fails: topk_v2.cuh carries per-arch + # defaults, and a 0 would size the persistent pool to an empty grid. + extra_cuda_cflags = [] + if is_arch_support_pdl(): # set the persistent cluster size after hopper + try: + occ_8_2 = get_max_active_clusters(8, occupancy=2) + except Exception: + pass + else: + if occ_8_2 > 0: + extra_cuda_cflags.append(f"-DSGL_TOPK_V2_MAX_C8_OCC2={occ_8_2}") + try: + occ_16_1 = get_max_active_clusters(16, occupancy=1) + except Exception: + pass + else: + if occ_16_1 > 0: + extra_cuda_cflags.append(f"-DSGL_TOPK_V2_MAX_C16_OCC1={occ_16_1}") + kernel = f"TopKKernel<{args}>" return load_jit( make_name("topk_v2"), + *args, + extra_cuda_cflags=extra_cuda_cflags, cuda_files=["deepseek_v4/topk_v2.cuh"], cuda_wrappers=[ - ("topk_transform_paged", "TopKKernel::transform_paged"), - ("topk_transform_ragged", "TopKKernel::transform_ragged"), - ("topk_plan", "TopKKernel::plan"), + ("topk_transform_paged", f"{kernel}::transform_paged"), + ("topk_transform_ragged", f"{kernel}::transform_ragged"), + ("topk_plan", f"{kernel}::plan"), ], ) +@cache_once +def _jit_topk_bf16_small_module(): + args = make_cpp_args(is_arch_support_pdl()) + return load_jit( + make_name("topk_bf16_small"), + *args, + cuda_files=["deepseek_v4/topk_bf16_small.cuh"], + cuda_wrappers=[("topk_transform", f"TopKBF16Kernel<{args}>::transform")], + ) + + +def topk_transform_bf16_small( + scores: torch.Tensor, + seq_lens: torch.Tensor, + page_table: torch.Tensor, + out_page_indices: torch.Tensor, + page_size: int, +) -> None: + """bf16 top-k for rows of at most 16384 scores (the DeepSeek-V4.1 sparse + indexer's consumer rows), fused with a page-table transform. + + Row ``b`` selects the ``k = out_page_indices.shape[1]`` best of its first + ``seq_lens[b]`` scores (``k`` at most 2048); a selected index ``i`` is + written as ``page_table[b, i // page_size] * page_size + i % page_size``, + in no particular order, and ``-1`` fills the slots past + ``min(k, seq_lens[b])``. Selection is exact (two radix passes over the raw + bf16 bytes locate the k-th largest value); which of the elements equal to + it fill the last slots is arbitrary. NaN scores are not supported. + """ + _jit_topk_bf16_small_module().topk_transform( + scores, seq_lens, page_table, out_page_indices, page_size + ) + + def topk_transform_paged( scores: torch.Tensor, seq_lens: torch.Tensor, @@ -75,15 +126,14 @@ def topk_transform_paged( _PLAN_METADATA_INTS_PER_BATCH = 2 -def plan_topk_v2(seq_lens: torch.Tensor, static_threshold: int = 0) -> torch.Tensor: - """Preprocess the per-batch routing plan for :func:`topk_transform_paged_v2`. +def plan_topk_v2(seq_lens: torch.Tensor, static_threshold: int = -1) -> torch.Tensor: + """ + Preprocess the per-batch routing plan for :func:`topk_transform_paged_v2`. + NOTE: every entry of ``seq_lens`` must be NON-NEGATIVE. - IMPORTANT: every entry of ``seq_lens`` must be NON-NEGATIVE. The device - kernel reads the int32 buffer as ``uint32_t``, so a negative length (e.g. - -4 from a DP-padded / idle-companion row) reinterprets as ~4e9, poisons - the plan, and drives the transform kernel into an illegal memory access. - Producers of padded rows must clamp their lengths to 0 (0 selects the - trivial all-(-1) output path, which is safe). + :param static_threshold: If a batch item has `seq_len` > `static_threshold`, + prefer the cluster implementation. + Negative number means internal heuristic. """ module = _jit_topk_v2_module() bs = seq_lens.shape[0] @@ -92,6 +142,19 @@ def plan_topk_v2(seq_lens: torch.Tensor, static_threshold: int = 0) -> torch.Ten return metadata +def topk_v2_plan_is_written(seq_lens: torch.Tensor) -> bool: + """Whether :func:`plan_topk_v2` writes a plan for these lengths. Small + batches and devices without clusters leave the plan buffer untouched.""" + probe = torch.full( + (seq_lens.shape[0] + 1, _PLAN_METADATA_INTS_PER_BATCH), + -1, + dtype=torch.int32, + device=seq_lens.device, + ) + _jit_topk_v2_module().topk_plan(seq_lens, probe, -1) + return probe[0, 1].item() != -1 + + def topk_transform_ragged_v2( scores: torch.Tensor, seq_lens: torch.Tensor, @@ -111,7 +174,7 @@ def topk_transform_ragged_v2( Unlike :func:`topk_transform_paged_v2` this needs no page table and no plan (the cluster path only pays off for very few rows, and prefill has many). - IMPORTANT: ``scores`` is written in place -- the <= 3 columns ahead of each + NOTE: ``scores`` is written in place -- the <= 3 columns ahead of each row's window that the 16-byte-aligned read base pulls in are masked out. They are invalid for that row and the buffer must have no other consumer. ``seq_lens`` entries must be NON-NEGATIVE, as for the paged entry point. @@ -151,14 +214,10 @@ def topk_transform_paged_v2( * Both outputs given -- ``out_page_indices`` receives the page-table transform and ``out_raw_indices`` receives the selected raw indices. - IMPORTANT: every entry of ``seq_lens`` must be NON-NEGATIVE, and - ``metadata`` must come from :func:`plan_topk_v2` over the same ``seq_lens`` - values. The kernel reads lengths as ``uint32_t``: a negative entry - reinterprets as a ~4e9-token sequence, sending the row down the cluster - path over garbage scores and crashing with an illegal memory access - (GLM 5.2 MTP DP-idle companion rows hit exactly this). A length of 0 is - the valid way to express "no tokens": the row takes the trivial path and - the output is all -1. + NOTE: every entry of `seq_lens` must be NON-NEGATIVE, and `metadata` must + come from :func:`plan_topk_v2` over the same `seq_lens` values. + A length of 0 is the valid way to express "no tokens": the row takes the + trivial path and the output is guaranteed to be all -1. """ if is_xpu(): if out_raw_indices is not None: diff --git a/python/sglang/test/kits/dsa_metadata_kit.py b/python/sglang/test/kits/dsa_metadata_kit.py index a3e5c3629..45de95799 100644 --- a/python/sglang/test/kits/dsa_metadata_kit.py +++ b/python/sglang/test/kits/dsa_metadata_kit.py @@ -5,6 +5,7 @@ from types import SimpleNamespace import torch +from sglang.kernels.ops.attention.dsv4.topk import topk_v2_plan_is_written from sglang.srt.environ import envs from sglang.srt.layers.attention.dsa.dsa_topk_backend import DSATopKBackend from sglang.srt.layers.attention.dsa_backend import DeepseekSparseAttnBackend @@ -101,6 +102,8 @@ def assert_metadata_equal(test, actual, expected): for name, value in actual_buffers.items(): reference = expected_buffers[name] if name == "topk_v2_plan": + if not topk_v2_plan_is_written(expected.dsa_seqlens_expanded): + continue # Unused plan rows are intentionally uninitialized. Active rows are # compacted by atomicAdd, so compare them in request order. torch.testing.assert_close(value[0], reference[0]) diff --git a/test/registered/kernels/benchmark/attention/bench_topk.py b/test/registered/kernels/benchmark/attention/bench_topk.py index a1bbb9815..c76b2e78d 100644 --- a/test/registered/kernels/benchmark/attention/bench_topk.py +++ b/test/registered/kernels/benchmark/attention/bench_topk.py @@ -104,14 +104,16 @@ if not DISABLE_TORCH: PRROVIDERS.append("torch") +@marker.parametrize("page_size", [1, 64], [1, 64]) @marker.parametrize("k", [512, 1024, 2048], [512]) @marker.parametrize("seq_len", [2**x for x in range(10, 19)], [4096, 65536]) @marker.parametrize("batch_size", [2**x for x in range(13)], [1, 128, 1024]) -@marker.parametrize("page_size", [1, 64], [1, 64]) @marker.benchmark("provider", PRROVIDERS) def benchmark_paged( seq_len: int, batch_size: int, k: int, page_size: int, provider: str ): + seed = seq_len ^ (batch_size << 16) ^ (k << 32) ^ (page_size << 48) + torch.random.manual_seed(seed) if k > seq_len: marker.skip("k cannot be larger than seq_len") if k == 2048 and provider == "jit_v1": @@ -127,6 +129,8 @@ def benchmark_paged( @marker.parametrize("batch_size", [2**x for x in range(7, 14)], [128, 1024]) @marker.benchmark("provider", PRROVIDERS) def benchmark_ragged(seq_len: int, batch_size: int, k: int, provider: str): + seed = seq_len ^ (batch_size << 16) ^ (k << 32) + torch.random.manual_seed(seed) if k > seq_len: marker.skip("k cannot be larger than seq_len") if k != 2048 and provider == "jit_v1": diff --git a/test/registered/kernels/ops/attention/test_topk_v2.py b/test/registered/kernels/ops/attention/test_topk_v2.py index 86adf27d9..c1fa207b4 100644 --- a/test/registered/kernels/ops/attention/test_topk_v2.py +++ b/test/registered/kernels/ops/attention/test_topk_v2.py @@ -13,13 +13,15 @@ boundaries are exercised: trivial seq <= k Register2 k < seq <= 8192 max_seq <= 8192 (level 0) Register4 8192 < seq <= 16384 max_seq <= 16384 (level 1) - Streaming 16384 < seq <= floor max_seq > 16384, non-cluster (level 2) - Cluster seq > floor(=65536) max_seq > floor and batch <= 128 + Streaming seq > 16384 max_seq > 16384, below the cluster floor (level 2) + Cluster seq above the floor the arch has clusters and batch <= 512 -and two cluster dispatch shapes: the fused small-batch kernel (batch <= 30) and -the persistent-pool + main kernel (30 < batch <= 128). Boundary seq lengths -(8192/8193, 16384/16385, 65535/65536/65537) and batch sizes (30/31, 128/129) are -included explicitly, across k in {512,1024,2048} and identity/perm page tables. +and two cluster dispatch shapes: the fused small-batch kernel (batch up to the +probed persistent-pool size) and the persistent pool + main kernel above it. The +cluster floor and pool size are per-arch (see topk_v2.cuh), so the (batch, seq) +grid below brackets the fixed boundaries (8192/8193, 16384/16385) exactly and +spans the arch-dependent ones, across k in {512,1024,2048} and identity/perm +page tables. """ from __future__ import annotations @@ -43,7 +45,6 @@ PAGE_SIZE = 64 # c4 page size = 256 // 4 PAGE_BITS = PAGE_SIZE.bit_length() - 1 PAGE_MASK = PAGE_SIZE - 1 MAX_PERMIT_ERROR = 5 -FLOOR = 65536 # kClusterFloor # (batch, seq) chosen to land on each template and each dispatch boundary. FIXED_CONFIGS = [ @@ -60,22 +61,22 @@ FIXED_CONFIGS = [ (64, 16384), # reg4 upper boundary (256, 16384), # batch > 128 # --- Streaming (level 2: max_seq > 16384, non-cluster) --- - (8, 16385), # just over reg4 (small batch, seq < floor => non-cluster) + (8, 16385), # just over reg4 (4, 32768), - (16, 65535), # just under floor - (4, 65536), # at floor (seq == floor => non-cluster) + (16, 65535), + (4, 65536), (100, 65536), - # --- Cluster, fused small-batch kernel (batch <= 30, max_seq > floor) --- - (1, 65537), # single row just over floor + # --- long rows, small batch: fused cluster kernel where the arch has clusters --- + (1, 65537), (2, 131072), (8, 98304), - (30, 131072), # batch == pool boundary - # --- Cluster, persistent pool + main kernel (30 < batch <= 128) --- - (31, 131072), # just over small-batch - (40, 262144), # N > pool of 30 => round-robin + (30, 131072), + # --- long rows, mid batch: persistent cluster pool + main kernel --- + (31, 131072), + (40, 262144), # more items than the pool => round-robin (64, 196608), - (128, 131072), # cluster batch upper boundary - # --- batch > 128 => non-cluster streaming even at long ctx --- + (128, 131072), + # --- long rows, large batch --- (129, 131072), (200, 262144), ] @@ -238,7 +239,7 @@ def test_topk_v2_ragged(batch: int, shape: str, k: int, per_row_pt: bool) -> Non device = "cuda" seq = 262144 scores = torch.randn(batch, seq, dtype=torch.float32, device=device) - # span every path; guarantee at least one > floor row so cluster dispatch fires + # span every path, including rows long enough for the cluster dispatch buckets = [max(1, k // 2), k, 4096, 12000, 40000, 65536, 98304, 262144] g = torch.Generator(device="cpu").manual_seed(batch + k) lengths = torch.tensor( @@ -387,8 +388,9 @@ def test_topk_v2_ragged_window(name: str, rows, k: int, offset_shift: int) -> No ref_raw = _reference(windows, lengths.cpu(), k) _assert_topk_close(windows, ref_raw, our_raw, len(rows), lengths.cpu(), k) - # the only legal in-place write is the <=3 masked columns ahead of a window - # that the kernel actually reads (trivial rows read nothing) + # The kernel may mask the at-most-three alignment columns immediately + # before a window. Everything else, including other rows' storage, must be + # left untouched. changed = (scores != before).cpu() for i, (start, length) in enumerate(rows): allowed = torch.zeros(scores.shape[1], dtype=torch.bool) @@ -398,6 +400,57 @@ def test_topk_v2_ragged_window(name: str, rows, k: int, offset_shift: int) -> No assert not stray, f"row {i} ({name}) wrote outside its masked head: {stray[:8]}" +def _assert_topk_values(window, indices, k): + indices = indices.cpu().long() + window = window.cpu() + assert indices.numel() == k + assert ((indices >= 0) & (indices < window.numel())).all(), indices + assert indices.unique().numel() == k + expected = window.topk(k).values.sort().values + actual = window[indices].sort().values + assert torch.equal(actual, expected) + + +@pytest.mark.parametrize("num_ties", [48, 96]) +@torch.inference_mode() +def test_topk_v2_negative_infinity_ties(num_ties: int) -> None: + """Inactive entries must not displace valid -inf scores or leave slots unwritten.""" + k = 16 + length = num_ties + 3 + scores = torch.full((1, (length + 3) & ~3), -torch.inf, device="cuda") + scores[0, :3] = torch.tensor([1.0, 2.0, 3.0], device="cuda") + lengths = torch.tensor([length], dtype=torch.int32, device="cuda") + out = torch.full((1, k), -2, dtype=torch.int32, device="cuda") + + topk_transform_paged_v2(scores, lengths, None, out, PAGE_SIZE, _plan(lengths)) + + _assert_topk_values(scores[0, :length], out[0], k) + + +@pytest.mark.parametrize("length", [257, 8193, 16385]) +@torch.inference_mode() +def test_topk_v2_ragged_negative_infinity(length: int) -> None: + """Columns before an unaligned window must never beat its valid -inf scores.""" + k = 16 + scores = torch.full((3, (length + 6) & ~3), OUTSIDE_SCORE, device="cuda") + starts = torch.tensor([1, 2, 3], dtype=torch.int32, device="cuda") + lengths = torch.full((3,), length, dtype=torch.int32, device="cuda") + offsets = starts + 1024 + out = torch.full((3, k), -2, dtype=torch.int32, device="cuda") + for row, start in enumerate((1, 2, 3)): + scores[row, start : start + length] = -torch.inf + scores[row, start : start + 3] = torch.tensor([1.0, 2.0, 3.0], device="cuda") + + topk_transform_ragged_v2( + scores, lengths, out_offsets=offsets, out_indices=out, row_starts=starts + ) + + for row, start in enumerate((1, 2, 3)): + _assert_topk_values( + scores[row, start : start + length], out[row] - offsets[row], k + ) + + @pytest.mark.parametrize("k", [512, 2048]) @torch.inference_mode() def test_topk_v2_ragged_no_row_starts(k: int) -> None: