From bdaf2e2f5709d83121cab6e4d468654d7ad4b22e Mon Sep 17 00:00:00 2001 From: hnyls2002 Date: Tue, 15 Sep 2026 15:21:04 -0700 Subject: [PATCH 01/30] dsv4.1: extract Top-k kernels and candidate helpers --- .../jit/csrc/deepseek_v4/amax_copy.cuh | 161 +++++ .../kernels/jit/csrc/deepseek_v4/sort_idx.cuh | 263 +++++++ .../jit/csrc/deepseek_v4/topk_bf16_small.cuh | 450 ++++++++++++ .../kernels/jit/csrc/deepseek_v4/topk_v2.cuh | 401 ++++++----- python/sglang/kernels/jit/csrc/misc/probe.cuh | 57 ++ .../sgl_kernel/deepseek_v4/topk_impl.cuh | 651 +++++++++--------- .../sglang/kernels/ops/attention/dsv4/topk.py | 193 +++++- python/sglang/kernels/ops/misc.py | 61 ++ .../kernel/attention/test_amax_copy.py | 125 ++++ .../test_dsv4_indexer_postprocess.py | 222 ++++++ .../kernel/attention/test_sort_idx.py | 140 ++++ .../kernel/attention/test_topk_bf16.py | 264 +++++++ .../kernels/benchmark/attention/bench_topk.py | 6 +- .../kernels/ops/attention/test_topk_v2.py | 59 +- 14 files changed, 2546 insertions(+), 507 deletions(-) create mode 100644 python/sglang/kernels/jit/csrc/deepseek_v4/amax_copy.cuh create mode 100644 python/sglang/kernels/jit/csrc/deepseek_v4/sort_idx.cuh create mode 100644 python/sglang/kernels/jit/csrc/deepseek_v4/topk_bf16_small.cuh create mode 100644 python/sglang/kernels/jit/csrc/misc/probe.cuh create mode 100644 python/sglang/kernels/ops/misc.py create mode 100644 test/registered/kernel/attention/test_amax_copy.py create mode 100644 test/registered/kernel/attention/test_dsv4_indexer_postprocess.py create mode 100644 test/registered/kernel/attention/test_sort_idx.py create mode 100644 test/registered/kernel/attention/test_topk_bf16.py diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/amax_copy.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/amax_copy.cuh new file mode 100644 index 000000000000..aa5b22dfbdbd --- /dev/null +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/amax_copy.cuh @@ -0,0 +1,161 @@ +#pragma once + +#include +#include + +#include +#include + +#include +#include + +#include +#include + +namespace sglang { + +/// Level-one keys of the DeepSeek-V4.1 two-level indexer: one key per block of +/// kBlockTokens consecutive scores, the block's maximum score (`amax` in the +/// model code, a plain max). Row `b` has `ceil(seq_len[b] / kBlockTokens)` keys; +/// its newest block is written as +inf so the block top-k can never drop it, +/// which also makes the scores past `seq_len` inside that block irrelevant. +/// Nothing is written past the key count: the consumer takes the same count as +/// the row length. Rows with at most `topk` blocks are skipped entirely (every +/// block is selected anyway; `topk = 0` disables the skip). +struct AmaxConfig { + using DType = float; // TODO: support bf16 + 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 AmaxParams { + const AmaxConfig::DType* __restrict__ scores; + AmaxConfig::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__(AmaxConfig::kBlockSize, AmaxConfig::kOccupancy) // + void amax8_varlen_kernel(const __grid_constant__ AmaxParams params) { + using namespace device; + using C = AmaxConfig; + 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 AmaxCopyKernel { + 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 = AmaxConfig; + 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 = AmaxParams{ + .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/sort_idx.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/sort_idx.cuh new file mode 100644 index 000000000000..2c68167bd9ca --- /dev/null +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/sort_idx.cuh @@ -0,0 +1,263 @@ +#pragma once + +#include +#include + +#include +#include +#include + +#include +#include + +#include +#include +#include + +namespace sglang { + +/// Finalises the block table layer 20 publishes for DeepGEMM's sparse indexer: +/// 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 keeps every block and gets +/// the identity table without reading its input. +/// +/// Counting sort over a bitmap of the row's blocks (one bit per block, 16 KiB +/// for the 128K blocks of a 1M-token row): set the selected bits, exclusive-scan +/// the popcounts, emit every set bit at its rank. A word with a single bit is +/// emitted by its owner (one `ffs`, no loop); a word with more goes to a +/// block-wide queue that the warps drain one word per step, one lane per bit, +/// so a dense cluster of selected blocks is spread over all warps. +struct SortConfig { + 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 SortParams { + 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__(SortConfig::kBlockSize, SortConfig::kOccupancy) // + void sort_128k_transform(const __grid_constant__ SortParams params) { + using namespace device; + using C = SortConfig; + __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); + } + } +} + +/// The page transform alone, for a block top-k that already emits its ids +/// ascending (e.g. DeepSelect with `sorted_index`): `out_pages[t]` is the pool +/// slot / 8 of `indices[t]` for `t < min(topk, ceil(seq_len / 8))`, INT32_MAX +/// past that; `indices` is left as it is. Those first entries must be valid +/// block ids of the row. +template +__global__ __launch_bounds__(SortConfig::kBlockSize, SortConfig::kOccupancy) // + void page_transform_128k(const __grid_constant__ SortParams params) { + using namespace device; + using C = SortConfig; + const auto bx = blockIdx.x; + const auto tx = threadIdx.x; + PDLWaitPrimary(); + const auto seq_len = params.seq_len[bx]; + const auto nblocks = (seq_len + C::kBlockTokens - 1) / C::kBlockTokens; + const auto num_valid = min(nblocks, params.topk); + const auto* __restrict__ table = params.page_table + bx * params.page_table_stride; + const 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; + for (uint32_t t = tx; t < params.topk; t += C::kBlockSize) { + if (t < num_valid) { + const auto id = static_cast(indices[t]); + pages[t] = (table[id >> params.page_bits] << params.page_bits) | static_cast(id & bpp_mask); + } else { + pages[t] = C::kPad; + } + } + PDLTriggerSecondary(); +} + +/// 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 SortIdxKernel { + /// Sort + page transform, in place on `indices`. + 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); + } + + /// Page transform only, `indices` already ascending and left untouched. + static void transform_pages( + 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 = SortConfig; + 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 = SortParams{ + .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 000000000000..e42b9f70bb0a --- /dev/null +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/topk_bf16_small.cuh @@ -0,0 +1,450 @@ +/** + * \brief DeepSeek-V4.1's bf16 top-k kernel for short rows (<= 16384 scores) + * Adapted from https://github.com/deepseek-ai/DeepSelect + * We only tuned for 16384 in + k=512 + * Rewrite in SIMT for better architecture portability (AMD team should thank me) + */ +#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. The row is split into contiguous per-thread slices of up to + * 32 scores held in registers for the whole kernel. Two radix passes (the raw high byte, then + * the raw low byte among the elements sharing the pivot's high byte) locate the k-th largest + * value exactly, a census then tells every thread how many of its elements are above / equal + * to it and where they go, and the selected indices are staged in shared memory before one + * coalesced, page-transformed copy to the output. 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 +}; + +/** + * \brief Locate the bin holding the k-th largest element in a 256-bin histogram indexed by a + * raw byte. Called by one whole warp; exactly one lane finds it and writes the answer + * to `smem.pivot_*` (the block reads it behind the caller's barrier, so there is no + * point in broadcasting it inside the warp first). + * \param msb_mode The raw byte is the high byte: lanes < 16 cover raw 0xFF..0x80 (negatives, + * reversed), lanes >= 16 cover raw 0x00..0x7F. + * \param 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; + } +} + +/// \brief One byte of hit bits for the 8 elements of a vector, element `e` at bit `e`. +/// \param m Per-pair 16-bit masks (0xFFFF / 0) as produced by `__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. That only happens with NaN scores (the + // histogram counts them, no ordered compare ever selects them); write -1 there rather than + // whatever shared memory held before. + const uint32_t totals = smem.count_gt_eq; + const uint32_t num_staged = (totals >> 16) + min(totals & 0xFFFFu, eq_total); + // TODO: pragma unroll this one, if real topk > 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 426f047881d0..bd30e5d50cc8 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,16 @@ 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 + batch_id * page_table_stride, page_bits, 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 +123,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 +137,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 +224,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 +259,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 +268,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 +287,135 @@ 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) { +#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) + +/** + * \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(); - 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)); + 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/misc/probe.cuh b/python/sglang/kernels/jit/csrc/misc/probe.cuh new file mode 100644 index 000000000000..63019b6a3ac3 --- /dev/null +++ b/python/sglang/kernels/jit/csrc/misc/probe.cuh @@ -0,0 +1,57 @@ +#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)); + + // Threads alone cannot pin the probe to `num_waves` blocks per SM: a block + // caps at 1024, so num_waves == 1 still leaves room for a second block. Spend + // the shared budget instead -- what one block gets at this occupancy, less the + // per-block driver reserve, floored to the 1 KiB allocation granularity so the + // driver cannot round it back up and squeeze out a block. + 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 7dedd89a56e1..0b09b4335cd2 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 @@ -27,7 +27,15 @@ #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 +43,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 +76,85 @@ 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); +// } + 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); +} + +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 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. +// 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`. (Measured faster than a per-key dispatch: the two + // conversions are independent and issue in parallel.) + 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 +187,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 +195,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 +231,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 +258,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 +325,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 +346,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 +401,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 +437,119 @@ 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 + // 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 + // footprint. 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 + 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); + } + } while (vi < num_full); + } - 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; + if (vi == num_full) { const auto base = vi * kVecSize; - vi += kBlockSize; - if (vi < num_full) next_vec.load(in, vi); + if (base == seq_len) return; + vec_t cur; + cur.load(in, vi); #pragma unroll for (uint32_t j = 0; j < kVecSize; ++j) { - fn(cur[j], base + j); + if (base + j < seq_len) fn(cur[j], base + j); } } + } - // 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); - } + 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); } - SGL_DEVICE static void find_threshold(const uint32_t topk, const uint32_t seq_len, Smem* smem) { + /// 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 +568,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 +580,82 @@ 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) { + if (j >= tail_start) local_vecs[i][j] = padding_value(); + } + } #pragma unroll - for (uint32_t j = 0; j < kVecSize; ++j) + 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) { + atomicSub(&smem->histogram[kHistSize - 1], 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 +663,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 +685,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[kHistSize - 1], 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 +715,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 +728,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 +744,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); - - // 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; + barrier_cluster_arrive_relaxed(); // bar-2 arrive + barrier_cluster_wait(); // bar-2 wait + + // 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/ops/attention/dsv4/topk.py b/python/sglang/kernels/ops/attention/dsv4/topk.py index e5f0d0fd1889..7a0186e0ed02 100644 --- a/python/sglang/kernels/ops/attention/dsv4/topk.py +++ b/python/sglang/kernels/ops/attention/dsv4/topk.py @@ -18,11 +18,6 @@ @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,168 @@ 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.ops.misc 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 + occ_8_2, occ_16_1 = 0, 0 + try: + occ_8_2 = get_max_active_clusters(8, occupancy=2) + # NOTE: cluster 16 might fail, but at least cluster 8 is ok + occ_16_1 = get_max_active_clusters(16, occupancy=1) + except Exception: + pass + extra_cuda_cflags = [ + f"-DSGL_TOPK_V2_MAX_C8_OCC2={occ_8_2}", + 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 + ) + + +@cache_once +def _jit_amax_copy_module(): + args = make_cpp_args(is_arch_support_pdl()) + return load_jit( + make_name("amax_copy"), + *args, + cuda_files=["deepseek_v4/amax_copy.cuh"], + cuda_wrappers=[("amax8_varlen", f"AmaxCopyKernel<{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_amax_copy_module().amax8_varlen(scores, seq_lens, out, topk) + return out + + +@cache_once +def _jit_sort_idx_module(): + args = make_cpp_args(is_arch_support_pdl()) + return load_jit( + make_name("sort_idx"), + *args, + cuda_files=["deepseek_v4/sort_idx.cuh"], + cuda_wrappers=[ + ("transform", f"SortIdxKernel<{args}>::transform"), + ("transform_pages", f"SortIdxKernel<{args}>::transform_pages"), ], ) +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_sort_idx_module().transform(blocks, seq_lens, page_table, out_pages, page_size) + return out_pages + + +def transform_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 page transform of ``sort_candidate_blocks`` alone, for a block top-k + that already emits ascending ids: ``out_pages[b, t]`` is the pool slot / 8 of + ``blocks[b, t]`` for ``t < min(k, ceil(seq_lens[b] / 8))`` (which must be + valid block ids), ``INT32_MAX`` past that; ``blocks`` is not modified. + Returns ``out_pages``. + """ + if out_pages is None: + out_pages = torch.empty_like(blocks) + _jit_sort_idx_module().transform_pages( + blocks, seq_lens, page_table, out_pages, page_size + ) + return out_pages + + def topk_transform_paged( scores: torch.Tensor, seq_lens: torch.Tensor, @@ -75,15 +219,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] @@ -111,7 +254,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 +294,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/kernels/ops/misc.py b/python/sglang/kernels/ops/misc.py new file mode 100644 index 000000000000..4e056f1078aa --- /dev/null +++ b/python/sglang/kernels/ops/misc.py @@ -0,0 +1,61 @@ +"""Device probes that a launch configuration depends on. + +Not an operator group -- these answer "what will the hardware actually schedule", +which a host-side dispatch needs before it can size a grid. Kept out of +``ops/__init__``'s eager group import for that reason; import it directly. +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +from sglang.kernels.jit.utils import cache_once, 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( + "misc_probe", + cuda_files=["misc/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 ``num_sms * occupancy / cluster_size`` once + the cluster stops dividing a GPC evenly. On B200 (148 SMs) at occupancy 2 the + driver reports 33 clusters of 8 where the division says 37, and 14 of 16 + where it says 18. + + Probed with an empty kernel pinned to ``occupancy`` blocks per SM, so the + answer is the *scheduling* limit at that occupancy and nothing else. Pass the + occupancy the real kernel reaches (the second ``__launch_bounds__`` + argument), not the one it asks for. + + :param cluster_size: Blocks per cluster. + :param occupancy: Blocks per SM (``num_waves`` in ``csrc/misc/probe.cuh``). + :raises RuntimeError: On pre-sm90 devices, which have no clusters. + :raises ValueError: If nothing is schedulable, which a real device should + never report for a cluster width it supports. + """ + 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/test/registered/kernel/attention/test_amax_copy.py b/test/registered/kernel/attention/test_amax_copy.py new file mode 100644 index 000000000000..603eabaac6eb --- /dev/null +++ b/test/registered/kernel/attention/test_amax_copy.py @@ -0,0 +1,125 @@ +"""Correctness tests for the DeepSeek-V4.1 JIT block-max ("amax") copy. + +``amax8_varlen`` writes, per row, the maximum of each block of 8 consecutive +fp32 scores for the first ``ceil(seq_lens[b] / 8)`` blocks, the last of them +``+inf`` (the newest block is always selected), and leaves everything else +untouched. It is the level-one key computation of the two-level sparse indexer: +the keys feed a block top-k, which is why rows with at most ``topk`` blocks may +be skipped (every block of such a row is selected anyway). +""" + +from __future__ import annotations + +import sys + +import pytest +import torch + +from sglang.kernels.ops.attention.dsv4.topk import amax8_varlen +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=20, stage="base-b-kernel-unit", runner_config="1-gpu-large") + +BLOCK = 8 +SENTINEL = -12345.0 +CONFIGS = [ + # (batch, seq): a single partial block, exact blocks, one CTA, many CTAs + (1, 1), + (1, 8), + (1, 9), + (3, 1000), + (4, 16389), + (8, 131072), + (2, 1048576), +] + + +def _keys_ref(scores: torch.Tensor, lens: torch.Tensor) -> torch.Tensor: + """Per row the block maxima of the first ceil(len / 8) blocks, the last +inf.""" + batch, width = scores.shape + nblocks = (width + BLOCK - 1) // BLOCK + padded = torch.nn.functional.pad( + scores, (0, nblocks * BLOCK - width), value=-torch.inf + ) + keys = padded.view(batch, nblocks, BLOCK).amax(-1) + last = (lens.long() + BLOCK - 1) // BLOCK - 1 + keys[torch.arange(batch, device=scores.device), last] = torch.inf + return keys + + +def _check( + out: torch.Tensor, scores: torch.Tensor, lens: torch.Tensor, skip=None +) -> None: + ref = _keys_ref(scores, lens) + for b in range(scores.shape[0]): + n = (int(lens[b]) + BLOCK - 1) // BLOCK + if skip is not None and skip[b]: + assert torch.all(out[b] == SENTINEL), f"row {b} was skipped but written" + continue + assert torch.equal(out[b, :n], ref[b, :n]), f"row {b} keys differ" + assert torch.all(out[b, n:] == SENTINEL), f"row {b} written past its keys" + + +def _inputs(batch: int, seq: int, ragged: bool, stride_pad: int = 0): + torch.manual_seed(batch * 977 + seq) + lens = torch.full((batch,), seq, dtype=torch.int32, device="cuda") + if ragged: + lens = torch.randint(1, seq + 1, (batch,), dtype=torch.int32, device="cuda") + lens[0] = seq + stride = (seq + BLOCK - 1) // BLOCK * BLOCK + stride_pad + storage = torch.randn(batch, stride, device="cuda") * 10 + scores = storage[:, :seq] + return scores, lens + + +@pytest.mark.parametrize("ragged", [False, True]) +@pytest.mark.parametrize("batch,seq", CONFIGS) +def test_matches_block_max(batch: int, seq: int, ragged: bool): + scores, lens = _inputs(batch, seq, ragged) + out = torch.full((batch, (seq + BLOCK - 1) // BLOCK), SENTINEL, device="cuda") + assert amax8_varlen(scores, lens, out=out) is out + _check(out, scores, lens) + + +def test_strided_views(): + """Score rows padded past the row width and a key buffer wider than needed.""" + scores, lens = _inputs(4, 16389, ragged=True, stride_pad=248) + nblocks = (16389 + BLOCK - 1) // BLOCK + out = torch.full((4, nblocks + 5), SENTINEL, device="cuda") + amax8_varlen(scores, lens, out=out[:, :nblocks]) + _check(out, scores, lens) + + +def test_allocates_output(): + scores, lens = _inputs(3, 20000, ragged=True) + out = amax8_varlen(scores, lens) + assert out.shape == (3, (20000 + BLOCK - 1) // BLOCK) and out.dtype == torch.float32 + ref = _keys_ref(scores, lens) + for b in range(3): + n = (int(lens[b]) + BLOCK - 1) // BLOCK + assert torch.equal(out[b, :n], ref[b, :n]) + + +def test_max_seqlen_sizes_output(): + """A batch whose rows are all shorter than the padded width only needs + ceil(max_seqlen / 8) keys per row.""" + scores, _ = _inputs(4, 65536, ragged=False) + lens = torch.tensor([40000, 1, 16384, 39999], dtype=torch.int32, device="cuda") + out = amax8_varlen(scores, lens, max_seqlen=40000) + assert out.shape == (4, 5000) + ref = _keys_ref(scores, lens) + for b in range(4): + n = (int(lens[b]) + BLOCK - 1) // BLOCK + assert torch.equal(out[b, :n], ref[b, :n]) + + +def test_topk_skips_rows_that_fit(): + scores, _ = _inputs(4, 40000, ragged=False) + lens = torch.tensor([40000, 16384, 16385, 100], dtype=torch.int32, device="cuda") + out = torch.full((4, 5000), SENTINEL, device="cuda") + amax8_varlen(scores, lens, 2048, out=out) + _check(out, scores, lens, skip=[False, True, False, True]) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/kernel/attention/test_dsv4_indexer_postprocess.py b/test/registered/kernel/attention/test_dsv4_indexer_postprocess.py new file mode 100644 index 000000000000..af9da90c4203 --- /dev/null +++ b/test/registered/kernel/attention/test_dsv4_indexer_postprocess.py @@ -0,0 +1,222 @@ +import unittest + +import torch + +from sglang.kernels.ops.attention.dsv4.candidate_blocks import ( + candidate_block_logits, + candidate_row_lens, +) +from sglang.kernels.ops.attention.dsv4.indexer_postprocess import filter_topk_pages +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci(est_time=20, stage="base-b-kernel-unit", runner_config="1-gpu-large") + + +def reference_pages(scores, indices, pages, page_size): + cols = indices.to(torch.int64) + value = scores.gather(1, cols.clamp(0, scores.shape[1] - 1)) + valid = (cols >= 0) & (cols < scores.shape[1]) & (value > -torch.inf) + raw = indices.masked_fill(~valid, -1) + safe = raw.clamp_min(0).to(torch.int64) + slots = pages.gather(1, safe // page_size) * page_size + safe % page_size + return torch.where(raw >= 0, slots, -1).to(torch.int32), raw + + +class TestIndexerPostprocess(CustomTestCase): + def test_unfiltered_verify_matches_candidate_chain(self): + from sglang.kernels.ops.attention.dsv4.topk import ( + plan_topk_v2, + topk_transform_paged_v2, + ) + + # Six causal rows per request. The last request ends exactly at the + # candidate budget; capacity and unread logits extend well past it. + torch.manual_seed(123) + rows, width, page_size = 384, 32768, 64 + base = torch.tensor([0, 506, 4096, 16378], device="cuda", dtype=torch.int32) + lens = (base.repeat(16)[:, None] + torch.arange(1, 7, device="cuda")).flatten() + lens = lens.to(torch.int32) + source = torch.randn(rows, width, device="cuda").relu_() + consumer = torch.randn_like(source).relu_() + cols = torch.arange(width, device="cuda")[None, :] + source.masked_fill_(cols >= lens[:, None], 1e6) + consumer.masked_fill_(cols >= lens[:, None], 1e6) + pages = torch.randint( + 0, 100000, (rows, width // page_size), device="cuda", dtype=torch.int32 + ) + source_masked, keep = candidate_block_logits( + source, lens, topk_blocks=2048, block_size=8, published=None + ) + consumer_masked, _ = candidate_block_logits( + consumer, lens, topk_blocks=2048, block_size=8, published=keep + ) + plan = plan_topk_v2(lens) + for original, masked in ((source, source_masked), (consumer, consumer_masked)): + with self.subTest(source=original is source): + old = torch.empty((rows, 512), dtype=torch.int32, device="cuda") + new = torch.empty_like(old) + raw = torch.empty_like(old) + new_raw = torch.empty_like(old) + topk_transform_paged_v2(masked, lens, pages, old, page_size, plan, raw) + filter_topk_pages(masked, raw, pages, old, page_size) + topk_transform_paged_v2( + original, lens, pages, new, page_size, plan, new_raw + ) + # Top-k v2 uses a persistent work queue and does not promise + # output order. Compare the selected multiset, including -1 + # padding, in both logical-position and physical-page space. + torch.testing.assert_close( + new.sort(-1).values, old.sort(-1).values, rtol=0, atol=0 + ) + torch.testing.assert_close( + new_raw.sort(-1).values, raw.sort(-1).values, rtol=0, atol=0 + ) + + def test_candidate_row_lens(self): + lens = torch.tensor( + [1, 7, 8, 9, 300, 16383, 16384, 16385, 16392, 40000, 1048576, 1048571], + dtype=torch.int32, + device="cuda", + ) + for topk in (2048, 256, 1): + nblocks, valid = candidate_row_lens(lens, topk) + ref_blocks = (lens + 7) // 8 + kept = ref_blocks.clamp_max(topk) + ref_valid = 8 * (kept - 1) + (lens - 1) % 8 + 1 + self.assertEqual(nblocks.dtype, torch.int32) + self.assertTrue(torch.equal(nblocks, ref_blocks)) + self.assertTrue(torch.equal(valid, ref_valid)) + # a row with all its blocks kept is the whole row + fits = ref_blocks <= topk + self.assertTrue(torch.equal(valid[fits], lens[fits])) + nblocks, valid = candidate_row_lens( + torch.zeros(3, dtype=torch.int32, device="cuda"), 2048 + ) + self.assertTrue(torch.equal(nblocks, torch.zeros_like(nblocks))) + self.assertTrue(torch.equal(valid, torch.zeros_like(valid))) + + def test_filter_and_page_mapping(self): + torch.manual_seed(91) + for rows in (1, 6, 64): + for width in (1, 65, 4096): + for dtype in (torch.int32, torch.int64): + # Row-strided tensors match sliced metadata buffers. + scores = torch.randn(rows, width + 7, device="cuda")[:, :width] + indices = torch.randint( + -2, width + 2, (rows, 520), device="cuda", dtype=dtype + )[:, :512] + indices[:, :5] = 0 + scores[0, 0] = -torch.inf + if rows > 1: + scores[1, 0] = torch.nan + scores[2, 0] = torch.inf + scores[3, :] = -torch.inf + pages = torch.randint( + 0, + 100000, + (rows, (width + 63) // 64 + 3), + device="cuda", + dtype=torch.int32, + )[:, :-3] + out = torch.empty(rows, 520, device="cuda", dtype=torch.int32)[ + :, :512 + ] + raw = torch.empty_like(indices) + expected, expected_raw = reference_pages(scores, indices, pages, 64) + for write_raw in (False, True): + filter_topk_pages( + scores, indices, pages, out, 64, raw if write_raw else None + ) + torch.testing.assert_close(out, expected, rtol=0, atol=0) + if write_raw: + torch.testing.assert_close( + raw, expected_raw, rtol=0, atol=0 + ) + + def test_graph_replay(self): + rows, width = 6, 4096 + scores = torch.randn(rows, width, device="cuda") + indices = torch.randint( + -1, width, (rows, 512), device="cuda", dtype=torch.int32 + ) + pages = torch.randint( + 0, 100000, (rows, width // 64), device="cuda", dtype=torch.int32 + ) + out, raw = torch.empty_like(indices), torch.empty_like(indices) + filter_topk_pages(scores, indices, pages, out, 64, raw) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + filter_topk_pages(scores, indices, pages, out, 64, raw) + for _ in range(3): + scores.normal_() + scores[:, :32] = -torch.inf + indices.random_(-1, width) + pages.random_(0, 100000) + graph.replay() + expected, expected_raw = reference_pages(scores, indices, pages, 64) + torch.testing.assert_close(out, expected, rtol=0, atol=0) + torch.testing.assert_close(raw, expected_raw, rtol=0, atol=0) + + def test_candidate_publication(self): + torch.manual_seed(912) + for width, group, topk in ( + (65, 32, 8), + (4096, 128, 8), + (32768, 2048, 8), + (1048576, 2048, 8), + (4096, 8, 2048), + (1048576, 8, 2048), + ): + rows = 6 + x = torch.randn(rows, width + 9, device="cuda")[:, :width] + lengths = torch.tensor( + [0, 1, width // 2, width - 1, width, width], + device="cuda", + dtype=torch.int32, + ) + # Include ties, NaN and a final partial block. + x[3, :] = 0 + x[4, 0] = torch.nan + x[5, :] = -torch.inf + masked = x.masked_fill( + torch.arange(width, device="cuda")[None, :] >= lengths[:, None], + -torch.inf, + ) + blocks = (width + group - 1) // group + padded = torch.nn.functional.pad( + masked, (0, blocks * group - width), value=-torch.inf + ) + scores = padded.reshape(rows, blocks, group).max(-1).values + last = (lengths - 1) // group + scores = torch.where( + (torch.arange(blocks, device="cuda")[None, :] == last[:, None]) + & (lengths[:, None] > 0), + torch.inf, + scores, + ) + selected = scores.topk(min(topk, blocks), dim=-1) + keep = ( + torch.zeros_like(scores, dtype=torch.bool) + .scatter_(1, selected.indices, selected.values > -torch.inf) + .repeat_interleave(group, dim=-1)[:, :width] + ) + got, published = candidate_block_logits( + x, lengths, topk_blocks=topk, block_size=group, published=None + ) + torch.testing.assert_close(got, masked, rtol=0, atol=0, equal_nan=True) + torch.testing.assert_close(published, keep, rtol=0, atol=0) + consumer, _ = candidate_block_logits( + x, lengths, topk_blocks=topk, block_size=group, published=published + ) + torch.testing.assert_close( + consumer, + masked.masked_fill(~keep, -torch.inf), + rtol=0, + atol=0, + equal_nan=True, + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/kernel/attention/test_sort_idx.py b/test/registered/kernel/attention/test_sort_idx.py new file mode 100644 index 000000000000..8b3bef80122e --- /dev/null +++ b/test/registered/kernel/attention/test_sort_idx.py @@ -0,0 +1,140 @@ +"""Correctness tests for the DeepSeek-V4.1 JIT candidate block table sort. + +``sort_candidate_blocks`` turns a row's selected block ids (any order, ``-1`` +padded) in place into the ascending, ``INT32_MAX``-padded table DeepGEMM's +sparse indexer schedule reads, and writes the same blocks as pool slots / 8 +through the row's index page table. Rows with at most ``k`` blocks get the +identity table. The kernel is a bitmap counting sort whose dense words are +drained by a block-wide queue, so the id distributions below cover sparse, +clustered and completely full words. +""" + +from __future__ import annotations + +import sys + +import pytest +import torch + +from sglang.kernels.ops.attention.dsv4.topk import ( + sort_candidate_blocks, + transform_candidate_blocks, +) +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=20, stage="base-b-kernel-unit", runner_config="1-gpu-large") + +BLOCK = 8 +PAD = torch.iinfo(torch.int32).max +MAX_BLOCKS = 1 << 17 + + +def _select(nblocks: int, k: int, dist: str, gen: torch.Generator) -> torch.Tensor: + """k distinct block ids below nblocks (all of them if nblocks <= k), shuffled.""" + n = min(k, nblocks) + if dist == "uniform": + ids = torch.randperm(nblocks, device="cuda", generator=gen)[:n] + elif dist == "clustered": # runs of 64 consecutive blocks mixed with scattered ones + runs = max(n // 96, 1) + starts = torch.randint( + 0, max(nblocks - 64, 1), (runs,), device="cuda", generator=gen + ) + runs_ids = ( + starts[:, None] + torch.arange(64, device="cuda")[None, :] + ).flatten() + rest = torch.randperm(nblocks, device="cuda", generator=gen)[:n] + ids = torch.unique(torch.cat([runs_ids, rest]))[:n] + elif dist == "newest": # the tail of the row: every word full + ids = torch.arange(nblocks - n, nblocks, device="cuda") + else: + raise ValueError(dist) + ids = ids.to(torch.int32) + ids = ids[torch.randperm(ids.numel(), device="cuda", generator=gen)] + return torch.nn.functional.pad(ids, (0, k - n), value=-1) + + +def _case(lens, k, dist, page_size, seed=0): + gen = torch.Generator(device="cuda").manual_seed(seed) + lens_t = torch.tensor(lens, dtype=torch.int32, device="cuda") + max_pages = (max(lens) + page_size - 1) // page_size + page_table = torch.stack( + [torch.randperm(max_pages, device="cuda", generator=gen) for _ in lens] + ).to(torch.int32) + blocks = torch.stack( + [_select((l + BLOCK - 1) // BLOCK, k, dist, gen) for l in lens] + ) + return lens_t, page_table, blocks + + +def _check(blocks, pages, selected, lens, page_table, page_size, k): + bpp = page_size // BLOCK + for b, seq in enumerate(lens): + nblocks = (seq + BLOCK - 1) // BLOCK + n = min(k, nblocks) + if nblocks <= k: + ref = torch.arange(n, device="cuda", dtype=torch.int32) + else: + ref = selected[b][selected[b] >= 0].sort().values + assert torch.equal(blocks[b, :n], ref), f"row {b} ids" + assert torch.all(blocks[b, n:] == PAD), f"row {b} padding" + ref_pages = page_table[b][ref // bpp] * bpp + ref % bpp + assert torch.equal(pages[b, :n], ref_pages), f"row {b} slots" + assert torch.all(pages[b, n:] == PAD), f"row {b} slot padding" + + +@pytest.mark.parametrize("dist", ["uniform", "clustered", "newest"]) +@pytest.mark.parametrize("page_size", [64, 128]) +@pytest.mark.parametrize( + "lens,k", + [ + # rows that fit (identity), rows just above k, long rows up to 1M tokens + ([1, 8, 16384, 16385, 16392], 2048), + ([40000, 131072, 65537, 1048576], 2048), + ([1048576, 1048571, 300000], 2048), + ([2049, 9000, 20000], 256), + ], +) +def test_matches_sorted_selection(lens, k, dist, page_size): + lens_t, page_table, blocks = _case(lens, k, dist, page_size) + selected = blocks.clone() + pages = sort_candidate_blocks(blocks, lens_t, page_table, page_size) + _check(blocks, pages, selected, lens, page_table, page_size, k) + + +def test_dense_words_only(): + """Every selected id inside k / 32 full words, in random order: the queue path only.""" + lens = [MAX_BLOCKS * BLOCK, 100000] + lens_t, page_table, _ = _case(lens, 2048, "uniform", 128) + gen = torch.Generator(device="cuda").manual_seed(3) + rows = [] + for seq in lens: + nblocks = (seq + BLOCK - 1) // BLOCK + start = (nblocks - 2048) // 2 // 32 * 32 + ids = torch.arange(start, start + 2048, device="cuda", dtype=torch.int32) + rows.append(ids[torch.randperm(2048, device="cuda", generator=gen)]) + blocks = torch.stack(rows) + selected = blocks.clone() + pages = torch.empty_like(blocks) + assert ( + sort_candidate_blocks(blocks, lens_t, page_table, 128, out_pages=pages) is pages + ) + _check(blocks, pages, selected, lens, page_table, 128, 2048) + + +@pytest.mark.parametrize("page_size", [64, 128]) +def test_page_transform_of_sorted_ids(page_size): + """Ascending ids in, INT32_MAX past the row's count (a DeepSelect-style + input): the same pages as the sort, blocks untouched.""" + lens = [1, 16384, 16385, 131072, 1048576] + k = 2048 + lens_t, page_table, blocks = _case(lens, k, "clustered", page_size, seed=11) + ref_pages = sort_candidate_blocks(blocks.clone(), lens_t, page_table, page_size) + ascending = torch.where(blocks < 0, PAD, blocks).sort(dim=1).values + kept = ascending.clone() + pages = transform_candidate_blocks(ascending, lens_t, page_table, page_size) + assert torch.equal(ascending, kept) + assert torch.equal(pages, ref_pages) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/kernel/attention/test_topk_bf16.py b/test/registered/kernel/attention/test_topk_bf16.py new file mode 100644 index 000000000000..6cdce0fa47d5 --- /dev/null +++ b/test/registered/kernel/attention/test_topk_bf16.py @@ -0,0 +1,264 @@ +"""Correctness tests for the DeepSeek-V4.1 JIT bf16 top-k transform. + +``topk_transform_bf16_small`` selects the per-row top-k of bf16 ``scores`` within each +row's ``seq_lens`` (rows of at most 16384, ``k`` at most 2048) and writes the +page-table transform of the selected indices, ``-1`` past ``min(k, seq_len)``, in +no particular order. It is the consumer selection of the two-level sparse +indexer: there the "page table" is the per-row physical block table at page +size 8. + +Selection is exact, so it is validated against ``torch.topk`` on the multiset of +selected values (elements of equal value may swap). NaN scores are not +supported and never generated here. +""" + +from __future__ import annotations + +import sys + +import pytest +import torch + +from sglang.kernels.ops.attention.dsv4.topk import topk_transform_bf16_small +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="1-gpu-large") + +MAX_SEQ = 16384 +CONFIGS = [ + # (batch, seq): short rows select everything, full rows the 16384 maximum + (1, 1), + (4, 300), + (8, 512), + (8, 513), + (16, 4097), + (32, 16000), + (32, MAX_SEQ), + (148, MAX_SEQ), + (300, MAX_SEQ), +] + + +def _rows(batch: int, seq: int, k: int, ties: bool, ragged: bool): + torch.manual_seed(batch * 131 + seq) + lens = torch.full((batch,), seq, dtype=torch.int32, device="cuda") + if ragged: + lens = torch.randint(1, seq + 1, (batch,), dtype=torch.int32, device="cuda") + lens[0] = seq + scores = torch.randn(batch, MAX_SEQ, device="cuda") * 2 + if ties: + scores = (scores * 2).round() / 2 # a handful of distinct values per row + # everything past a row is garbage that must never be selected + scores.masked_fill_( + torch.arange(MAX_SEQ, device="cuda")[None, :] >= lens[:, None], 1e4 + ) + return scores.to(torch.bfloat16), lens + + +def _identity_table(batch: int, page_size: int = 8) -> torch.Tensor: + return torch.arange(MAX_SEQ // page_size, device="cuda", dtype=torch.int32).repeat( + batch, 1 + ) + + +def _check(scores, lens, table, page_size, out): + """Every row: min(k, len) valid slots, the rest -1; slots invert to unique + in-row indices whose values are torch's top-k values.""" + k = out.shape[1] + nblocks = MAX_SEQ // page_size + inv = torch.empty_like(table) + inv.scatter_( + 1, + table.long(), + torch.arange(nblocks, device="cuda", dtype=torch.int32).expand_as(table), + ) + for b in range(scores.shape[0]): + n = min(k, int(lens[b])) + chosen = out[b] >= 0 + assert int(chosen.sum()) == n, ( + f"row {b}: {int(chosen.sum())} selected, want {n}" + ) + slots = out[b][chosen].long() + idx = inv[b][slots // page_size].long() * page_size + slots % page_size + assert bool((idx < int(lens[b])).all()), f"row {b}: index past the row" + assert idx.unique().numel() == n, f"row {b}: duplicate index" + got = scores[b, idx].float().sort(descending=True).values + ref = scores[b, : int(lens[b])].float().topk(n).values + assert torch.equal(got, ref), f"row {b}: selected values differ from torch.topk" + + +def _run_full_rows(scores: torch.Tensor, k: int) -> None: + """Whole rows of `scores` (any width up to MAX_SEQ) through the identity table.""" + batch, seq = scores.shape + lens = torch.full((batch,), seq, dtype=torch.int32, device="cuda") + table = _identity_table(batch) + out = torch.full((batch, k), -7, dtype=torch.int32, device="cuda") + topk_transform_bf16_small(scores, lens, table, out, 8) + torch.cuda.synchronize() + _check(scores, lens, table, 8, out) + + +def _randn_aligned(batch: int, seq: int) -> torch.Tensor: + # rows must start vector-aligned, so odd widths come from a wider padded tensor + padded = (seq + 7) // 8 * 8 + return torch.randn(batch, padded, device="cuda", dtype=torch.bfloat16)[:, :seq] + + +@pytest.mark.parametrize("page_mode", ["identity", "perm"]) +@pytest.mark.parametrize("ties", [False, True]) +@pytest.mark.parametrize("k", [512, 2048]) +@pytest.mark.parametrize("batch,seq", CONFIGS) +def test_topk_bf16(batch: int, seq: int, k: int, ties: bool, page_mode: str) -> None: + page_size = 8 + scores, lens = _rows(batch, seq, k, ties, ragged=False) + nblocks = MAX_SEQ // page_size + if page_mode == "identity": + table = _identity_table(batch, page_size) + else: + table = torch.stack( + [torch.randperm(nblocks, device="cuda") for _ in range(batch)] + ).to(torch.int32) + out = torch.full((batch, k), -7, dtype=torch.int32, device="cuda") + topk_transform_bf16_small(scores, lens, table, out, page_size) + torch.cuda.synchronize() + _check(scores, lens, table, page_size, out) + + +@pytest.mark.parametrize("page_size", [8, 64]) +def test_topk_bf16_ragged_lengths(page_size: int) -> None: + batch, k = 64, 512 + scores, lens = _rows(batch, MAX_SEQ, k, ties=False, ragged=True) + nblocks = MAX_SEQ // page_size + table = torch.stack( + [torch.randperm(nblocks, device="cuda") for _ in range(batch)] + ).to(torch.int32) + out = torch.full((batch, k), -7, dtype=torch.int32, device="cuda") + topk_transform_bf16_small(scores, lens, table, out, page_size) + torch.cuda.synchronize() + _check(scores, lens, table, page_size, out) + + +def test_topk_bf16_padded_output() -> None: + """The output may be a view whose last dimension was padded.""" + batch, k, page_size = 3, 512, 8 + scores, lens = _rows(batch, MAX_SEQ, k, ties=False, ragged=False) + table = _identity_table(batch, page_size) + buf = torch.full((batch, k + 64), -7, dtype=torch.int32, device="cuda") + out = buf[:, :k] + topk_transform_bf16_small(scores, lens, table, out, page_size) + torch.cuda.synchronize() + assert bool((buf[:, k:] == -7).all()), "wrote past the view" + _check(scores, lens, table, page_size, out) + + +@pytest.mark.parametrize("k", [1, 512, 2048]) +@pytest.mark.parametrize("seq", [1000, 3000, 8191, 12345]) +def test_topk_bf16_odd_widths(seq: int, k: int) -> None: + """Rows whose width is not a multiple of the vector: the last vector is partial.""" + torch.manual_seed(seq * 7 + k) + _run_full_rows(_randn_aligned(33, seq), k) + + +@pytest.mark.parametrize( + "seq,k", [(4096, 1024), (16384, 2048), (16384, 512), (1000, 512)] +) +def test_topk_bf16_heavy_ties(seq: int, k: int) -> None: + """Heavy ties, all-equal rows and all -inf rows exercise the equal-quota path.""" + torch.manual_seed(seq + k) + _run_full_rows(torch.randint(0, 5, (32, seq), device="cuda").bfloat16(), k) + _run_full_rows(torch.randint(-3, 3, (32, seq), device="cuda").bfloat16(), k) + _run_full_rows(torch.full((32, seq), 1.5, device="cuda", dtype=torch.bfloat16), k) + _run_full_rows( + torch.full((32, seq), float("-inf"), device="cuda", dtype=torch.bfloat16), k + ) + + +@pytest.mark.parametrize("seq,k", [(4096, 1024), (16384, 2048), (16384, 1024)]) +def test_topk_bf16_signed_zero(seq: int, k: int) -> None: + """+0 and -0 compare equal as floats but sit in different histogram bins: every + mix of pivot sign / neighbours must still fill exactly k slots.""" + torch.manual_seed(seq + k) + x = torch.randn(32, seq, device="cuda", dtype=torch.bfloat16) + zero = torch.zeros_like(x) + neg_zero = -zero + c = torch.rand(32, seq, device="cuda") + _run_full_rows( + torch.where(c < 0.45, zero, torch.where(c < 0.9, neg_zero, x.abs())), k + ) + _run_full_rows(torch.where(c < 0.98, neg_zero, x.abs()), k) # pivot is -0 + _run_full_rows(torch.where(c < 0.98, zero, x.abs()), k) # pivot is +0 + _run_full_rows( + torch.where(c < 0.5, neg_zero, torch.where(c < 0.98, zero, x.abs())), k + ) + _run_full_rows( + torch.where(c < 0.9, -x.abs(), torch.where(c < 0.95, neg_zero, zero)), k + ) + + +@pytest.mark.parametrize("seq,k", [(4096, 1024), (16384, 2048)]) +def test_topk_bf16_bit_patterns(seq: int, k: int) -> None: + """Denormals, and the full bf16 range minus NaN.""" + torch.manual_seed(seq + k) + _run_full_rows( + torch.randint(0, 64, (32, seq), dtype=torch.int16, device="cuda").view( + torch.bfloat16 + ), + k, + ) + x = torch.randint( + -(2**15), 2**15, (32, seq), dtype=torch.int16, device="cuda" + ).view(torch.bfloat16) + _run_full_rows(torch.where(x.isnan(), torch.zeros_like(x), x), k) + _run_full_rows( + torch.randint(0, 0x7F80, (32, seq), dtype=torch.int16, device="cuda").view( + torch.bfloat16 + ), + k, + ) + + +@pytest.mark.parametrize( + "nan_bits,n_nan", [(0x7FC0, 5), (0x7FC0, 100), (0x7FC0, 600), (0xFFC0, 600)] +) +def test_topk_bf16_nan_scores(nan_bits: int, n_nan: int) -> None: + """NaN scores are never selected: the slots they would have taken are -1, the + rest are the top of the real scores, and nothing reads stale shared memory.""" + torch.manual_seed(nan_bits + n_nan) + batch, seq, k = 8, MAX_SEQ, 512 + scores = (torch.randn(batch, seq, device="cuda") * 2).to(torch.bfloat16) + bits = scores.view( + torch.int16 + ) # write the NaN by bit pattern: .item() would lose its sign + for b in range(batch): + bits[b, torch.randperm(seq, device="cuda")[:n_nan]] = nan_bits - ( + 0x10000 if nan_bits >= 0x8000 else 0 + ) + lens = torch.full((batch,), seq, dtype=torch.int32, device="cuda") + table = _identity_table(batch) + out = torch.full((batch, k), -7, dtype=torch.int32, device="cuda") + topk_transform_bf16_small(scores, lens, table, out, 8) + torch.cuda.synchronize() + for b in range(batch): + real = scores[b][~scores[b].isnan()].float() + n_real = ( + max(0, k - n_nan) if nan_bits < 0x8000 else k + ) # positive NaNs eat slots + chosen = out[b] >= 0 + assert int(chosen.sum()) == n_real, ( + f"row {b}: {int(chosen.sum())} selected, want {n_real}" + ) + assert bool((out[b][~chosen] == -1).all()), ( + f"row {b}: unselected slots are not -1" + ) + idx = out[b][chosen].long() + assert bool((idx < seq).all()) and idx.unique().numel() == n_real, ( + f"row {b}: bad index" + ) + got = scores[b, idx].float().sort(descending=True).values + assert torch.equal(got, real.topk(n_real).values), ( + f"row {b}: not the top real scores" + ) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-q"])) diff --git a/test/registered/kernels/benchmark/attention/bench_topk.py b/test/registered/kernels/benchmark/attention/bench_topk.py index a1bbb9815495..c76b2e78d02d 100644 --- a/test/registered/kernels/benchmark/attention/bench_topk.py +++ b/test/registered/kernels/benchmark/attention/bench_topk.py @@ -104,14 +104,16 @@ def fn(scores, seq_lens, offsets): 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 86adf27d9a4d..5a7e5e96d4d1 100644 --- a/test/registered/kernels/ops/attention/test_topk_v2.py +++ b/test/registered/kernels/ops/attention/test_topk_v2.py @@ -387,15 +387,56 @@ 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) - changed = (scores != before).cpu() - for i, (start, length) in enumerate(rows): - allowed = torch.zeros(scores.shape[1], dtype=torch.bool) - if length > k: - allowed[start - start % 4 : start] = True - stray = (changed[i] & ~allowed).nonzero().flatten().tolist() - 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]) From 9b17c0de2191563ff279342bf4207f35c1875c23 Mon Sep 17 00:00:00 2001 From: hnyls2002 Date: Tue, 15 Sep 2026 20:20:21 -0700 Subject: [PATCH 02/30] dsv4.1: skip unwritten Top-k plans in metadata comparisons --- python/sglang/test/kits/dsa_metadata_kit.py | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/python/sglang/test/kits/dsa_metadata_kit.py b/python/sglang/test/kits/dsa_metadata_kit.py index a3e5c36293bd..1568e8ac27a2 100644 --- a/python/sglang/test/kits/dsa_metadata_kit.py +++ b/python/sglang/test/kits/dsa_metadata_kit.py @@ -5,6 +5,7 @@ import torch +from sglang.kernels.ops.attention.dsv4.topk import _jit_topk_v2_module 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,12 @@ def assert_metadata_equal(test, actual, expected): for name, value in actual_buffers.items(): reference = expected_buffers[name] if name == "topk_v2_plan": + # Small-batch and non-cluster paths leave the entire plan unwritten. + # Probe the planner instead of duplicating its device-specific limits. + probe = torch.full_like(reference, -1) + _jit_topk_v2_module().topk_plan(expected.dsa_seqlens_expanded, probe, -1) + if probe[0, 1].item() == -1: + 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]) From f51894f2cbb72ae49593307ba6af1a061958862e Mon Sep 17 00:00:00 2001 From: BBuf <1182563586@qq.com> Date: Wed, 16 Sep 2026 12:21:40 +0800 Subject: [PATCH 03/30] fix dsv4 top-k edge cases and trim tests --- .../kernels/jit/csrc/deepseek_v4/topk_v2.cuh | 5 +- .../sgl_kernel/deepseek_v4/topk_impl.cuh | 16 +- .../sglang/kernels/ops/attention/dsv4/topk.py | 15 +- .../kernel/attention/test_amax_copy.py | 125 --------- .../test_dsv4_indexer_postprocess.py | 222 --------------- .../kernel/attention/test_sort_idx.py | 140 ---------- .../kernel/attention/test_topk_bf16.py | 264 ------------------ .../kernels/ops/attention/test_topk_v2.py | 11 + 8 files changed, 28 insertions(+), 770 deletions(-) delete mode 100644 test/registered/kernel/attention/test_amax_copy.py delete mode 100644 test/registered/kernel/attention/test_dsv4_indexer_postprocess.py delete mode 100644 test/registered/kernel/attention/test_sort_idx.py delete mode 100644 test/registered/kernel/attention/test_topk_bf16.py 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 bd30e5d50cc8..f79c7d99698f 100644 --- a/python/sglang/kernels/jit/csrc/deepseek_v4/topk_v2.cuh +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/topk_v2.cuh @@ -97,7 +97,10 @@ struct TopKPagedParams { return page_indices + batch_id * static_cast(topk); } SGL_DEVICE PageTransform get_transform(uint32_t batch_id) const { - return {page_table + batch_id * page_table_stride, page_bits, raw_indices + batch_id * static_cast(topk)}; + 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); 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 0b09b4335cd2..ab5fee576ca4 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,6 +23,7 @@ #include #include +#include #include #include #include @@ -84,15 +85,6 @@ 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); -// } - template SGL_DEVICE uint32_t extract_coarse_bin(float x) { static_assert(0 < kBits && kBits < 15); diff --git a/python/sglang/kernels/ops/attention/dsv4/topk.py b/python/sglang/kernels/ops/attention/dsv4/topk.py index 7a0186e0ed02..7193edb56097 100644 --- a/python/sglang/kernels/ops/attention/dsv4/topk.py +++ b/python/sglang/kernels/ops/attention/dsv4/topk.py @@ -36,17 +36,20 @@ def _jit_topk_v2_module(): # 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 - occ_8_2, occ_16_1 = 0, 0 try: occ_8_2 = get_max_active_clusters(8, occupancy=2) - # NOTE: cluster 16 might fail, but at least cluster 8 is ok + 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 - extra_cuda_cflags = [ - f"-DSGL_TOPK_V2_MAX_C8_OCC2={occ_8_2}", - f"-DSGL_TOPK_V2_MAX_C16_OCC1={occ_16_1}", - ] + 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"), diff --git a/test/registered/kernel/attention/test_amax_copy.py b/test/registered/kernel/attention/test_amax_copy.py deleted file mode 100644 index 603eabaac6eb..000000000000 --- a/test/registered/kernel/attention/test_amax_copy.py +++ /dev/null @@ -1,125 +0,0 @@ -"""Correctness tests for the DeepSeek-V4.1 JIT block-max ("amax") copy. - -``amax8_varlen`` writes, per row, the maximum of each block of 8 consecutive -fp32 scores for the first ``ceil(seq_lens[b] / 8)`` blocks, the last of them -``+inf`` (the newest block is always selected), and leaves everything else -untouched. It is the level-one key computation of the two-level sparse indexer: -the keys feed a block top-k, which is why rows with at most ``topk`` blocks may -be skipped (every block of such a row is selected anyway). -""" - -from __future__ import annotations - -import sys - -import pytest -import torch - -from sglang.kernels.ops.attention.dsv4.topk import amax8_varlen -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=20, stage="base-b-kernel-unit", runner_config="1-gpu-large") - -BLOCK = 8 -SENTINEL = -12345.0 -CONFIGS = [ - # (batch, seq): a single partial block, exact blocks, one CTA, many CTAs - (1, 1), - (1, 8), - (1, 9), - (3, 1000), - (4, 16389), - (8, 131072), - (2, 1048576), -] - - -def _keys_ref(scores: torch.Tensor, lens: torch.Tensor) -> torch.Tensor: - """Per row the block maxima of the first ceil(len / 8) blocks, the last +inf.""" - batch, width = scores.shape - nblocks = (width + BLOCK - 1) // BLOCK - padded = torch.nn.functional.pad( - scores, (0, nblocks * BLOCK - width), value=-torch.inf - ) - keys = padded.view(batch, nblocks, BLOCK).amax(-1) - last = (lens.long() + BLOCK - 1) // BLOCK - 1 - keys[torch.arange(batch, device=scores.device), last] = torch.inf - return keys - - -def _check( - out: torch.Tensor, scores: torch.Tensor, lens: torch.Tensor, skip=None -) -> None: - ref = _keys_ref(scores, lens) - for b in range(scores.shape[0]): - n = (int(lens[b]) + BLOCK - 1) // BLOCK - if skip is not None and skip[b]: - assert torch.all(out[b] == SENTINEL), f"row {b} was skipped but written" - continue - assert torch.equal(out[b, :n], ref[b, :n]), f"row {b} keys differ" - assert torch.all(out[b, n:] == SENTINEL), f"row {b} written past its keys" - - -def _inputs(batch: int, seq: int, ragged: bool, stride_pad: int = 0): - torch.manual_seed(batch * 977 + seq) - lens = torch.full((batch,), seq, dtype=torch.int32, device="cuda") - if ragged: - lens = torch.randint(1, seq + 1, (batch,), dtype=torch.int32, device="cuda") - lens[0] = seq - stride = (seq + BLOCK - 1) // BLOCK * BLOCK + stride_pad - storage = torch.randn(batch, stride, device="cuda") * 10 - scores = storage[:, :seq] - return scores, lens - - -@pytest.mark.parametrize("ragged", [False, True]) -@pytest.mark.parametrize("batch,seq", CONFIGS) -def test_matches_block_max(batch: int, seq: int, ragged: bool): - scores, lens = _inputs(batch, seq, ragged) - out = torch.full((batch, (seq + BLOCK - 1) // BLOCK), SENTINEL, device="cuda") - assert amax8_varlen(scores, lens, out=out) is out - _check(out, scores, lens) - - -def test_strided_views(): - """Score rows padded past the row width and a key buffer wider than needed.""" - scores, lens = _inputs(4, 16389, ragged=True, stride_pad=248) - nblocks = (16389 + BLOCK - 1) // BLOCK - out = torch.full((4, nblocks + 5), SENTINEL, device="cuda") - amax8_varlen(scores, lens, out=out[:, :nblocks]) - _check(out, scores, lens) - - -def test_allocates_output(): - scores, lens = _inputs(3, 20000, ragged=True) - out = amax8_varlen(scores, lens) - assert out.shape == (3, (20000 + BLOCK - 1) // BLOCK) and out.dtype == torch.float32 - ref = _keys_ref(scores, lens) - for b in range(3): - n = (int(lens[b]) + BLOCK - 1) // BLOCK - assert torch.equal(out[b, :n], ref[b, :n]) - - -def test_max_seqlen_sizes_output(): - """A batch whose rows are all shorter than the padded width only needs - ceil(max_seqlen / 8) keys per row.""" - scores, _ = _inputs(4, 65536, ragged=False) - lens = torch.tensor([40000, 1, 16384, 39999], dtype=torch.int32, device="cuda") - out = amax8_varlen(scores, lens, max_seqlen=40000) - assert out.shape == (4, 5000) - ref = _keys_ref(scores, lens) - for b in range(4): - n = (int(lens[b]) + BLOCK - 1) // BLOCK - assert torch.equal(out[b, :n], ref[b, :n]) - - -def test_topk_skips_rows_that_fit(): - scores, _ = _inputs(4, 40000, ragged=False) - lens = torch.tensor([40000, 16384, 16385, 100], dtype=torch.int32, device="cuda") - out = torch.full((4, 5000), SENTINEL, device="cuda") - amax8_varlen(scores, lens, 2048, out=out) - _check(out, scores, lens, skip=[False, True, False, True]) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/kernel/attention/test_dsv4_indexer_postprocess.py b/test/registered/kernel/attention/test_dsv4_indexer_postprocess.py deleted file mode 100644 index af9da90c4203..000000000000 --- a/test/registered/kernel/attention/test_dsv4_indexer_postprocess.py +++ /dev/null @@ -1,222 +0,0 @@ -import unittest - -import torch - -from sglang.kernels.ops.attention.dsv4.candidate_blocks import ( - candidate_block_logits, - candidate_row_lens, -) -from sglang.kernels.ops.attention.dsv4.indexer_postprocess import filter_topk_pages -from sglang.test.ci.ci_register import register_cuda_ci -from sglang.test.test_utils import CustomTestCase - -register_cuda_ci(est_time=20, stage="base-b-kernel-unit", runner_config="1-gpu-large") - - -def reference_pages(scores, indices, pages, page_size): - cols = indices.to(torch.int64) - value = scores.gather(1, cols.clamp(0, scores.shape[1] - 1)) - valid = (cols >= 0) & (cols < scores.shape[1]) & (value > -torch.inf) - raw = indices.masked_fill(~valid, -1) - safe = raw.clamp_min(0).to(torch.int64) - slots = pages.gather(1, safe // page_size) * page_size + safe % page_size - return torch.where(raw >= 0, slots, -1).to(torch.int32), raw - - -class TestIndexerPostprocess(CustomTestCase): - def test_unfiltered_verify_matches_candidate_chain(self): - from sglang.kernels.ops.attention.dsv4.topk import ( - plan_topk_v2, - topk_transform_paged_v2, - ) - - # Six causal rows per request. The last request ends exactly at the - # candidate budget; capacity and unread logits extend well past it. - torch.manual_seed(123) - rows, width, page_size = 384, 32768, 64 - base = torch.tensor([0, 506, 4096, 16378], device="cuda", dtype=torch.int32) - lens = (base.repeat(16)[:, None] + torch.arange(1, 7, device="cuda")).flatten() - lens = lens.to(torch.int32) - source = torch.randn(rows, width, device="cuda").relu_() - consumer = torch.randn_like(source).relu_() - cols = torch.arange(width, device="cuda")[None, :] - source.masked_fill_(cols >= lens[:, None], 1e6) - consumer.masked_fill_(cols >= lens[:, None], 1e6) - pages = torch.randint( - 0, 100000, (rows, width // page_size), device="cuda", dtype=torch.int32 - ) - source_masked, keep = candidate_block_logits( - source, lens, topk_blocks=2048, block_size=8, published=None - ) - consumer_masked, _ = candidate_block_logits( - consumer, lens, topk_blocks=2048, block_size=8, published=keep - ) - plan = plan_topk_v2(lens) - for original, masked in ((source, source_masked), (consumer, consumer_masked)): - with self.subTest(source=original is source): - old = torch.empty((rows, 512), dtype=torch.int32, device="cuda") - new = torch.empty_like(old) - raw = torch.empty_like(old) - new_raw = torch.empty_like(old) - topk_transform_paged_v2(masked, lens, pages, old, page_size, plan, raw) - filter_topk_pages(masked, raw, pages, old, page_size) - topk_transform_paged_v2( - original, lens, pages, new, page_size, plan, new_raw - ) - # Top-k v2 uses a persistent work queue and does not promise - # output order. Compare the selected multiset, including -1 - # padding, in both logical-position and physical-page space. - torch.testing.assert_close( - new.sort(-1).values, old.sort(-1).values, rtol=0, atol=0 - ) - torch.testing.assert_close( - new_raw.sort(-1).values, raw.sort(-1).values, rtol=0, atol=0 - ) - - def test_candidate_row_lens(self): - lens = torch.tensor( - [1, 7, 8, 9, 300, 16383, 16384, 16385, 16392, 40000, 1048576, 1048571], - dtype=torch.int32, - device="cuda", - ) - for topk in (2048, 256, 1): - nblocks, valid = candidate_row_lens(lens, topk) - ref_blocks = (lens + 7) // 8 - kept = ref_blocks.clamp_max(topk) - ref_valid = 8 * (kept - 1) + (lens - 1) % 8 + 1 - self.assertEqual(nblocks.dtype, torch.int32) - self.assertTrue(torch.equal(nblocks, ref_blocks)) - self.assertTrue(torch.equal(valid, ref_valid)) - # a row with all its blocks kept is the whole row - fits = ref_blocks <= topk - self.assertTrue(torch.equal(valid[fits], lens[fits])) - nblocks, valid = candidate_row_lens( - torch.zeros(3, dtype=torch.int32, device="cuda"), 2048 - ) - self.assertTrue(torch.equal(nblocks, torch.zeros_like(nblocks))) - self.assertTrue(torch.equal(valid, torch.zeros_like(valid))) - - def test_filter_and_page_mapping(self): - torch.manual_seed(91) - for rows in (1, 6, 64): - for width in (1, 65, 4096): - for dtype in (torch.int32, torch.int64): - # Row-strided tensors match sliced metadata buffers. - scores = torch.randn(rows, width + 7, device="cuda")[:, :width] - indices = torch.randint( - -2, width + 2, (rows, 520), device="cuda", dtype=dtype - )[:, :512] - indices[:, :5] = 0 - scores[0, 0] = -torch.inf - if rows > 1: - scores[1, 0] = torch.nan - scores[2, 0] = torch.inf - scores[3, :] = -torch.inf - pages = torch.randint( - 0, - 100000, - (rows, (width + 63) // 64 + 3), - device="cuda", - dtype=torch.int32, - )[:, :-3] - out = torch.empty(rows, 520, device="cuda", dtype=torch.int32)[ - :, :512 - ] - raw = torch.empty_like(indices) - expected, expected_raw = reference_pages(scores, indices, pages, 64) - for write_raw in (False, True): - filter_topk_pages( - scores, indices, pages, out, 64, raw if write_raw else None - ) - torch.testing.assert_close(out, expected, rtol=0, atol=0) - if write_raw: - torch.testing.assert_close( - raw, expected_raw, rtol=0, atol=0 - ) - - def test_graph_replay(self): - rows, width = 6, 4096 - scores = torch.randn(rows, width, device="cuda") - indices = torch.randint( - -1, width, (rows, 512), device="cuda", dtype=torch.int32 - ) - pages = torch.randint( - 0, 100000, (rows, width // 64), device="cuda", dtype=torch.int32 - ) - out, raw = torch.empty_like(indices), torch.empty_like(indices) - filter_topk_pages(scores, indices, pages, out, 64, raw) - graph = torch.cuda.CUDAGraph() - with torch.cuda.graph(graph): - filter_topk_pages(scores, indices, pages, out, 64, raw) - for _ in range(3): - scores.normal_() - scores[:, :32] = -torch.inf - indices.random_(-1, width) - pages.random_(0, 100000) - graph.replay() - expected, expected_raw = reference_pages(scores, indices, pages, 64) - torch.testing.assert_close(out, expected, rtol=0, atol=0) - torch.testing.assert_close(raw, expected_raw, rtol=0, atol=0) - - def test_candidate_publication(self): - torch.manual_seed(912) - for width, group, topk in ( - (65, 32, 8), - (4096, 128, 8), - (32768, 2048, 8), - (1048576, 2048, 8), - (4096, 8, 2048), - (1048576, 8, 2048), - ): - rows = 6 - x = torch.randn(rows, width + 9, device="cuda")[:, :width] - lengths = torch.tensor( - [0, 1, width // 2, width - 1, width, width], - device="cuda", - dtype=torch.int32, - ) - # Include ties, NaN and a final partial block. - x[3, :] = 0 - x[4, 0] = torch.nan - x[5, :] = -torch.inf - masked = x.masked_fill( - torch.arange(width, device="cuda")[None, :] >= lengths[:, None], - -torch.inf, - ) - blocks = (width + group - 1) // group - padded = torch.nn.functional.pad( - masked, (0, blocks * group - width), value=-torch.inf - ) - scores = padded.reshape(rows, blocks, group).max(-1).values - last = (lengths - 1) // group - scores = torch.where( - (torch.arange(blocks, device="cuda")[None, :] == last[:, None]) - & (lengths[:, None] > 0), - torch.inf, - scores, - ) - selected = scores.topk(min(topk, blocks), dim=-1) - keep = ( - torch.zeros_like(scores, dtype=torch.bool) - .scatter_(1, selected.indices, selected.values > -torch.inf) - .repeat_interleave(group, dim=-1)[:, :width] - ) - got, published = candidate_block_logits( - x, lengths, topk_blocks=topk, block_size=group, published=None - ) - torch.testing.assert_close(got, masked, rtol=0, atol=0, equal_nan=True) - torch.testing.assert_close(published, keep, rtol=0, atol=0) - consumer, _ = candidate_block_logits( - x, lengths, topk_blocks=topk, block_size=group, published=published - ) - torch.testing.assert_close( - consumer, - masked.masked_fill(~keep, -torch.inf), - rtol=0, - atol=0, - equal_nan=True, - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/kernel/attention/test_sort_idx.py b/test/registered/kernel/attention/test_sort_idx.py deleted file mode 100644 index 8b3bef80122e..000000000000 --- a/test/registered/kernel/attention/test_sort_idx.py +++ /dev/null @@ -1,140 +0,0 @@ -"""Correctness tests for the DeepSeek-V4.1 JIT candidate block table sort. - -``sort_candidate_blocks`` turns a row's selected block ids (any order, ``-1`` -padded) in place into the ascending, ``INT32_MAX``-padded table DeepGEMM's -sparse indexer schedule reads, and writes the same blocks as pool slots / 8 -through the row's index page table. Rows with at most ``k`` blocks get the -identity table. The kernel is a bitmap counting sort whose dense words are -drained by a block-wide queue, so the id distributions below cover sparse, -clustered and completely full words. -""" - -from __future__ import annotations - -import sys - -import pytest -import torch - -from sglang.kernels.ops.attention.dsv4.topk import ( - sort_candidate_blocks, - transform_candidate_blocks, -) -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=20, stage="base-b-kernel-unit", runner_config="1-gpu-large") - -BLOCK = 8 -PAD = torch.iinfo(torch.int32).max -MAX_BLOCKS = 1 << 17 - - -def _select(nblocks: int, k: int, dist: str, gen: torch.Generator) -> torch.Tensor: - """k distinct block ids below nblocks (all of them if nblocks <= k), shuffled.""" - n = min(k, nblocks) - if dist == "uniform": - ids = torch.randperm(nblocks, device="cuda", generator=gen)[:n] - elif dist == "clustered": # runs of 64 consecutive blocks mixed with scattered ones - runs = max(n // 96, 1) - starts = torch.randint( - 0, max(nblocks - 64, 1), (runs,), device="cuda", generator=gen - ) - runs_ids = ( - starts[:, None] + torch.arange(64, device="cuda")[None, :] - ).flatten() - rest = torch.randperm(nblocks, device="cuda", generator=gen)[:n] - ids = torch.unique(torch.cat([runs_ids, rest]))[:n] - elif dist == "newest": # the tail of the row: every word full - ids = torch.arange(nblocks - n, nblocks, device="cuda") - else: - raise ValueError(dist) - ids = ids.to(torch.int32) - ids = ids[torch.randperm(ids.numel(), device="cuda", generator=gen)] - return torch.nn.functional.pad(ids, (0, k - n), value=-1) - - -def _case(lens, k, dist, page_size, seed=0): - gen = torch.Generator(device="cuda").manual_seed(seed) - lens_t = torch.tensor(lens, dtype=torch.int32, device="cuda") - max_pages = (max(lens) + page_size - 1) // page_size - page_table = torch.stack( - [torch.randperm(max_pages, device="cuda", generator=gen) for _ in lens] - ).to(torch.int32) - blocks = torch.stack( - [_select((l + BLOCK - 1) // BLOCK, k, dist, gen) for l in lens] - ) - return lens_t, page_table, blocks - - -def _check(blocks, pages, selected, lens, page_table, page_size, k): - bpp = page_size // BLOCK - for b, seq in enumerate(lens): - nblocks = (seq + BLOCK - 1) // BLOCK - n = min(k, nblocks) - if nblocks <= k: - ref = torch.arange(n, device="cuda", dtype=torch.int32) - else: - ref = selected[b][selected[b] >= 0].sort().values - assert torch.equal(blocks[b, :n], ref), f"row {b} ids" - assert torch.all(blocks[b, n:] == PAD), f"row {b} padding" - ref_pages = page_table[b][ref // bpp] * bpp + ref % bpp - assert torch.equal(pages[b, :n], ref_pages), f"row {b} slots" - assert torch.all(pages[b, n:] == PAD), f"row {b} slot padding" - - -@pytest.mark.parametrize("dist", ["uniform", "clustered", "newest"]) -@pytest.mark.parametrize("page_size", [64, 128]) -@pytest.mark.parametrize( - "lens,k", - [ - # rows that fit (identity), rows just above k, long rows up to 1M tokens - ([1, 8, 16384, 16385, 16392], 2048), - ([40000, 131072, 65537, 1048576], 2048), - ([1048576, 1048571, 300000], 2048), - ([2049, 9000, 20000], 256), - ], -) -def test_matches_sorted_selection(lens, k, dist, page_size): - lens_t, page_table, blocks = _case(lens, k, dist, page_size) - selected = blocks.clone() - pages = sort_candidate_blocks(blocks, lens_t, page_table, page_size) - _check(blocks, pages, selected, lens, page_table, page_size, k) - - -def test_dense_words_only(): - """Every selected id inside k / 32 full words, in random order: the queue path only.""" - lens = [MAX_BLOCKS * BLOCK, 100000] - lens_t, page_table, _ = _case(lens, 2048, "uniform", 128) - gen = torch.Generator(device="cuda").manual_seed(3) - rows = [] - for seq in lens: - nblocks = (seq + BLOCK - 1) // BLOCK - start = (nblocks - 2048) // 2 // 32 * 32 - ids = torch.arange(start, start + 2048, device="cuda", dtype=torch.int32) - rows.append(ids[torch.randperm(2048, device="cuda", generator=gen)]) - blocks = torch.stack(rows) - selected = blocks.clone() - pages = torch.empty_like(blocks) - assert ( - sort_candidate_blocks(blocks, lens_t, page_table, 128, out_pages=pages) is pages - ) - _check(blocks, pages, selected, lens, page_table, 128, 2048) - - -@pytest.mark.parametrize("page_size", [64, 128]) -def test_page_transform_of_sorted_ids(page_size): - """Ascending ids in, INT32_MAX past the row's count (a DeepSelect-style - input): the same pages as the sort, blocks untouched.""" - lens = [1, 16384, 16385, 131072, 1048576] - k = 2048 - lens_t, page_table, blocks = _case(lens, k, "clustered", page_size, seed=11) - ref_pages = sort_candidate_blocks(blocks.clone(), lens_t, page_table, page_size) - ascending = torch.where(blocks < 0, PAD, blocks).sort(dim=1).values - kept = ascending.clone() - pages = transform_candidate_blocks(ascending, lens_t, page_table, page_size) - assert torch.equal(ascending, kept) - assert torch.equal(pages, ref_pages) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/kernel/attention/test_topk_bf16.py b/test/registered/kernel/attention/test_topk_bf16.py deleted file mode 100644 index 6cdce0fa47d5..000000000000 --- a/test/registered/kernel/attention/test_topk_bf16.py +++ /dev/null @@ -1,264 +0,0 @@ -"""Correctness tests for the DeepSeek-V4.1 JIT bf16 top-k transform. - -``topk_transform_bf16_small`` selects the per-row top-k of bf16 ``scores`` within each -row's ``seq_lens`` (rows of at most 16384, ``k`` at most 2048) and writes the -page-table transform of the selected indices, ``-1`` past ``min(k, seq_len)``, in -no particular order. It is the consumer selection of the two-level sparse -indexer: there the "page table" is the per-row physical block table at page -size 8. - -Selection is exact, so it is validated against ``torch.topk`` on the multiset of -selected values (elements of equal value may swap). NaN scores are not -supported and never generated here. -""" - -from __future__ import annotations - -import sys - -import pytest -import torch - -from sglang.kernels.ops.attention.dsv4.topk import topk_transform_bf16_small -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="1-gpu-large") - -MAX_SEQ = 16384 -CONFIGS = [ - # (batch, seq): short rows select everything, full rows the 16384 maximum - (1, 1), - (4, 300), - (8, 512), - (8, 513), - (16, 4097), - (32, 16000), - (32, MAX_SEQ), - (148, MAX_SEQ), - (300, MAX_SEQ), -] - - -def _rows(batch: int, seq: int, k: int, ties: bool, ragged: bool): - torch.manual_seed(batch * 131 + seq) - lens = torch.full((batch,), seq, dtype=torch.int32, device="cuda") - if ragged: - lens = torch.randint(1, seq + 1, (batch,), dtype=torch.int32, device="cuda") - lens[0] = seq - scores = torch.randn(batch, MAX_SEQ, device="cuda") * 2 - if ties: - scores = (scores * 2).round() / 2 # a handful of distinct values per row - # everything past a row is garbage that must never be selected - scores.masked_fill_( - torch.arange(MAX_SEQ, device="cuda")[None, :] >= lens[:, None], 1e4 - ) - return scores.to(torch.bfloat16), lens - - -def _identity_table(batch: int, page_size: int = 8) -> torch.Tensor: - return torch.arange(MAX_SEQ // page_size, device="cuda", dtype=torch.int32).repeat( - batch, 1 - ) - - -def _check(scores, lens, table, page_size, out): - """Every row: min(k, len) valid slots, the rest -1; slots invert to unique - in-row indices whose values are torch's top-k values.""" - k = out.shape[1] - nblocks = MAX_SEQ // page_size - inv = torch.empty_like(table) - inv.scatter_( - 1, - table.long(), - torch.arange(nblocks, device="cuda", dtype=torch.int32).expand_as(table), - ) - for b in range(scores.shape[0]): - n = min(k, int(lens[b])) - chosen = out[b] >= 0 - assert int(chosen.sum()) == n, ( - f"row {b}: {int(chosen.sum())} selected, want {n}" - ) - slots = out[b][chosen].long() - idx = inv[b][slots // page_size].long() * page_size + slots % page_size - assert bool((idx < int(lens[b])).all()), f"row {b}: index past the row" - assert idx.unique().numel() == n, f"row {b}: duplicate index" - got = scores[b, idx].float().sort(descending=True).values - ref = scores[b, : int(lens[b])].float().topk(n).values - assert torch.equal(got, ref), f"row {b}: selected values differ from torch.topk" - - -def _run_full_rows(scores: torch.Tensor, k: int) -> None: - """Whole rows of `scores` (any width up to MAX_SEQ) through the identity table.""" - batch, seq = scores.shape - lens = torch.full((batch,), seq, dtype=torch.int32, device="cuda") - table = _identity_table(batch) - out = torch.full((batch, k), -7, dtype=torch.int32, device="cuda") - topk_transform_bf16_small(scores, lens, table, out, 8) - torch.cuda.synchronize() - _check(scores, lens, table, 8, out) - - -def _randn_aligned(batch: int, seq: int) -> torch.Tensor: - # rows must start vector-aligned, so odd widths come from a wider padded tensor - padded = (seq + 7) // 8 * 8 - return torch.randn(batch, padded, device="cuda", dtype=torch.bfloat16)[:, :seq] - - -@pytest.mark.parametrize("page_mode", ["identity", "perm"]) -@pytest.mark.parametrize("ties", [False, True]) -@pytest.mark.parametrize("k", [512, 2048]) -@pytest.mark.parametrize("batch,seq", CONFIGS) -def test_topk_bf16(batch: int, seq: int, k: int, ties: bool, page_mode: str) -> None: - page_size = 8 - scores, lens = _rows(batch, seq, k, ties, ragged=False) - nblocks = MAX_SEQ // page_size - if page_mode == "identity": - table = _identity_table(batch, page_size) - else: - table = torch.stack( - [torch.randperm(nblocks, device="cuda") for _ in range(batch)] - ).to(torch.int32) - out = torch.full((batch, k), -7, dtype=torch.int32, device="cuda") - topk_transform_bf16_small(scores, lens, table, out, page_size) - torch.cuda.synchronize() - _check(scores, lens, table, page_size, out) - - -@pytest.mark.parametrize("page_size", [8, 64]) -def test_topk_bf16_ragged_lengths(page_size: int) -> None: - batch, k = 64, 512 - scores, lens = _rows(batch, MAX_SEQ, k, ties=False, ragged=True) - nblocks = MAX_SEQ // page_size - table = torch.stack( - [torch.randperm(nblocks, device="cuda") for _ in range(batch)] - ).to(torch.int32) - out = torch.full((batch, k), -7, dtype=torch.int32, device="cuda") - topk_transform_bf16_small(scores, lens, table, out, page_size) - torch.cuda.synchronize() - _check(scores, lens, table, page_size, out) - - -def test_topk_bf16_padded_output() -> None: - """The output may be a view whose last dimension was padded.""" - batch, k, page_size = 3, 512, 8 - scores, lens = _rows(batch, MAX_SEQ, k, ties=False, ragged=False) - table = _identity_table(batch, page_size) - buf = torch.full((batch, k + 64), -7, dtype=torch.int32, device="cuda") - out = buf[:, :k] - topk_transform_bf16_small(scores, lens, table, out, page_size) - torch.cuda.synchronize() - assert bool((buf[:, k:] == -7).all()), "wrote past the view" - _check(scores, lens, table, page_size, out) - - -@pytest.mark.parametrize("k", [1, 512, 2048]) -@pytest.mark.parametrize("seq", [1000, 3000, 8191, 12345]) -def test_topk_bf16_odd_widths(seq: int, k: int) -> None: - """Rows whose width is not a multiple of the vector: the last vector is partial.""" - torch.manual_seed(seq * 7 + k) - _run_full_rows(_randn_aligned(33, seq), k) - - -@pytest.mark.parametrize( - "seq,k", [(4096, 1024), (16384, 2048), (16384, 512), (1000, 512)] -) -def test_topk_bf16_heavy_ties(seq: int, k: int) -> None: - """Heavy ties, all-equal rows and all -inf rows exercise the equal-quota path.""" - torch.manual_seed(seq + k) - _run_full_rows(torch.randint(0, 5, (32, seq), device="cuda").bfloat16(), k) - _run_full_rows(torch.randint(-3, 3, (32, seq), device="cuda").bfloat16(), k) - _run_full_rows(torch.full((32, seq), 1.5, device="cuda", dtype=torch.bfloat16), k) - _run_full_rows( - torch.full((32, seq), float("-inf"), device="cuda", dtype=torch.bfloat16), k - ) - - -@pytest.mark.parametrize("seq,k", [(4096, 1024), (16384, 2048), (16384, 1024)]) -def test_topk_bf16_signed_zero(seq: int, k: int) -> None: - """+0 and -0 compare equal as floats but sit in different histogram bins: every - mix of pivot sign / neighbours must still fill exactly k slots.""" - torch.manual_seed(seq + k) - x = torch.randn(32, seq, device="cuda", dtype=torch.bfloat16) - zero = torch.zeros_like(x) - neg_zero = -zero - c = torch.rand(32, seq, device="cuda") - _run_full_rows( - torch.where(c < 0.45, zero, torch.where(c < 0.9, neg_zero, x.abs())), k - ) - _run_full_rows(torch.where(c < 0.98, neg_zero, x.abs()), k) # pivot is -0 - _run_full_rows(torch.where(c < 0.98, zero, x.abs()), k) # pivot is +0 - _run_full_rows( - torch.where(c < 0.5, neg_zero, torch.where(c < 0.98, zero, x.abs())), k - ) - _run_full_rows( - torch.where(c < 0.9, -x.abs(), torch.where(c < 0.95, neg_zero, zero)), k - ) - - -@pytest.mark.parametrize("seq,k", [(4096, 1024), (16384, 2048)]) -def test_topk_bf16_bit_patterns(seq: int, k: int) -> None: - """Denormals, and the full bf16 range minus NaN.""" - torch.manual_seed(seq + k) - _run_full_rows( - torch.randint(0, 64, (32, seq), dtype=torch.int16, device="cuda").view( - torch.bfloat16 - ), - k, - ) - x = torch.randint( - -(2**15), 2**15, (32, seq), dtype=torch.int16, device="cuda" - ).view(torch.bfloat16) - _run_full_rows(torch.where(x.isnan(), torch.zeros_like(x), x), k) - _run_full_rows( - torch.randint(0, 0x7F80, (32, seq), dtype=torch.int16, device="cuda").view( - torch.bfloat16 - ), - k, - ) - - -@pytest.mark.parametrize( - "nan_bits,n_nan", [(0x7FC0, 5), (0x7FC0, 100), (0x7FC0, 600), (0xFFC0, 600)] -) -def test_topk_bf16_nan_scores(nan_bits: int, n_nan: int) -> None: - """NaN scores are never selected: the slots they would have taken are -1, the - rest are the top of the real scores, and nothing reads stale shared memory.""" - torch.manual_seed(nan_bits + n_nan) - batch, seq, k = 8, MAX_SEQ, 512 - scores = (torch.randn(batch, seq, device="cuda") * 2).to(torch.bfloat16) - bits = scores.view( - torch.int16 - ) # write the NaN by bit pattern: .item() would lose its sign - for b in range(batch): - bits[b, torch.randperm(seq, device="cuda")[:n_nan]] = nan_bits - ( - 0x10000 if nan_bits >= 0x8000 else 0 - ) - lens = torch.full((batch,), seq, dtype=torch.int32, device="cuda") - table = _identity_table(batch) - out = torch.full((batch, k), -7, dtype=torch.int32, device="cuda") - topk_transform_bf16_small(scores, lens, table, out, 8) - torch.cuda.synchronize() - for b in range(batch): - real = scores[b][~scores[b].isnan()].float() - n_real = ( - max(0, k - n_nan) if nan_bits < 0x8000 else k - ) # positive NaNs eat slots - chosen = out[b] >= 0 - assert int(chosen.sum()) == n_real, ( - f"row {b}: {int(chosen.sum())} selected, want {n_real}" - ) - assert bool((out[b][~chosen] == -1).all()), ( - f"row {b}: unselected slots are not -1" - ) - idx = out[b][chosen].long() - assert bool((idx < seq).all()) and idx.unique().numel() == n_real, ( - f"row {b}: bad index" - ) - got = scores[b, idx].float().sort(descending=True).values - assert torch.equal(got, real.topk(n_real).values), ( - f"row {b}: not the top real scores" - ) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-q"])) diff --git a/test/registered/kernels/ops/attention/test_topk_v2.py b/test/registered/kernels/ops/attention/test_topk_v2.py index 5a7e5e96d4d1..fea799717541 100644 --- a/test/registered/kernels/ops/attention/test_topk_v2.py +++ b/test/registered/kernels/ops/attention/test_topk_v2.py @@ -387,6 +387,17 @@ 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 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) + if length > k: + allowed[start - start % 4 : start] = True + stray = (changed[i] & ~allowed).nonzero().flatten().tolist() + 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() From ed1c649cb48fcc39fce0b356523ef92a0d66c318 Mon Sep 17 00:00:00 2001 From: hnyls2002 Date: Tue, 15 Sep 2026 15:21:04 -0700 Subject: [PATCH 04/30] dsv4.1: extract communication kernels and wrappers --- .../csrc/distributed/all_reduce_fusion.cuh | 767 ++++++++++++++++ .../csrc/distributed/custom_all_reduce.cuh | 11 +- .../jit/csrc/distributed/nvlink_comm.cuh | 851 ++++++++++++++++++ .../include/sgl_kernel/distributed/ptx.cuh | 8 +- .../ops/communication/all_reduce_fusion.py | 256 ++++++ .../ops/communication/all_reduce_mhc.py | 318 +++++++ .../kernels/ops/communication/nvlink_comm.py | 282 ++++++ .../test_moe_finalize_all_reduce.py | 312 +++++++ .../kernel/communication/test_nvlink_comm.py | 202 +++++ 9 files changed, 3002 insertions(+), 5 deletions(-) create mode 100644 python/sglang/kernels/jit/csrc/distributed/all_reduce_fusion.cuh create mode 100644 python/sglang/kernels/jit/csrc/distributed/nvlink_comm.cuh create mode 100644 python/sglang/kernels/ops/communication/all_reduce_fusion.py create mode 100644 python/sglang/kernels/ops/communication/all_reduce_mhc.py create mode 100644 python/sglang/kernels/ops/communication/nvlink_comm.py create mode 100644 test/registered/kernel/communication/test_moe_finalize_all_reduce.py create mode 100644 test/registered/kernel/communication/test_nvlink_comm.py diff --git a/python/sglang/kernels/jit/csrc/distributed/all_reduce_fusion.cuh b/python/sglang/kernels/jit/csrc/distributed/all_reduce_fusion.cuh new file mode 100644 index 000000000000..efdb57202ffd --- /dev/null +++ b/python/sglang/kernels/jit/csrc/distributed/all_reduce_fusion.cuh @@ -0,0 +1,767 @@ +// Fused deferred-MoE finalize -> 1shot lamport push all-reduce [-> RMSNorm] +// over the CustomAllReduceV2 push plane, for decode-sized batches (bf16). +// +// A generalisation of the K3 `finalize_push_norm` kernel +// (csrc/kimi_k3/comm/ar_fusion.cuh): the hidden width, top_k and cluster +// geometry are template parameters chosen from Python, the shared-expert add +// and the RMSNorm epilogue are optional, and the result goes to a separate +// output tensor. Per token row t: +// +// local[t] = sum_k expert_weights[t, k] * gemm2_out[idx[t * top_k + k]] +// (+ shared_output[t]) -- stage 1, registers only +// out[t] = sum over ranks of local[t] -- stage 2 +// out[t] = out[t] * rsqrt(mean(out[t]^2) + eps) * w -- kNorm only +// +// `idx == -1` marks a dropped slot (EP: the token was routed to an expert +// that is not local) and contributes nothing. Accumulation is fp32; the bf16 +// rounding points are exactly the unfused path's: the routed combine (what +// TRT-LLM's finalize returns), the `+ shared` (torch's bf16 add) and the +// all-reduce output, so the plain (kNorm=false) result equals TRT-LLM finalize +// -> `shared.add_(routed)` -> fp32-accumulating bf16 all-reduce in rank order, +// and the staged vector is bit-identical to what the unfused path reduces. +// +// The rank-local finalize never materializes in global memory: each thread +// computes one 16B vector of it and pushes it straight into every peer's push +// slot with unicast `st.relaxed.sys` stores, exactly like the generic +// `all_reduce_1shot_push_kernel`, so no multicast mapping is required. +// +// Push-plane protocol (see include/sgl_kernel/distributed/communicator.cuh): +// * every rank owns 2 phases x kWorldSize slots of `slot_bytes`; a round +// uses phase `counter & 1`, producer r writes slot r of every peer, the +// consumer polls its own kWorldSize slots until no +0.0 marker remains, +// reduces, and restores the +0.0 markers before it exits; +// * +0.0 payload words are remapped to -0.0 (numerically identical) so a +// written word is never 0 and `word == 0` means "not arrived yet"; +// * the phase counters are per block of the GENERIC push kernel, which +// launches `num_blocks` blocks and flips one counter each. This kernel +// uses one counter per row cluster (flipped by the cluster's leader block +// after a cluster barrier, since every block of the cluster reads it) and +// a trailing "bumper" cluster flips every remaining one, so the whole +// array keeps one parity and the two kernel families can share the plane +// freely (single-stream calls are serialized); +// * every rank must call with the same num_tokens / hidden / top_k / epilogue: +// slots are addressed by 16B vector index of the [T, hidden] row view. +#include +#include +#include + +#include +#include +#include +#include +#include +#include + +#include +#include + +#include + +#include +#include +#include +#include + +namespace sglang { + +using device::distributed::PushWorkSpace; +using host::distributed::CommunicatorRef; + +/// One 16B staging vector (8 bf16) viewed as the 4 u32 words the lamport marker +/// protocol tests, matching the generic push kernel's `LamportTrait`. +using Lamport = device::distributed::LamportTrait; +using StageVec = device::AlignedVector; + +SGL_DEVICE void barrier_cluster_arrive_relaxed() { + asm volatile("barrier.cluster.arrive.relaxed.aligned;" ::: "memory"); +} + +SGL_DEVICE void barrier_cluster_wait() { + asm volatile("barrier.cluster.wait.aligned;" ::: "memory"); +} + +template +struct FinalizeAllReduceParams { + bf16_t* out; // [num_tokens, kHiddenDim], output-only + const bf16_t* gemm2; // [P, kHiddenDim], permuted / padded rows + const int32_t* idx; // [num_tokens * kTopK], -1 = dropped slot + const WeightT* weights; // [num_tokens, kTopK], scaling already folded in + const bf16_t* shared; // [num_tokens, kHiddenDim] (kHasShared only) + const bf16_t* norm_weight; // [kHiddenDim] (kNorm only) + float norm_eps; // kNorm only + // Caller's promise that everything this kernel reads before its PDL wait is + // complete when the preceding kernel merely *triggers*: no all-reduce on this + // plane right before it, and the routing metadata's producers finished (not + // just the immediate predecessor -- PDL completion is not transitive through + // early-triggering kernels). False (the safe default) waits first. + bool prefetch_metadata; + uint32_t rank; + uint32_t num_tokens; + uint32_t num_push_counters; // full counter array size (bumper range end) + PushWorkSpace ws; + bf16_t* mhc_out = nullptr; + const bf16_t* residual = nullptr; + const float* post = nullptr; + const float* comb = nullptr; + const float* pre = nullptr; + bf16_t* normalized = nullptr; + fp8_e4m3_t* quantized = nullptr; + uint8_t* scales = nullptr; +}; + +template +SGL_DEVICE void mhc_quant_vec( + const FinalizeAllReduceParams& params, const StageVec& value, uint32_t token, uint32_t hvec) { + using namespace device; + fp32x2_t v[4]; + float amax = 0.0f; +#pragma unroll + for (uint32_t j = 0; j < 4; ++j) { + v[j] = cast(value[j]); + amax = fmaxf(amax, fmaxf(fabsf(v[j].x), fabsf(v[j].y))); + } + amax = fmaxf(amax, __shfl_xor_sync(0xffffffff, amax, 1, 4)); + amax = fmaxf(amax, __shfl_xor_sync(0xffffffff, amax, 2, 4)); + const float normalized = amax * (1.0f / 448.0f); + const uint32_t bits = __float_as_uint(normalized); + const uint32_t exponent = (bits >> 23) & 255; + const uint32_t mantissa = bits & 0x7fffff; + const bool bump = mantissa != 0 && !(exponent == 0 && mantissa <= 0x400000); + const uint32_t sf = normalized <= 0 ? 0 : min(exponent + uint32_t(bump), 254u); + const float inv_scale = __uint_as_float(sf == 0 ? 0 : (254 - sf) << 23); + AlignedVector q; +#pragma unroll + for (uint32_t j = 0; j < 4; ++j) { + q[j] = cast( + fp32x2_t{fminf(fmaxf(v[j].x * inv_scale, -448.0f), 448.0f), fminf(fmaxf(v[j].y * inv_scale, -448.0f), 448.0f)}); + } + q.store(params.quantized + static_cast(token) * kHiddenDim, hvec); + if (hvec % 4 == 0) { + const uint32_t g = hvec / 4; + const uint32_t off = (g / 4) * 512 + ((token % 32) * 4 + (token / 32) % 4) * 4 + g % 4; + params.scales[off] = sf; + } +} + +/// Apply HC=4 post mixing to an already BF16-rounded all-reduce vector. +/// Match mhc_post_split_h: round comb[0]*residual[0], then FMA post*x, +/// then FMA the remaining three residual streams in order. +template +SGL_DEVICE StageVec mhc_post_vec( + const FinalizeAllReduceParams& params, const StageVec& red, uint32_t token, uint32_t hvec) { + using namespace device; + StageVec residual[4]; + fp32x2_t collapsed[4] = {}; +#pragma unroll + for (uint32_t c = 0; c < 4; ++c) { + residual[c].load(params.residual + (static_cast(token) * 4 + c) * kHiddenDim, hvec); + } +#pragma unroll + for (uint32_t c = 0; c < 4; ++c) { + const float post = params.post[token * 4 + c]; + float comb[4]; +#pragma unroll + for (uint32_t r = 0; r < 4; ++r) + comb[r] = params.comb[token * 16 + r * 4 + c]; + StageVec out; +#pragma unroll + for (uint32_t j = 0; j < 4; ++j) { + const auto x = cast(red[j]); + const auto r0 = cast(residual[0][j]); + fp32x2_t acc{fmaf(post, x.x, __fmul_rn(comb[0], r0.x)), fmaf(post, x.y, __fmul_rn(comb[0], r0.y))}; +#pragma unroll + for (uint32_t r = 1; r < 4; ++r) { + const auto v = cast(residual[r][j]); + acc.x = fmaf(comb[r], v.x, acc.x); + acc.y = fmaf(comb[r], v.y, acc.y); + } + out[j] = cast(acc); + if constexpr (kCollapse) { + const auto rounded = cast(out[j]); + const float pre = params.pre[token * 4 + c]; + collapsed[j].x = fmaf(rounded.x, pre, collapsed[j].x); + collapsed[j].y = fmaf(rounded.y, pre, collapsed[j].y); + } + } + out.store(params.mhc_out + (static_cast(token) * 4 + c) * kHiddenDim, hvec); + } + StageVec result; + if constexpr (kCollapse) { +#pragma unroll + for (uint32_t j = 0; j < 4; ++j) + result[j] = cast(collapsed[j]); + } + return result; +} + +/// Row geometry: one 16B vector per thread, one cluster per row, so the block +/// size follows from the hidden width and the cluster size. The cluster size +/// is the tuning knob (dims per block = kHiddenDim / kClusterSize). +template +struct AllReduceNormTrait { + static constexpr uint32_t kRowVecs = kHiddenDim / 8; // 16B vectors per row + static constexpr uint32_t kBlockSize = kRowVecs / kClusterSize; // threads per block + static constexpr uint32_t kNumWarps = kBlockSize / device::kWarpThreads; + static_assert(kHiddenDim % 8 == 0, "hidden must be a whole number of 16B vectors"); + static_assert(1 <= kClusterSize && kClusterSize <= 8, "portable cluster sizes only"); + static_assert(kRowVecs % kClusterSize == 0, "cluster size must divide the row's vector count"); + static_assert(kBlockSize % device::kWarpThreads == 0, "block must be whole warps"); + static_assert(kBlockSize <= 1024, "block too large: raise the cluster size"); +}; + +// --- stage 1: the deferred finalize of one 16B vector ------------------------ +// The shared-expert vector is loaded first so that load is in flight while the +// routing rows and the kTopK gathers are fetched; the routed combine is then +// accumulated in ascending k from zero, rounded to bf16, and the shared vector +// is added with one more bf16 rounding (see the header: the unfused path's +// numerics, preserving the rank-local rounding points). +// Threads of the same token read the same kTopK indices / weights (a broadcast +// load per warp). FP32 routing weights retain their precision in the multiply; +// bf16 weights can use Blackwell's mixed-precision FMA. +template +SGL_DEVICE StageVec +finalize_vec(const FinalizeAllReduceParams& params, uint32_t token, uint32_t hvec) { + using namespace device; + const auto* idx = params.idx + static_cast(token) * kTopK; + const auto* weights = params.weights + static_cast(token) * kTopK; + int32_t rows[kTopK]; + WeightT w[kTopK]; +#pragma unroll + for (uint32_t k = 0; k < kTopK; ++k) { + rows[k] = idx[k]; + w[k] = weights[k]; + } + + // delay PDL wait until here + PDLWaitPrimary(); + + StageVec shared_in; + if constexpr (kHasShared) { + shared_in.load(params.shared + static_cast(token) * kHiddenDim, hvec); + } + + StageVec in[kTopK]; +#pragma unroll + for (uint32_t k = 0; k < kTopK; ++k) { + if (rows[k] >= 0) in[k].load(params.gemm2 + static_cast(rows[k]) * kHiddenDim, hvec); + } + + fp32x2_t acc[4]; +#pragma unroll + for (uint32_t j = 0; j < 4; ++j) { + acc[j] = fp32x2_t{0.0f, 0.0f}; + } + +#pragma unroll + for (uint32_t k = 0; k < kTopK; ++k) { + if (rows[k] < 0) continue; +#if SGL_ARCH_BLACKWELL_OR_GREATER + if constexpr (std::is_same_v) { +#pragma unroll + for (uint32_t j = 0; j < 4; ++j) { + acc[j].x = math::fma_f32_bf16(in[k][j].x, w[k], acc[j].x); + acc[j].y = math::fma_f32_bf16(in[k][j].y, w[k], acc[j].y); + } + } else +#endif + { + const auto w_fp32 = cast(w[k]); +#pragma unroll + for (uint32_t j = 0; j < 4; ++j) { + const auto [x, y] = cast(in[k][j]); + acc[j].x = fmaf(x, w_fp32, acc[j].x); + acc[j].y = fmaf(y, w_fp32, acc[j].y); + } + } + } + // Same rounding as the unfused path: TRT-LLM's finalize returns the routed + // combine rounded to bf16, and `shared.add_(routed)` then rounds the bf16 + + // bf16 sum once more (torch adds in fp32). Reproducing both roundings keeps + // the staged vector bit-identical to the unfused rank-local result. + StageVec out; +#pragma unroll + for (uint32_t j = 0; j < 4; ++j) { + if constexpr (kHasShared) { + const auto routed = cast(cast(acc[j])); + const auto sh = cast(shared_in[j]); + out[j] = cast(fp32x2_t{routed.x + sh.x, routed.y + sh.y}); + } else { + out[j] = cast(acc[j]); + } + } + return out; +} + +// --- the kernel -------------------------------------------------------------- +// Grid: dim3(num_tokens [+ 1], kClusterSize) with the cluster laid along y, so +// blockIdx.x is the token row (and its phase counter: PushEpoch's default) and +// blockIdx.y the rank inside the cluster. When rows do not own every counter +// of the plane, one extra cluster (blockIdx.x == num_tokens) is the bumper: it +// only flips the counters [num_tokens, num_push_counters) and exits; with +// num_tokens == num_push_counters no bumper is launched. The plane holds +// num_sm counters, so decode batches always fit. +template < + uint32_t kWorldSize, + uint32_t kHiddenDim, + uint32_t kTopK, + uint32_t kClusterSize, + bool kUsePDL, + bool kHasShared, + bool kNorm, + typename WeightT, + bool kMhc = false, + bool kQuant = false> +__global__ __launch_bounds__(AllReduceNormTrait::kBlockSize) + __cluster_dims__(1, kClusterSize, 1) void moe_finalize_all_reduce_kernel( + const __grid_constant__ FinalizeAllReduceParams params) { + namespace cg = cooperative_groups; + using namespace device; + using T = AllReduceNormTrait; + constexpr uint32_t kRowVecs = T::kRowVecs; + constexpr uint32_t kBlockSize = T::kBlockSize; + constexpr uint32_t kNumWarps = T::kNumWarps; + + const auto tx = threadIdx.x; + const auto row_idx = blockIdx.x; + const auto cluster_rank = blockIdx.y; + // this thread's vector within a row: cluster rank picks the block's chunk + const auto hvec = cluster_rank * kBlockSize + tx; + + // Under PDL this grid may start while the preceding kernel is still running. + // If that kernel is an all-reduce on this plane, it is still flipping the + // phase counters and resetting slot markers: reading the epoch now would see + // a half-done state and the poll below would never complete. Only a caller + // who knows the predecessor is a compute kernel (the MoE GEMM in the model) + // may defer the wait to finalize_vec, past the routing-metadata prefetch; + // the second wait there is then a no-op. + if (!params.prefetch_metadata) PDLWaitPrimary(); + + if (row_idx == params.num_tokens) { + PDLWaitPrimary(); + if constexpr (kQuant) { + // The SF buffer is padded to 128 rows. Active rows are written by the + // norm epilogue; the existing bumper zeros only the disjoint padding. + for (uint32_t off = hvec; off < (kHiddenDim / 32) * 128; off += kRowVecs) { + const uint32_t swizzled_row = (off % 512) / 4; + const uint32_t row = swizzled_row / 4 + (swizzled_row % 4) * 32; + if (row >= params.num_tokens) params.scales[off] = 0; + } + } + if (cluster_rank == 0) { + const auto epoch = distributed::PushEpoch{params.ws}; + __syncthreads(); + epoch.unsafe_flip_range(row_idx, params.num_push_counters); + } + return PDLTriggerSecondary(); + } + + // this cluster's epoch: the counter at blockIdx.x, one per row cluster + // (every block of the cluster reads the same one) + const auto epoch = distributed::PushEpoch{params.ws}; + const auto r = params.rank; + // my slot (`src = r`) inside every peer's workspace, and every peer's slot + // inside mine (`dst = r`), for this epoch + void* push_ptrs[kWorldSize]; +#pragma unroll + for (uint32_t i = 0; i < kWorldSize; ++i) { + push_ptrs[i] = epoch.slot_ptr(/*dst=*/i, /*src=*/r); + } + + // stage 1: finalize this row's vector in registers and push it to every peer + const auto vid = row_idx * kRowVecs + hvec; + { + auto vec = finalize_vec(params, row_idx, hvec); + Lamport::clear_pos_zero(vec.data()); +#pragma unroll + for (uint32_t i = 0; i < kWorldSize; ++i) { + ptx::st_relaxed_16B(vec, push_ptrs[i], vid); + } + } + + // ensure epoch is consumed, so flipping it won't lead to error + if constexpr (!kNorm) barrier_cluster_arrive_relaxed(); + + // stage 2: poll own slots, reduce across ranks, [norm], write, reset markers + void* poll_ptrs[kWorldSize]; +#pragma unroll + for (uint32_t i = 0; i < kWorldSize; ++i) { + poll_ptrs[i] = epoch.slot_ptr(/*dst=*/r, /*src=*/i); + } + StageVec vec[kWorldSize]; + do { + bool has_zero = false; +#pragma unroll + for (uint32_t i = 0; i < kWorldSize; ++i) { + ptx::ld_relaxed_16B(vec[i], poll_ptrs[i], vid); + } +#pragma unroll + for (uint32_t i = 0; i < kWorldSize; ++i) { + // the producer remapped +0.0 words, so a written word is never 0: + // word == 0 <=> the slot still holds the empty marker + has_zero |= Lamport::has_pos_zero(vec[i].data()); + } + if (!has_zero) break; + } while (true); + + if constexpr (!kNorm) { + const auto red = reduce_vec(vec); + ptx::st_global_16B(red, params.out, vid); + if constexpr (kMhc) mhc_post_vec(params, red, row_idx, hvec); + // ensure epoch is consumed, so flipping it won't lead to error + barrier_cluster_wait(); + } else { + // push to peer + __shared__ float smem_sq[kClusterSize][kNumWarps]; + auto red = reduce_vec(vec); + if constexpr (kMhc) { + ptx::st_global_16B(red, params.out, vid); + red = mhc_post_vec(params, red, row_idx, hvec); + } + StageVec w; + w.load(params.norm_weight, hvec); + const auto cluster = cg::this_cluster(); + fp32x2_t acc[4]; + float sq = 0.0f; +#pragma unroll + for (uint32_t j = 0; j < 4; ++j) { + acc[j] = cast(red[j]); + sq = fmaf(acc[j].x, acc[j].x, sq); + sq = fmaf(acc[j].y, acc[j].y, sq); + } + sq = warp::reduce_sum(sq); + const auto lane = tx % kWarpThreads; + const auto warp = tx / kWarpThreads; + if (lane < kClusterSize) { + *cluster.map_shared_rank(&smem_sq[cluster_rank][warp], lane) = sq; + } + cluster.sync(); + float total = 0.0f; +#pragma unroll + for (uint32_t c = 0; c < kClusterSize; ++c) { +#pragma unroll + for (uint32_t wp = 0; wp < kNumWarps; ++wp) { + total += smem_sq[c][wp]; + } + } + const auto factor = math::rsqrt(total / static_cast(kHiddenDim) + params.norm_eps); + StageVec out; +#pragma unroll + for (uint32_t j = 0; j < 4; ++j) { + const auto [wa, wb] = cast(w[j]); + out[j] = cast(fp32x2_t{acc[j].x * factor * wa, acc[j].y * factor * wb}); + } + ptx::st_global_16B(out, kMhc ? params.normalized : params.out, vid); + if constexpr (kQuant) mhc_quant_vec(params, out, row_idx, hvec); + } + PDLTriggerSecondary(); + + // re-establish the empty markers for the next same-phase round + StageVec zero_vec; + Lamport::fill_pos_zero(zero_vec.data()); +#pragma unroll + for (uint32_t i = 0; i < kWorldSize; ++i) { + ptx::st_global_16B(zero_vec, poll_ptrs[i], vid); + } + + if (cluster_rank == 0) epoch.flip(); +} + +// --- host -------------------------------------------------------------------- + +template < + uint32_t kWorldSize, + uint32_t kHiddenDim, + uint32_t kTopK, + uint32_t kClusterSize, + bool kUsePDL, + typename WeightT, + bool kMhc = false, + bool kQuant = false> +struct MoeFinalizeAllReduceKernel { + private: + static_assert(std::is_same_v || std::is_same_v); + using TensorView = tvm::ffi::TensorView; + using Params = FinalizeAllReduceParams; + using Trait = AllReduceNormTrait; + + template + static constexpr auto kernel = moe_finalize_all_reduce_kernel< + kWorldSize, + kHiddenDim, + kTopK, + kClusterSize, + kUsePDL, + kHasShared, + kNorm, + WeightT, + kMhc, + kQuant>; + + public: + /// out = [allreduce over ranks of] finalize(gemm2_out, idx, weights) [+ shared] [-> RMSNorm(norm_weight, eps)]. + /// `out` ([T, kHiddenDim] bf16) is output-only. `shared_output` and `norm_weight` + /// select the epilogue at runtime (four kernel instantiations per module). + static void + run(CommunicatorRef ref, + TensorView out, + TensorView gemm2_out, + TensorView permuted_idx, + TensorView expert_weights, + std::optional shared_output, + std::optional norm_weight, + double eps, + bool prefetch_metadata) { + static_assert(!kMhc); + run_impl( + ref, + out, + gemm2_out, + permuted_idx, + expert_weights, + shared_output, + norm_weight, + eps, + prefetch_metadata, + std::nullopt, + std::nullopt, + std::nullopt, + std::nullopt); + } + + /// Finalize + all-reduce + HC=4 post; original reduced output is retained. + static void run_mhc( + CommunicatorRef ref, + TensorView out, + TensorView gemm2_out, + TensorView permuted_idx, + TensorView expert_weights, + std::optional shared_output, + TensorView mhc_out, + TensorView residual, + TensorView post, + TensorView comb) { + static_assert(kMhc); + run_impl( + ref, + out, + gemm2_out, + permuted_idx, + expert_weights, + shared_output, + std::nullopt, + 0.0, + false, + mhc_out, + residual, + post, + comb); + } + + static void run_mhc_norm( + CommunicatorRef ref, + TensorView out, + TensorView gemm2_out, + TensorView permuted_idx, + TensorView expert_weights, + std::optional shared_output, + TensorView mhc_out, + TensorView residual, + TensorView post, + TensorView comb, + TensorView pre, + TensorView norm_weight, + double eps, + TensorView normalized) { + static_assert(kMhc); + run_impl( + ref, + out, + gemm2_out, + permuted_idx, + expert_weights, + shared_output, + norm_weight, + eps, + false, + mhc_out, + residual, + post, + comb, + pre, + normalized); + } + + static void run_mhc_quant( + CommunicatorRef ref, + TensorView out, + TensorView gemm2_out, + TensorView permuted_idx, + TensorView expert_weights, + std::optional shared_output, + TensorView mhc_out, + TensorView residual, + TensorView post, + TensorView comb, + TensorView pre, + TensorView norm_weight, + double eps, + TensorView normalized, + TensorView quantized, + TensorView scales) { + static_assert(kMhc && kQuant); + run_impl( + ref, + out, + gemm2_out, + permuted_idx, + expert_weights, + shared_output, + norm_weight, + eps, + false, + mhc_out, + residual, + post, + comb, + pre, + normalized, + quantized, + scales); + } + + private: + static void run_impl( + CommunicatorRef ref, + TensorView out, + TensorView gemm2_out, + TensorView permuted_idx, + TensorView expert_weights, + std::optional shared_output, + std::optional norm_weight, + double eps, + bool prefetch_metadata, + std::optional mhc_out, + std::optional residual, + std::optional post, + std::optional comb, + std::optional pre = std::nullopt, + std::optional normalized = std::nullopt, + std::optional quantized = std::nullopt, + std::optional scales = std::nullopt) { + using namespace host; + const auto& comm = *ref.get(); + const auto& push = comm.get_push_obj(); + CHECK_HOST(push.world_size == kWorldSize) + << "communicator holds " << push.world_size << " ranks, kernel built for " << kWorldSize; + + auto T = SymbolicSize{"num_tokens"}; + auto P = SymbolicSize{"num_permuted_rows"}; + auto TK = SymbolicSize{"num_expanded"}; + SymbolicDevice device; + device.set_options(); + TensorMatcher({T, kHiddenDim}) + .with_strides({kHiddenDim, 1}) + .with_dtype() + .with_device(device) + .verify(out); + TensorMatcher({P, kHiddenDim}) + .with_strides({kHiddenDim, 1}) + .with_dtype() + .with_device(device) + .verify(gemm2_out); + TensorMatcher({T, kTopK}) + .with_strides({kTopK, 1}) + .with_dtype() + .template with_device(device) + .verify(expert_weights); + TK.set_value(T.unwrap() * kTopK); + TensorMatcher({TK}).with_strides({1}).with_dtype().with_device(device).verify(permuted_idx); + if (shared_output.has_value()) { + TensorMatcher({T, kHiddenDim}) + .with_strides({kHiddenDim, 1}) + .with_dtype() + .with_device(device) + .verify(shared_output.value()); + } + if (norm_weight.has_value()) { + TensorMatcher({kHiddenDim}) + .with_strides({1}) + .with_dtype() + .with_device(device) + .verify(norm_weight.value()); + } + const auto num_tokens = static_cast(T.unwrap()); + if constexpr (kQuant) { + CHECK_HOST(num_tokens <= 8); + CHECK_HOST(norm_weight.has_value()); + TensorMatcher({T, kHiddenDim}).with_dtype().with_device(device).verify(quantized.value()); + TensorMatcher({(kHiddenDim / 32) * 128}) + .with_dtype() + .with_device(device) + .verify(scales.value()); + } + if constexpr (kMhc) { + static_assert(kHiddenDim == 5120); + if (norm_weight.has_value()) { + TensorMatcher({T, 4}).with_dtype().with_device(device).verify(pre.value()); + TensorMatcher({T, kHiddenDim}).with_dtype().with_device(device).verify(normalized.value()); + } + TensorMatcher({T, 4, kHiddenDim}) + .with_dtype() + .with_device(device) + .verify(mhc_out.value()) + .verify(residual.value()); + TensorMatcher({T, 4}).with_dtype().with_device(device).verify(post.value()); + TensorMatcher({T, 4, 4}).with_dtype().with_device(device).verify(comb.value()); + } + CHECK_HOST(num_tokens > 0) << "num_tokens must be positive"; + CHECK_HOST(reinterpret_cast(gemm2_out.data_ptr()) % 16 == 0) << "gemm2_out must be 16B aligned"; + CHECK_HOST(reinterpret_cast(out.data_ptr()) % 16 == 0) << "out must be 16B aligned"; + + // the whole [T, hidden] row view is staged by vector index, so it must fit + // one push slot; the generic push kernel's callers pick the slot size + const int64_t nbytes = num_tokens * int64_t(kHiddenDim) * sizeof(bf16_t); + CHECK_HOST(nbytes <= push.slot_bytes) << "num_tokens * hidden * 2 = " << nbytes << " bytes exceeds the " + << push.slot_bytes << "-byte push slot (reduce the batch or enlarge " + << "max_push_size)"; + // one cluster (and phase counter) per row; the bumper cluster is launched + // only when counters are left over for it to flip (num_tokens < num_blocks), + // so a batch that owns every counter runs without it + CHECK_HOST(num_tokens <= push.num_blocks) + << "num_tokens = " << num_tokens << " exceeds the " << push.num_blocks << " push phase counters of the plane"; + const uint32_t num_clusters = num_tokens + (num_tokens < push.num_blocks ? 1 : 0); + + const auto params = Params{ + .out = static_cast(out.data_ptr()), + .gemm2 = static_cast(gemm2_out.data_ptr()), + .idx = static_cast(permuted_idx.data_ptr()), + .weights = static_cast(expert_weights.data_ptr()), + .shared = shared_output.has_value() ? static_cast(shared_output.value().data_ptr()) : nullptr, + .norm_weight = norm_weight.has_value() ? static_cast(norm_weight.value().data_ptr()) : nullptr, + .norm_eps = static_cast(eps), + .prefetch_metadata = prefetch_metadata, + .rank = push.rank, + .num_tokens = num_tokens, + .num_push_counters = push.num_blocks, + .ws = push.get_workspace(nbytes), + .mhc_out = mhc_out.has_value() ? static_cast(mhc_out.value().data_ptr()) : nullptr, + .residual = residual.has_value() ? static_cast(residual.value().data_ptr()) : nullptr, + .post = post.has_value() ? static_cast(post.value().data_ptr()) : nullptr, + .comb = comb.has_value() ? static_cast(comb.value().data_ptr()) : nullptr, + .pre = pre.has_value() ? static_cast(pre.value().data_ptr()) : nullptr, + .normalized = normalized.has_value() ? static_cast(normalized.value().data_ptr()) : nullptr, + .quantized = quantized.has_value() ? static_cast(quantized.value().data_ptr()) : nullptr, + .scales = scales.has_value() ? static_cast(scales.value().data_ptr()) : nullptr, + }; + + const auto has_shared = shared_output.has_value(); + const auto has_norm = norm_weight.has_value(); + const auto kern = has_shared ? (has_norm ? kernel : kernel) + : (has_norm ? kernel : kernel); + // __cluster_dims__(1, kClusterSize, 1) is compiled in, so a plain launch + // already forms the clusters along y + LaunchKernel(dim3(num_clusters, kClusterSize), Trait::kBlockSize, out.device()).enable_pdl(kUsePDL)(kern, params); + } +}; + +} // namespace sglang diff --git a/python/sglang/kernels/jit/csrc/distributed/custom_all_reduce.cuh b/python/sglang/kernels/jit/csrc/distributed/custom_all_reduce.cuh index 627218dfdaf6..b048d2cb789a 100644 --- a/python/sglang/kernels/jit/csrc/distributed/custom_all_reduce.cuh +++ b/python/sglang/kernels/jit/csrc/distributed/custom_all_reduce.cuh @@ -142,7 +142,16 @@ ALL_REDUCE_KERNEL void all_reduce_1shot_push_kernel(const __grid_constant__ AllR const auto r = params.rank; const auto num_vecs = params.num_vecs; const auto num_threads = blockDim.x * gridDim.x; - const auto global_tid = blockIdx.x * blockDim.x + threadIdx.x; + // Round-robin warps to blocks rather than giving each block a contiguous run. + // The grid is pinned to the counter array, so once `num_vecs` stops filling + // `gridDim * blockDim` a block-major index leaves the tail CTAs with nothing + // to do; that happens over a whole 2x band of sizes, between the point where + // `choose_block_size` gives up on 512 and the point where 1024 threads fill + // the grid again. + const auto warp_in_block = threadIdx.x / kWarpThreads; + const auto lane_id = threadIdx.x % kWarpThreads; + const auto global_warp_id = blockIdx.x + gridDim.x * warp_in_block; + const auto global_tid = global_warp_id * kWarpThreads + lane_id; PDLWaitPrimary(); const auto epoch = distributed::PushEpoch{params.ws}; diff --git a/python/sglang/kernels/jit/csrc/distributed/nvlink_comm.cuh b/python/sglang/kernels/jit/csrc/distributed/nvlink_comm.cuh new file mode 100644 index 000000000000..1254f4fce1b7 --- /dev/null +++ b/python/sglang/kernels/jit/csrc/distributed/nvlink_comm.cuh @@ -0,0 +1,851 @@ +#include +#include +#include + +#include +#include +#include +#include + +#include + +#include +#include +#include + +#include +#include +#include +#include +#include + +namespace sglang { + +using device::distributed::PushWorkSpace; +using device::distributed::Semaphore; + +// Runtime uint32 division as one 32x32->64 multiply and a shift (the round-up +// magic number, exact for dividends below 2^31; a vector index is far smaller). +// Self-contained so the header builds with the CCCL bundled in every CUDA 13 +// toolkit: cuda::fast_mod_div only arrived in a later CCCL. +struct fast_mod_div_u32_t { + uint32_t divisor; + uint32_t magic; + uint32_t shift; + + __host__ explicit fast_mod_div_u32_t(uint32_t d) : divisor(d), magic(0), shift(0) { + if (d > 1) { + const uint32_t log2_ceil = 32 - std::countl_zero(d - 1); + const uint32_t p = 31 + log2_ceil; + magic = static_cast(((uint64_t{1} << p) + d - 1) / d); + shift = p - 32; + } + } + + __device__ friend uint32_t operator/(uint32_t n, const fast_mod_div_u32_t& fd) { + return fd.divisor == 1 ? n : __umulhi(n, fd.magic) >> fd.shift; + } + + __device__ friend uint32_t operator%(uint32_t n, const fast_mod_div_u32_t& fd) { + return n - (n / fd) * fd.divisor; + } +}; + +template +struct NVLinkCommPushParams { + const void* __restrict__ input; + const void* __restrict__ residual; + void* __restrict__ output; + uint32_t dst_offset; // AR = rank slot stride; AG = packed token prefix + uint32_t rank; + uint32_t num_push_vecs; + uint32_t num_poll_vecs; + uint32_t num_vecs_per_token; + // Ragged split, reduce-scatter only: rank r owns `avg + (r < rem)` tokens of + // the input starting at `r * avg + min(r, rem)`. + uint32_t tokens_avg; + uint32_t tokens_rem; + fast_mod_div_u32_t vecs_per_token_div; + PushWorkSpace ws; +}; + +struct NVLinkCommPullParams { + const void* __restrict__ input; + const void* __restrict__ residual; + void* __restrict__ output; + uint32_t num_vecs; + // multicast buffer + uint8_t* input_mc; + uint8_t* output_mc; + Semaphore* sem_local; + Semaphore* sem_mc; + uint32_t rank; + uint32_t world_size; +}; + +template +inline constexpr uint32_t get_poll_group(uint32_t world_size) { + if (world_size <= 8) return world_size; + return kHasResidual ? 6 : 8; +} + +inline constexpr uint32_t kPushCTASize = 1024; // max value +inline constexpr uint32_t kPullCTASize = 512; // max value + +#define PUSH_KERNEL __global__ __launch_bounds__(kPushCTASize, 1) +#define PULL_KERNEL __global__ __launch_bounds__(kPullCTASize, 1) + +enum Primitive { + RS = 0b01, // Reduce-Scatter + AG = 0b10, // All-Gather + AR = RS | AG, // All-Reduce = RS + AG +}; + +template +SGL_DEVICE vec_t reduce_vec(vec_t x, vec_t y) { + vec_t arr[2] = {x, y}; + return device::reduce_vec(arr); +} + +/** + * \brief Layout: + * 1. `AG`/`RS`: each rank push to its own slot + * [rank0] | [rank 1] | [rank 2] | ... + * 2. `AG`: each rank push to a contiguous region + * [rank0, rank1, rank2, ...] + * + * `RS` use swizzle layout for push kernel \n + * `AG` use normal linear layout for push kernel + */ +template +PUSH_KERNEL void nvlink_push_kernel(const __grid_constant__ NVLinkCommPushParams params) { + using namespace device; + enable_smem_spilling(); + constexpr uint32_t kVecSize = 16 / sizeof(T); // 16 bytes per vector + using vec_t = device::AlignedVector, kVecSize / 2>; + using Lamport = distributed::LamportTrait; + constexpr uint32_t kGroup = get_poll_group(kWorldSize); + + // Round-robin warps to blocks rather than giving each block a contiguous run. + // The poll domain is this rank's shard for the reduce-scatter, `world_size` + // times smaller than what the push loop walks, so a block-major index parks + // all of it on the first `num_poll_vecs / blockDim` CTAs and idles the rest + // of the SMs; with the grid pinned to the SM count that is most of them. + const auto warp_in_block = threadIdx.x / kWarpThreads; + const auto lane_id = threadIdx.x % kWarpThreads; + const auto global_warp_id = blockIdx.x + gridDim.x * warp_in_block; + const auto global_tid = global_warp_id * kWarpThreads + lane_id; + const auto num_threads = blockDim.x * gridDim.x; + + PDLWaitPrimary(); + const auto epoch = distributed::PushEpoch{params.ws}; + + void* push_ptrs[kWorldSize]; + /// NOTE: broadcast write is only fast when world size is large + if constexpr (kWorldSize < 8 && (kPrim & Primitive::AG)) { +#pragma unroll + for (uint32_t i = 0; i < kWorldSize; ++i) { + // Same address arithmetic as the multicast branch below, so the two agree + // on where a sender's shard lands. `dst_offset` is a slot stride for the + // all-reduce but a packed token prefix for the gather, whose consumer + // reads the plane linearly; `slot_ptr(i, rank)` would put the gather's + // senders `slot_bytes` apart and the poll loop would never see them. + push_ptrs[i] = static_cast(epoch.slot_ptr(/*dst=*/i)) + params.dst_offset; + } + } + + const auto dst_ptr_mc = params.ws.mc_workspace + params.dst_offset + epoch.slot_offset(); + const auto vpt = params.num_vecs_per_token; + + for (auto vid = global_tid; vid < params.num_push_vecs; vid += num_threads) { + if constexpr (kPrim & Primitive::AG) { + vec_t vec; + vec.load(params.input, vid); + if constexpr (kHasResidual && kPrim == Primitive::AG) { + vec_t res; + res.load(params.residual, vid); + vec = reduce_vec(vec, res); + } + Lamport::clear_pos_zero(vec.data()); + if constexpr (kWorldSize < 8) { +#pragma unroll + for (uint32_t i = 0; i < kWorldSize; ++i) { + ptx::st_relaxed_16B(vec, push_ptrs[i], vid); + } + } else { + ptx::st_multimem_16B(vec, dst_ptr_mc, vid); + } + } else /* reduce-scatter only */ { + const auto token_id = vid / params.vecs_per_token_div; + const auto offset = vid % params.vecs_per_token_div; + // Both by a compile-time constant, so this is a mask and a shift. + const auto dst_rank = token_id % kWorldSize; + const auto dst_token_id = token_id / kWorldSize; + // The walk is round-robin so neighbouring work lands on different peers + // and every link stays busy instead congestion on 1 rank + const auto avg_tokens = params.tokens_avg; + const auto rem_tokens = params.tokens_rem; + const auto rank_prefix = dst_rank * avg_tokens + std::min(dst_rank, rem_tokens); + const auto src_token = rank_prefix + dst_token_id; + vec_t vec; + vec.load(params.input, src_token * vpt + offset); + const auto dst_ptr = epoch.slot_ptr(dst_rank, params.rank); + Lamport::clear_pos_zero(vec.data()); + ptx::st_relaxed_16B(vec, dst_ptr, dst_token_id * vpt + offset); + } + } + + // Poll addresses are linear in the source rank -- one base, `slot_bytes` + // apart -- so a base plus a vector-index bias replaces a kWorldSize-wide + // pointer table: 2 registers instead of 2 per peer. (The push side cannot do + // this; `workspaces[i]` genuinely varies per peer.) + const auto poll_base = epoch.slot_ptr(params.rank); + const auto slot_vecs = params.ws.slot_bytes / sizeof(vec_t); + vec_t pos_zero_vec; + Lamport::fill_pos_zero(pos_zero_vec.data()); + PDLTriggerSecondary(); + + for (auto vid = global_tid; vid < params.num_poll_vecs; vid += num_threads) { + if constexpr (kPrim & Primitive::RS) { + constexpr uint32_t kNumPairs = kVecSize / 2; + vec_t out_vec; + + if constexpr (kGroup >= kWorldSize) { + vec_t vec[kWorldSize + kHasResidual]; + if constexpr (kHasResidual) vec[kWorldSize].load(params.residual, vid); + do { + bool has_zero = false; +#pragma unroll + for (uint32_t i = 0; i < kWorldSize; ++i) { + ptx::ld_relaxed_16B(vec[i], poll_base, i * slot_vecs + vid); + } +#pragma unroll + for (uint32_t i = 0; i < kWorldSize; ++i) { + has_zero |= Lamport::has_pos_zero(vec[i].data()); + } + if (!has_zero) break; + } while (true); + out_vec = reduce_vec(vec); +#pragma unroll + for (uint32_t i = 0; i < kWorldSize; ++i) { + ptx::st_global_16B(pos_zero_vec, poll_base, i * slot_vecs + vid); + } + } else /* > 1 group: divide into chunks */ { + fp32x2_t acc[kNumPairs]; + constexpr uint32_t kNumGroups = div_ceil(kWorldSize, kGroup); + vec_t vec[kGroup]; + vec_t res; + +#pragma unroll + for (uint32_t g = 0; g < kNumGroups; ++g) { + const auto for_each = [&](auto&& fn) { +#pragma unroll + for (uint32_t j = 0; j < kGroup; ++j) { + const auto i = g * kGroup + j; + if (i >= kWorldSize) continue; + fn(i, j); + } + }; + + // Loaded a group early so the fetch overlaps the last poll; it is + // folded into the accumulator once the groups are done. + if constexpr (kHasResidual) { + if (g + 1 == kNumGroups) res.load(params.residual, vid); + } + + do { + bool has_zero = false; + for_each([&](uint32_t i, uint32_t j) { + // load all the vectors + ptx::ld_relaxed_16B(vec[j], poll_base, i * slot_vecs + vid); + }); + for_each([&](uint32_t, uint32_t j) { + // check for zeros + has_zero |= Lamport::has_pos_zero(vec[j].data()); + }); + if (!has_zero) break; + } while (true); + + for_each([&](uint32_t i, uint32_t j) { +#pragma unroll + for (uint32_t k = 0; k < kNumPairs; ++k) { + const auto [x, y] = cast(vec[j][k]); + acc[k].x = i == 0 ? x : acc[k].x + x; + acc[k].y = i == 0 ? y : acc[k].y + y; + } + ptx::st_global_16B(pos_zero_vec, poll_base, i * slot_vecs + vid); + }); + } + if constexpr (kHasResidual) { +#pragma unroll + for (uint32_t k = 0; k < kNumPairs; ++k) { + const auto [x, y] = cast(res[k]); + acc[k].x += x; + acc[k].y += y; + } + } +#pragma unroll + for (uint32_t k = 0; k < kNumPairs; ++k) { + out_vec[k] = cast>(acc[k]); + } + } + + out_vec.store(params.output, vid); + } else /* all-gather only */ { + vec_t vec; + do { + ptx::ld_relaxed_16B(vec, poll_base, vid); + } while (Lamport::has_pos_zero(vec.data())); + vec.store(params.output, vid); + ptx::st_global_16B(pos_zero_vec, poll_base, vid); + } + } + + __syncthreads(); + epoch.flip(); +} + +template +PULL_KERNEL void nvlink_pull_kernel(const __grid_constant__ NVLinkCommPullParams params) { + using namespace device; + constexpr uint32_t kVecSize = 16 / sizeof(T); // 16 bytes per vector + using vec_t = device::AlignedVector, kVecSize / 2>; + constexpr uint32_t kNumWarpVecs = kPullUnroll * kWarpThreads; + + // Round-robin chunks to blocks rather than giving each block a contiguous + // run: the global warp index runs block-fastest, so neighbouring chunks are + // driven by different CTAs. + const auto warp_in_block = threadIdx.x / kWarpThreads; + const auto global_warp_id = blockIdx.x + gridDim.x * warp_in_block; + const auto lane_id = threadIdx.x % kWarpThreads; + const auto num_warps = gridDim.x * (kPullCTASize / kWarpThreads); + + PDLWaitPrimary(); + const auto barrier = distributed::McBarrier{params.sem_local, params.sem_mc, params.world_size, 2}; + barrier.arrive_relaxed(/*n=*/0); + __syncthreads(); + + const auto num_whole_chunks = params.num_vecs / kNumWarpVecs; + // warp uniform unrolled path, 0 predicate + for (auto chunk = global_warp_id; chunk < num_whole_chunks; chunk += num_warps) { + vec_t vecs[kPullUnroll]; + const auto base = chunk * kNumWarpVecs + lane_id; + +#pragma unroll + for (uint32_t i = 0; i < kPullUnroll; ++i) { + const auto vid = base + i * kWarpThreads; + if constexpr (kPrim & Primitive::RS) { + ptx::ld_multimem_16B(vecs[i], params.input_mc, vid); + } else { + ptx::ld_global_16B(vecs[i], params.input, vid); + } + } + + if constexpr (kHasResidual) { + vec_t residuals[kPullUnroll]; +#pragma unroll + for (uint32_t i = 0; i < kPullUnroll; ++i) { + residuals[i].load(params.residual, base + i * kWarpThreads); + } +#pragma unroll + for (uint32_t i = 0; i < kPullUnroll; ++i) { + vecs[i] = reduce_vec(vecs[i], residuals[i]); + } + } + +#pragma unroll + for (uint32_t i = 0; i < kPullUnroll; ++i) { + if constexpr (kPrim & Primitive::AG) { + ptx::st_multimem_16B(vecs[i], params.output_mc, base + i * kWarpThreads); + } else { + ptx::st_global_16B(vecs[i], params.output, base + i * kWarpThreads); + } + } + } + + const auto chunk_offset = num_whole_chunks * kNumWarpVecs; + const auto global_tid = global_warp_id * kWarpThreads + lane_id; + const auto global_threads = num_warps * kWarpThreads; + for (auto vid = chunk_offset + global_tid; vid < params.num_vecs; vid += global_threads) { + vec_t vec; + if constexpr (kPrim & Primitive::RS) { + ptx::ld_multimem_16B(vec, params.input_mc, vid); + } else { + ptx::ld_global_16B(vec, params.input, vid); + } + if constexpr (kHasResidual) { + vec_t res; + res.load(params.residual, vid); + vec = reduce_vec(vec, res); + } + if constexpr (kPrim & Primitive::AG) { + ptx::st_multimem_16B(vec, params.output_mc, vid); + } else { + ptx::st_global_16B(vec, params.output, vid); + } + } + + PDLTriggerSecondary(); + __syncthreads(); + if constexpr (kPrim & Primitive::AG) { + barrier.arrive_rel_acq(/*n=*/1); + } else { // no store multimem, only local store + barrier.arrive_relaxed(/*n=*/1); + } +} + +template +__global__ void nvlink_barrier_kernel(Semaphore* sem_local, Semaphore* sem_mc, uint32_t world_size) { + using device::distributed::McBarrier; + device::PDLWaitPrimary(); + const auto barrier = McBarrier{sem_local, sem_mc, world_size, 1}; + barrier.arrive_relaxed(0); + device::PDLTriggerSecondary(); +} + +/// Block size for the push kernel: the smallest that still spreads the work +/// over every SM, capped at the launch bound. +inline auto choose_push_block_size(uint32_t num_vecs) -> uint32_t { + static const uint32_t kNumSM = [] { + int device = 0; + CHECK_CUDA(cudaGetDevice(&device)); + return host::runtime::get_sm_count(device); + }(); + for (const uint32_t block_size : {128u, 256u, 384u, 512u}) { + if (host::div_ceil(num_vecs, block_size) <= kNumSM) return block_size; + } + return 1024u; +} + +template +struct NVLinkComm { + private: + using TensorView = tvm::ffi::TensorView; + using PushPlaneObj = host::distributed::PushPlaneObj; + using PullPlaneObj = host::distributed::PullPlaneObj; + using CommunicatorObj = host::distributed::CommunicatorObj; + using CommunicatorRef = host::distributed::CommunicatorRef; + static constexpr uint32_t kVecBytes = 16; + static constexpr uint32_t kVecSize = kVecBytes / sizeof(T); + + public: + struct RouteInfo { + uint32_t prefix_tokens; // exclusive prefix sum + uint32_t num_rank_tokens; // current rank + }; + + static RouteInfo get_routing(uint32_t num_tokens, uint32_t rank, uint32_t world_size) { + const auto avg = num_tokens / world_size; + const auto rem = num_tokens % world_size; + return {rank * avg + std::min(rank, rem), avg + (rank < rem ? 1 : 0)}; + } + + struct HostParams { + int64_t hidden_size; + DLDevice device; + }; + + /// \brief Base pointer of the residual, shifted onto this rank's slice when + /// the caller hands over the whole tensor. + /// + /// The kernels fold the residual in over their own working domain, which is + /// this rank's shard everywhere except the push all-reduce, where every rank + /// reduces the whole tensor. So a caller holding a shard-shaped residual + /// passes it straight through, and one holding the full tensor passes that + /// and gets sliced here -- which keeps ragged splits working, since the slice + /// comes from `get_routing` rather than a uniform stride. + static const void* get_residual_ptr( + const tvm::ffi::Optional& residual, + uint32_t domain_tokens, + uint32_t total_tokens, + uint32_t prefix_bytes) { + if (!residual.has_value()) return nullptr; + const auto tokens = static_cast(residual.value().size(0)); + const auto* base = static_cast(residual.value().data_ptr()); + if (tokens == domain_tokens) return base; + CHECK_HOST(tokens == total_tokens) << "residual has " << tokens << " tokens, expected " << domain_tokens + << " (this rank's shard) or " << total_tokens << " (the whole tensor)"; + return base + prefix_bytes; + } + + static HostParams check_params( + const TensorView in, + const TensorView out, + const tvm::ffi::Optional& residual = {}, + host::DebugInfo info = {}) { + using namespace host; + auto D = SymbolicSize{"hidden_size"}; + auto device_ = SymbolicDevice{}; + auto dtype_ = SymbolicDType{}; + if constexpr (!std::is_same_v) dtype_.set_options(); + device_.set_options(); + TensorMatcher({-1, D}) // + .with_dtype(dtype_) + .with_device(device_) + .verify(in, info); + TensorMatcher({-1, D}) // + .with_dtype(dtype_) + .with_device(device_) + .verify(out, info); + if (residual.has_value()) { + TensorMatcher({-1, D}) // + .with_dtype(dtype_) + .with_device(device_) + .verify(residual.value(), info); + } + return {D.unwrap(), device_.unwrap()}; + } + + private: + template + static void run_push( + const PushPlaneObj& push, + const TensorView in, + const TensorView out, + const tvm::ffi::Optional residual) { + CHECK_HOST(push.world_size == kWorldSize) << push.world_size << " != " << kWorldSize; + const auto [hidden_size, device] = check_params(in, out, residual); + const auto rank = push.rank; + const auto num_vecs_per_token = static_cast(hidden_size / kVecSize); + const auto num_push_vecs = static_cast(in.numel() / kVecSize); + const auto num_poll_vecs = static_cast(out.numel() / kVecSize); + const auto slot_bytes = static_cast(push.slot_bytes); + const auto out_nbytes = static_cast(out.numel() * sizeof(T)); + const auto num_tokens = static_cast(in.size(0)); + const auto out_tokens = static_cast(out.size(0)); + const auto total_tokens = kPrim == Primitive::AG ? out_tokens : num_tokens; + const auto routing = get_routing(total_tokens, rank, kWorldSize); + + uint32_t dst_offset = 0; + if constexpr (kPrim == Primitive::AR) { + CHECK_HOST(num_tokens == out_tokens); + CHECK_HOST(out_nbytes <= push.slot_bytes); + dst_offset = static_cast(rank * slot_bytes); + } else if constexpr (kPrim == Primitive::RS) { + CHECK_HOST(out_tokens == routing.num_rank_tokens); + CHECK_HOST(out_nbytes <= push.slot_bytes); + // dst_offset is not used for this case + } else { + static_assert(kPrim == Primitive::AG); + CHECK_HOST(num_tokens == routing.num_rank_tokens); + CHECK_HOST(out_nbytes <= slot_bytes * kWorldSize); + dst_offset = routing.prefix_tokens * static_cast(num_vecs_per_token * kVecBytes); + } + + // Slice the whole plane: the kernel reaches every slot, not just this rank's. + const auto block_size = choose_push_block_size(std::max(num_push_vecs, num_poll_vecs)); + CHECK_HOST(num_vecs_per_token > 0) << "fast div-mod rejects a zero divisor"; + const auto in_tokens_total = static_cast(in.size(0)); + // The all-reduce reduces the whole tensor on every rank; the other two work + // on this rank's shard, so a full-length residual is sliced. + const auto residual_domain = kPrim == Primitive::AR ? total_tokens : routing.num_rank_tokens; + const auto residual_ptr = get_residual_ptr( + residual, + residual_domain, + total_tokens, + routing.prefix_tokens * static_cast(num_vecs_per_token * kVecBytes)); + const auto params = NVLinkCommPushParams{ + .input = in.data_ptr(), + .residual = residual_ptr, + .output = out.data_ptr(), + .dst_offset = dst_offset, + .rank = rank, + .num_push_vecs = num_push_vecs, + .num_poll_vecs = num_poll_vecs, + .num_vecs_per_token = num_vecs_per_token, + .tokens_avg = in_tokens_total / kWorldSize, + .tokens_rem = in_tokens_total % kWorldSize, + .vecs_per_token_div = fast_mod_div_u32_t{num_vecs_per_token}, + .ws = push.get_workspace(/*size=*/0), + }; + const auto kernel = residual.has_value() ? nvlink_push_kernel + : nvlink_push_kernel; + host::LaunchKernel(push.num_blocks, block_size, device).enable_pdl(kUsePDL)(kernel, params); + } + + template + static void run_pull( + const PullPlaneObj& pull, + const TensorView in, + const TensorView out, + const tvm::ffi::Optional residual, + uintptr_t in_mc_ptr, + uintptr_t out_mc_ptr, + uint32_t num_blocks_hint) { + CHECK_HOST(pull.mc_semaphore != nullptr); + const auto [hidden_size, device] = check_params(in, out, residual); + const auto rank = pull.rank; + const auto world_size = pull.world_size; + const auto num_tokens = static_cast(in.size(0)); + const auto out_tokens = static_cast(out.size(0)); + const auto total_tokens = kPrim == Primitive::AG ? out_tokens : num_tokens; + const auto routing = get_routing(total_tokens, rank, world_size); + const auto num_vecs_per_token = static_cast(hidden_size / kVecSize); + const auto bytes_per_token = static_cast(num_vecs_per_token * kVecBytes); + const auto prefix_bytes = static_cast(routing.prefix_tokens * bytes_per_token); + + // 0 = no hint, autotune; > 0 always use hint but clip to upper bound + if constexpr (kPrim == Primitive::AR) { + CHECK_HOST(num_tokens == out_tokens && in_mc_ptr != 0 && out_mc_ptr != 0); + in_mc_ptr += prefix_bytes; + out_mc_ptr += prefix_bytes; + if (num_blocks_hint == 0) num_blocks_hint = host::div_ceil(256u, kPullUnroll * world_size); + } else if constexpr (kPrim == Primitive::RS) { + CHECK_HOST(out_tokens == routing.num_rank_tokens && in_mc_ptr != 0); + in_mc_ptr += prefix_bytes; + if (num_blocks_hint == 0) num_blocks_hint = pull.num_blocks; // use all the blocks for RS + } else { + static_assert(kPrim == Primitive::AG); + CHECK_HOST(num_tokens == routing.num_rank_tokens && out_mc_ptr != 0); + out_mc_ptr += prefix_bytes; + if (num_blocks_hint == 0) num_blocks_hint = host::div_ceil(128u, kPullUnroll * world_size); + } + /// NOTE: hard limit upper bound is `pull.num_blocks` + num_blocks_hint = std::min(num_blocks_hint, pull.num_blocks); + + // Every pull primitive works on this rank's shard, so a full-length + // residual is sliced onto it. + const auto residual_ptr = get_residual_ptr(residual, routing.num_rank_tokens, total_tokens, prefix_bytes); + const auto params = NVLinkCommPullParams{ + .input = in.data_ptr(), + .residual = residual_ptr, + .output = out.data_ptr(), + .num_vecs = static_cast(routing.num_rank_tokens * num_vecs_per_token), + .input_mc = std::bit_cast(in_mc_ptr), + .output_mc = std::bit_cast(out_mc_ptr), + .sem_local = pull.semaphores[rank], + .sem_mc = pull.mc_semaphore, + .rank = rank, + .world_size = pull.world_size, + }; + + /// NOTE: the final num_blocks resolution must be world unified, otherwise may deadlock + const auto max_vecs_in_world = host::div_ceil(total_tokens, world_size) * num_vecs_per_token; + const auto max_num_blocks = host::div_ceil(max_vecs_in_world, kPullUnroll * kPullCTASize); + const auto num_blocks = std::max(1u, std::min(max_num_blocks, num_blocks_hint)); + const auto kernel = residual.has_value() ? nvlink_pull_kernel + : nvlink_pull_kernel; + host::LaunchKernel(num_blocks, kPullCTASize, device).enable_pdl(kUsePDL)(kernel, params); + } + + public: + // specialized for each world size + template + static void + all_reduce_push(CommunicatorRef comm, TensorView in, TensorView out, tvm::ffi::Optional residual) { + return run_push(comm->get_push_obj(), in, out, residual); + } + template + static void + all_gather_push(CommunicatorRef comm, TensorView in, TensorView out, tvm::ffi::Optional residual) { + return run_push(comm->get_push_obj(), in, out, residual); + } + template + static void + reduce_scatter_push(CommunicatorRef comm, TensorView in, TensorView out, tvm::ffi::Optional residual) { + return run_push(comm->get_push_obj(), in, out, residual); + } + + // only compile once for each world size + template + static void all_reduce_pull( + CommunicatorRef comm, // only pull is needed + TensorView in, + TensorView out, + tvm::ffi::Optional residual, + int64_t in_mc_ptr, + int64_t out_mc_ptr, + uint32_t num_blocks_hint) { + return run_pull( + comm->get_pull_obj(), in, out, residual, in_mc_ptr, out_mc_ptr, num_blocks_hint); + } + template + static void all_gather_pull( + CommunicatorRef comm, // only pull is needed + TensorView in, + TensorView out, + tvm::ffi::Optional residual, + int64_t out_mc_ptr, + uint32_t num_blocks_hint) { + return run_pull( + comm->get_pull_obj(), in, out, residual, 0, out_mc_ptr, num_blocks_hint); + } + template + static void reduce_scatter_pull( + CommunicatorRef comm, // only pull is needed + TensorView in, + TensorView out, + tvm::ffi::Optional residual, + int64_t in_mc_ptr, + uint32_t num_blocks_hint) { + return run_pull(comm->get_pull_obj(), in, out, residual, in_mc_ptr, 0, num_blocks_hint); + } +}; + +/// The stream comes from the caller rather than the FFI environment: tvm-ffi +/// only publishes the framework stream when a call carries a DLPack tensor, and +/// this one carries none. Resolving it from the environment instead put the +/// launch on a stale stream, so under graph capture the barrier ran once at +/// capture time and every replay silently skipped it. +template +void nvlink_barrier(host::distributed::CommunicatorRef comm, int64_t stream_id) { + const auto& pull = comm->get_pull_obj(); + CHECK_HOST(pull.mc_semaphore); + const auto stream = std::bit_cast(stream_id); + const auto sem_local = pull.semaphores[pull.rank]; + host::LaunchKernel(1, device::kWarpThreads, stream) // + .enable_pdl(kUsePDL)(nvlink_barrier_kernel, sem_local, pull.mc_semaphore, pull.world_size); +} + +template +void all_gather_copy_engine( + const host::distributed::CommunicatorRef comm, + const tvm::ffi::TensorView in, + const tvm::ffi::TensorView out, + const int64_t out_mc_ptr) { + using Impl = NVLinkComm; + const auto& pull = comm->get_pull_obj(); + const auto [hidden_size, device] = Impl::check_params(in, out); + const auto total_tokens = out.size(0); + const auto routing = Impl::get_routing(total_tokens, pull.rank, pull.world_size); + CHECK_HOST(in.size(0) == routing.num_rank_tokens); + const auto element_bytes = host::dtype_bytes(in.dtype()); + const auto dst_ptr = out_mc_ptr + routing.prefix_tokens * hidden_size * element_bytes; + const auto stream = host::LaunchKernel::resolve_device(device); + const auto sem_local = pull.semaphores[pull.rank]; + const auto launch_barrier = [&] { + host::LaunchKernel(1, device::kWarpThreads, stream) // + .enable_pdl(kUsePDL)(nvlink_barrier_kernel, sem_local, pull.mc_semaphore, pull.world_size); + }; + + launch_barrier(); + CHECK_CUDA(cudaMemcpyAsync( + /*dst=*/std::bit_cast(dst_ptr), + /*src=*/in.data_ptr(), + /*count=*/in.numel() * element_bytes, + /*kind=*/cudaMemcpyDeviceToDevice, + /*stream=*/stream)); + launch_barrier(); +} + +/// Stream memory ops, resolved through the runtime so the module does not have +/// to link the driver library. +inline auto cu_stream_batch_mem_op() { + using Fn = CUresult (*)(CUstream, unsigned int, CUstreamBatchMemOpParams*, unsigned int); + static Fn fn = [] { + void* sym = nullptr; + cudaDriverEntryPointQueryResult found{}; + CHECK_CUDA(cudaGetDriverEntryPointByVersion("cuStreamBatchMemOp", &sym, 12030, cudaEnableDefault, &found)); + CHECK_HOST(found == cudaDriverEntryPointSuccess && sym != nullptr) + << "cuStreamBatchMemOp is unavailable; the copy-engine collectives need CUDA 12.3 or newer"; + return reinterpret_cast(sym); + }(); + return fn; +} + +/// Arrive-and-wait across the plane without launching anything. +/// +/// The arrive is a four-byte host-to-device copy to the flag array's multicast +/// alias, so one operation lands in every rank's array; the waits and the reset +/// writes go down as a single batched stream memory op, which the stream itself +/// blocks on. Nothing here occupies an SM. +/// +/// A sequence number would be baked into a graph at capture time and every +/// replay would then wait on a stale value, so the flag is a constant and the +/// same batch clears it again: each barrier walks its slots 0 -> 1 -> 0. That +/// makes graph and eager identical, at the cost of `world_size` extra writes. +/// +/// `slot` picks one of the two flag arrays. Consecutive barriers must alternate, +/// which is what keeps one round's arrive from being erased by the previous +/// round's reset: between two barriers on the same array there is always a +/// complete barrier on the other one. Callers therefore need an even number of +/// barriers per collective -- entry and exit. +inline void ce_barrier( + cudaStream_t stream, uint32_t* flag_local, uint32_t* flag_mc, uint32_t rank, uint32_t world_size, uint32_t slot) { + // A graph node keeps the source pointer, not the value, so this has to outlive + // the capture; pinned, because the copy engine stages a pageable source and + // that shows up as several microseconds on a four-byte transfer. + static const uint32_t* arrived = [] { + void* p = nullptr; + CHECK_CUDA(cudaHostAlloc(&p, sizeof(uint32_t), cudaHostAllocDefault)); + *static_cast(p) = 1; + return static_cast(p); + }(); + + const auto base = slot * world_size; + CHECK_CUDA(cudaMemcpyAsync(flag_mc + base + rank, arrived, sizeof(uint32_t), cudaMemcpyHostToDevice, stream)); + + std::vector ops; + ops.reserve(2 * world_size - 1); + for (uint32_t r = 0; r < world_size; ++r) { + if (r == rank) continue; + auto& op = ops.emplace_back(); + op.waitValue.operation = CU_STREAM_MEM_OP_WAIT_VALUE_32; + op.waitValue.address = std::bit_cast(flag_local + base + r); + op.waitValue.value = 1; + op.waitValue.flags = CU_STREAM_WAIT_VALUE_EQ; + } + // Clearing only touches this rank's copy, so it cannot erase an arrival a + // peer has yet to observe. Ordered after the waits within the batch. + for (uint32_t i = 0; i < world_size; ++i) { + auto& op = ops.emplace_back(); + op.writeValue.operation = CU_STREAM_MEM_OP_WRITE_VALUE_32; + op.writeValue.address = std::bit_cast(flag_local + base + i); + op.writeValue.value = 0; + op.writeValue.flags = CU_STREAM_WRITE_VALUE_DEFAULT; + } + const auto rc = + cu_stream_batch_mem_op()(std::bit_cast(stream), static_cast(ops.size()), ops.data(), 0); + CHECK_HOST(rc == CUDA_SUCCESS) << "cuStreamBatchMemOp failed with " << static_cast(rc); +} + +/// All-gather with no kernel at all: the copy engine writes this rank's shard +/// straight into every peer's output, and the two barriers are stream memory +/// ops. The walk starts at this rank so that at any step the senders are spread +/// across distinct destinations instead of converging on one. +/// +/// Unlike the multicast copy-engine gather, this injects `world_size` times the +/// payload but rides the unicast links, which is the better trade once the +/// multicast injection rate -- flat at roughly 100 GB/s regardless of fan-out -- +/// stops being amortised by a wide enough world. +inline void all_gather_copy_engine_unicast( + const host::distributed::CommunicatorRef comm, + const tvm::ffi::TensorView in, + const tvm::ffi::TensorView out, + const tvm::ffi::Array peer_out_ptrs, + const int64_t flag_ptr, + const int64_t flag_mc_ptr, + const int64_t stream_id) { + using Impl = NVLinkComm; + const auto& pull = comm->get_pull_obj(); + const auto [hidden_size, device] = Impl::check_params(in, out); + const auto world_size = pull.world_size; + const auto rank = pull.rank; + CHECK_HOST(static_cast(peer_out_ptrs.size()) == world_size) + << "need one output pointer per rank, got " << peer_out_ptrs.size(); + const auto routing = Impl::get_routing(static_cast(out.size(0)), rank, world_size); + CHECK_HOST(static_cast(in.size(0)) == routing.num_rank_tokens) + << "all_gather takes this rank's shard of " << out.size(0) << ", which is " << routing.num_rank_tokens + << " tokens, got " << in.size(0); + + const auto element_bytes = host::dtype_bytes(in.dtype()); + const auto shard_bytes = static_cast(in.numel()) * element_bytes; + const auto prefix_bytes = static_cast(routing.prefix_tokens) * hidden_size * element_bytes; + const auto stream = std::bit_cast(stream_id); + const auto flag_local = std::bit_cast(flag_ptr); + const auto flag_mc = std::bit_cast(flag_mc_ptr); + + ce_barrier(stream, flag_local, flag_mc, rank, world_size, /*slot=*/0); + for (uint32_t step = 0; step < world_size; ++step) { + const auto dst_rank = (rank + step) % world_size; + CHECK_CUDA(cudaMemcpyAsync( + /*dst=*/std::bit_cast(peer_out_ptrs[dst_rank] + prefix_bytes), + /*src=*/in.data_ptr(), + /*count=*/shard_bytes, + /*kind=*/cudaMemcpyDeviceToDevice, + /*stream=*/stream)); + } + ce_barrier(stream, flag_local, flag_mc, rank, world_size, /*slot=*/1); +} + +} // namespace sglang diff --git a/python/sglang/kernels/jit/include/sgl_kernel/distributed/ptx.cuh b/python/sglang/kernels/jit/include/sgl_kernel/distributed/ptx.cuh index 0b2983da9023..ae232abb502a 100644 --- a/python/sglang/kernels/jit/include/sgl_kernel/distributed/ptx.cuh +++ b/python/sglang/kernels/jit/include/sgl_kernel/distributed/ptx.cuh @@ -97,7 +97,7 @@ SGL_DEVICE void ld_multimem_16B(V& x, const void* mc_addr, int64_t vec_offset) { mc_addr = static_cast(mc_addr) + vec_offset * 16; if constexpr (std::is_same_v>) { float4 val; - asm volatile("multimem.ld_reduce.weak.add.v4.f32 {%0, %1, %2, %3}, [%4];" + asm volatile("multimem.ld_reduce.weak.global.add.v4.f32 {%0, %1, %2, %3}, [%4];" : "=f"(val.x), "=f"(val.y), "=f"(val.z), "=f"(val.w) : "l"(mc_addr)); x = *reinterpret_cast(&val); @@ -107,12 +107,12 @@ SGL_DEVICE void ld_multimem_16B(V& x, const void* mc_addr, int64_t vec_offset) { // rejects .f32 ("=f") destinations with "Arguments mismatch". uint4 val; if constexpr (std::is_same_v>) { - asm volatile("multimem.ld_reduce.weak.add.acc::f32.v4.f16x2 {%0, %1, %2, %3}, [%4];" + asm volatile("multimem.ld_reduce.weak.global.add.acc::f32.v4.f16x2 {%0, %1, %2, %3}, [%4];" : "=r"(val.x), "=r"(val.y), "=r"(val.z), "=r"(val.w) : "l"(mc_addr)); } else { static_assert(std::is_same_v>); // 4x bf16x2 - asm volatile("multimem.ld_reduce.weak.add.acc::f32.v4.bf16x2 {%0, %1, %2, %3}, [%4];" + asm volatile("multimem.ld_reduce.weak.global.add.acc::f32.v4.bf16x2 {%0, %1, %2, %3}, [%4];" : "=r"(val.x), "=r"(val.y), "=r"(val.z), "=r"(val.w) : "l"(mc_addr)); } @@ -150,7 +150,7 @@ SGL_DEVICE void st_multimem_16B(const V& x, void* mc_addr, int64_t vec_offset) { static_assert(alignof(V) == 16 && sizeof(V) == 16); const auto val = *reinterpret_cast(&x); mc_addr = static_cast(mc_addr) + vec_offset * 16; - asm volatile("multimem.st.weak.v4.f32 [%4], {%0, %1, %2, %3};" + asm volatile("multimem.st.weak.global.v4.f32 [%4], {%0, %1, %2, %3};" : : "f"(val.x), "f"(val.y), "f"(val.z), "f"(val.w), "l"(mc_addr)); #else diff --git a/python/sglang/kernels/ops/communication/all_reduce_fusion.py b/python/sglang/kernels/ops/communication/all_reduce_fusion.py new file mode 100644 index 000000000000..b2db0551168d --- /dev/null +++ b/python/sglang/kernels/ops/communication/all_reduce_fusion.py @@ -0,0 +1,256 @@ +"""Fused deferred-MoE finalize + 1shot push all-reduce [+ RMSNorm] (bf16). + +One entry point, :func:`moe_finalize_all_reduce`, over +``csrc/distributed/all_reduce_fusion.cuh``:: + + out[t] = allreduce( sum_k expert_weights[t, k] * gemm2_out[idx[t*top_k + k]] + (+ shared_output[t]) ) # then, optionally, + out[t] = out[t] * rsqrt(mean(out[t]^2) + eps) * norm_weight + +The rank-local finalize (the trtllm-gen ``do_finalize=False`` triple, see +``moe_runner/flashinfer_trtllm.py``) is computed in registers and pushed +straight into every peer's CustomAllReduceV2 push slot, so it never +materializes; ``idx == -1`` slots (EP: non-local expert) contribute nothing. +Small-batch only: the whole ``[T, hidden]`` bf16 row view must fit one push +slot (checked C++-side; :func:`fits_push_slot` lets callers pre-check). + +The un-normed result is what DeepSeek-V4.1 consumes (its consumer is the mHC +post-split, not an RMSNorm), so ``norm_weight=None`` is the primary +configuration; ``norm_weight`` + ``norm_eps`` give the K3-style fused norm. + +Geometry: one thread-block cluster per token row plus a bumper cluster that +keeps the plane's phase counters uniform; ``cluster_size`` blocks share a +row (``hidden / cluster_size`` dims each). :func:`default_cluster_size` holds +the tuned default per hidden size and can be overridden per call. + +Needs :func:`register_comm` once per process (the CustomAllReduceV2 +``Communicator``); the ops key on ``world_size`` alone, like the K3 ones. +""" + +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 sglang.srt.utils.custom_op import register_custom_op + +if TYPE_CHECKING: + from tvm_ffi.module import Module + + from sglang.kernels.ops.communication.all_reduce import Communicator + + +# Storage plane: the CustomAllReduceV2 Communicator (push plane only) + +_COMM_MAP: dict[int, Communicator] = {} + + +def register_comm(comm: Communicator) -> None: + """Register the CustomAllReduceV2 communicator whose push plane the fused + kernel stages through. + + ``world_size`` is the whole key (the custom op takes nothing else), so at + most one communicator per size may be registered in a process; a second + group of the same size would silently inherit the first one's peer + pointers and the symptom would be a hang, hence the assert. + """ + prev = _COMM_MAP.get(comm.world_size) + assert prev is None or prev is comm, ( + f"a different communicator is already registered for world_size=" + f"{comm.world_size}; these ops key only on world_size, so two groups of " + f"the same size cannot coexist in one process" + ) + _COMM_MAP[comm.world_size] = comm + + +def get_registered_comm(world_size: int) -> Optional[Communicator]: + return _COMM_MAP.get(world_size) + + +# Geometry + +_VEC_ELEMS = 8 # bf16 per 16B vector = per thread +_MAX_CLUSTER_SIZE = 8 # portable cluster size limit + + +def valid_cluster_sizes(hidden_dim: int) -> list[int]: + """Cluster sizes the kernel can be built for at this hidden width: whole + 16B vectors per row, whole warps per block, <= 1024 threads, <= 8 blocks.""" + if hidden_dim % _VEC_ELEMS != 0: + return [] + row_vecs = hidden_dim // _VEC_ELEMS + return [ + c + for c in range(1, _MAX_CLUSTER_SIZE + 1) + if row_vecs % c == 0 and (row_vecs // c) % 32 == 0 and row_vecs // c <= 1024 + ] + + +@cache_once +def default_cluster_size(hidden_dim: int) -> int: + if hidden_dim % 1024 == 0 and hidden_dim <= 8192: + return hidden_dim // 1024 + if hidden_dim % 512 == 0 and hidden_dim <= 3584: + return hidden_dim // 512 + candidates = valid_cluster_sizes(hidden_dim) + if not candidates: + raise ValueError( + f"hidden_dim={hidden_dim} has no valid cluster geometry (needs a " + f"multiple of {_VEC_ELEMS * 32} bf16)" + ) + # closest to 128 threads per block, larger block on ties + return min(candidates, key=lambda c: (abs(hidden_dim // _VEC_ELEMS // c - 128), c)) + + +def fits_push_slot(max_push_size: int, num_tokens: int, hidden_dim: int) -> bool: + """Whether a ``[num_tokens, hidden_dim]`` bf16 row view fits one push slot + (``CustomAllReduceV2.max_push_size``).""" + return 0 < num_tokens * hidden_dim * 2 <= max_push_size + + +# JIT module: one per (world_size, hidden_dim, top_k, cluster_size, weight_dtype); the +# shared-add and norm variants are compiled into it and picked at call time. + + +@cache_once +def _jit_module( + world_size: int, + hidden_dim: int, + top_k: int, + cluster_size: int, + weight_dtype: torch.dtype, +) -> Module: + assert cluster_size in valid_cluster_sizes(hidden_dim), ( + f"cluster_size={cluster_size} is not valid for hidden_dim={hidden_dim}; " + f"choose from {valid_cluster_sizes(hidden_dim)}" + ) + args = make_cpp_args( + world_size, hidden_dim, top_k, cluster_size, is_arch_support_pdl(), weight_dtype + ) + return load_jit( + "moe_finalize_all_reduce", + *args, + cuda_files=["distributed/all_reduce_fusion.cuh"], + cuda_wrappers=[("run", f"MoeFinalizeAllReduceKernel<{args}>::run")], + ) + + +def compile_moe_finalize_all_reduce( + world_size: int, + hidden_dim: int, + top_k: int, + cluster_size: Optional[int] = None, + weight_dtype: torch.dtype = torch.bfloat16, +) -> None: + """Warm the JIT module (tests / benches precompile in parallel).""" + _jit_module( + world_size, + hidden_dim, + top_k, + cluster_size or default_cluster_size(hidden_dim), + weight_dtype, + ) + + +@register_custom_op(mutates_args=["out"]) +def _moe_finalize_all_reduce_op( + world_size: int, + hidden_dim: int, + top_k: int, + cluster_size: int, + out: torch.Tensor, + gemm2_out: torch.Tensor, + expanded_idx_to_permuted_idx: torch.Tensor, + expert_weights: torch.Tensor, + shared_output: Optional[torch.Tensor], + norm_weight: Optional[torch.Tensor], + norm_eps: float, + prefetch_metadata: bool, +) -> None: + comm = _COMM_MAP.get(world_size) + assert comm is not None, ( + f"no communicator registered for world_size={world_size}; call " + "all_reduce_fusion.register_comm(comm.obj) first" + ) + _jit_module(world_size, hidden_dim, top_k, cluster_size, expert_weights.dtype).run( + comm, + out, + gemm2_out, + expanded_idx_to_permuted_idx, + expert_weights, + shared_output, + norm_weight, + norm_eps, + prefetch_metadata, + ) + + +def moe_finalize_all_reduce( + gemm2_out: torch.Tensor, + expanded_idx_to_permuted_idx: torch.Tensor, + expert_weights: torch.Tensor, + top_k: int, + shared_output: Optional[torch.Tensor] = None, + norm_weight: Optional[torch.Tensor] = None, + norm_eps: Optional[float] = None, + *, + world_size: int, + hidden_dim: int, + cluster_size: Optional[int] = None, + prefetch_metadata: bool = False, +) -> torch.Tensor: + """Deferred MoE finalize [+ shared add] -> 1shot push all-reduce [-> RMSNorm]. + + :param gemm2_out: ``[P, hidden_dim]`` bf16, trtllm-gen permuted / padded rows. + :param expanded_idx_to_permuted_idx: ``[T * top_k]`` int32, ``-1`` = dropped slot. + :param expert_weights: ``[T, top_k]`` bf16 or fp32; any routed scaling factor is + already folded in (nothing is rescaled here). + :param shared_output: optional ``[T, hidden_dim]`` bf16 added before the reduce. + :param norm_weight: optional ``[hidden_dim]`` bf16 RMSNorm weight; with + ``norm_eps`` it turns on the fused norm epilogue. + :param prefetch_metadata: let the kernel read the plane's phase counter and + the routing metadata before its PDL wait. Under + PDL the kernel may start as soon as the preceding + kernel *triggers*, and nothing earlier in the + stream is guaranteed complete until the wait: so + this is only valid when the preceding kernel is + not an all-reduce on the same plane AND the + producers of ``expanded_idx_to_permuted_idx`` / + ``expert_weights`` are known complete (a chain of + early-triggering kernels such as the TRT-LLM MoE + GEMMs is not). Defaults to False (wait first). + :returns: a new ``[T, hidden_dim]`` bf16 tensor (not in place). + """ + num_tokens = expert_weights.shape[0] + assert expert_weights.dtype in (torch.bfloat16, torch.float32) + assert expert_weights.shape[1] == top_k, (expert_weights.shape, top_k) + assert (norm_weight is None) == (norm_eps is None), ( + "norm_weight and norm_eps must be given together" + ) + out = torch.empty( + num_tokens, hidden_dim, dtype=torch.bfloat16, device=gemm2_out.device + ) + if num_tokens == 0: # nothing staged: no phase flip on any rank, stays in step + return out + _moe_finalize_all_reduce_op( + world_size, + hidden_dim, + top_k, + cluster_size or default_cluster_size(hidden_dim), + out, + gemm2_out, + expanded_idx_to_permuted_idx, + expert_weights, + shared_output, + norm_weight, + float(norm_eps) if norm_eps is not None else 0.0, + prefetch_metadata, + ) + return out diff --git a/python/sglang/kernels/ops/communication/all_reduce_mhc.py b/python/sglang/kernels/ops/communication/all_reduce_mhc.py new file mode 100644 index 000000000000..d17acb961bca --- /dev/null +++ b/python/sglang/kernels/ops/communication/all_reduce_mhc.py @@ -0,0 +1,318 @@ +"""MoE finalize and TP all-reduce with an HC=4 post-mixing epilogue.""" + +from typing import Optional + +import torch + +from sglang.kernels.jit.utils import ( + cache_once, + is_arch_support_pdl, + load_jit, + make_cpp_args, +) +from sglang.kernels.ops.communication.all_reduce_fusion import ( + default_cluster_size, + get_registered_comm, +) +from sglang.srt.utils.custom_op import register_custom_op + + +@cache_once +def _module(world_size, top_k, cluster_size, weight_dtype): + args = make_cpp_args( + world_size, 5120, top_k, cluster_size, is_arch_support_pdl(), weight_dtype, True + ) + return load_jit( + "moe_finalize_all_reduce_mhc", + *args, + cuda_files=["distributed/all_reduce_fusion.cuh"], + cuda_wrappers=[ + ("run", f"MoeFinalizeAllReduceKernel<{args}>::run_mhc"), + ("run_norm", f"MoeFinalizeAllReduceKernel<{args}>::run_mhc_norm"), + ], + ) + + +@register_custom_op(mutates_args=["out", "mhc_out"]) +def _run( + world_size: int, + top_k: int, + cluster_size: int, + out: torch.Tensor, + mhc_out: torch.Tensor, + gemm2: torch.Tensor, + idx: torch.Tensor, + weights: torch.Tensor, + shared: Optional[torch.Tensor], + residual: torch.Tensor, + post: torch.Tensor, + comb: torch.Tensor, +) -> None: + comm = get_registered_comm(world_size) + assert comm is not None + _module(world_size, top_k, cluster_size, weights.dtype).run( + comm, + out, + gemm2, + idx, + weights, + shared, + mhc_out, + residual, + post, + comb, + ) + + +def moe_finalize_all_reduce_mhc( + gemm2, + idx, + weights, + top_k, + shared, + residual, + post, + comb, + *, + world_size, + cluster_size=None, +): + out = torch.empty( + (weights.shape[0], 5120), dtype=torch.bfloat16, device=gemm2.device + ) + mhc_out = torch.empty_like(residual) + if weights.shape[0]: + _run( + world_size, + top_k, + cluster_size or default_cluster_size(5120), + out, + mhc_out, + gemm2, + idx, + weights, + shared, + residual, + post, + comb, + ) + return out, mhc_out + + +@register_custom_op(mutates_args=["out", "mhc_out", "normalized"]) +def _run_norm( + world_size: int, + top_k: int, + cluster_size: int, + out: torch.Tensor, + mhc_out: torch.Tensor, + normalized: torch.Tensor, + gemm2: torch.Tensor, + idx: torch.Tensor, + weights: torch.Tensor, + shared: Optional[torch.Tensor], + residual: torch.Tensor, + post: torch.Tensor, + comb: torch.Tensor, + pre: torch.Tensor, + norm_weight: torch.Tensor, + eps: float, +) -> None: + comm = get_registered_comm(world_size) + assert comm is not None + _module(world_size, top_k, cluster_size, weights.dtype).run_norm( + comm, + out, + gemm2, + idx, + weights, + shared, + mhc_out, + residual, + post, + comb, + pre, + norm_weight, + eps, + normalized, + ) + + +def moe_finalize_all_reduce_mhc_norm( + gemm2, + idx, + weights, + top_k, + shared, + residual, + post, + comb, + pre, + norm_weight, + eps, + *, + world_size, + cluster_size=None, +): + out = torch.empty( + (weights.shape[0], 5120), dtype=torch.bfloat16, device=gemm2.device + ) + mhc_out = torch.empty_like(residual) + normalized = torch.empty_like(out) + if weights.shape[0]: + _run_norm( + world_size, + top_k, + cluster_size or default_cluster_size(5120), + out, + mhc_out, + normalized, + gemm2, + idx, + weights, + shared, + residual, + post, + comb, + pre, + norm_weight, + eps, + ) + return out, mhc_out, normalized + + +@cache_once +def _identity_routing(rows, device): + return ( + torch.arange(rows, device=device, dtype=torch.int32), + torch.ones(rows, 1, device=device, dtype=torch.float32), + ) + + +def all_reduce_mhc_norm(x, residual, post, comb, pre, norm_weight, eps, *, world_size): + idx, weights = _identity_routing(x.shape[0], x.device) + return moe_finalize_all_reduce_mhc_norm( + x, + idx, + weights, + 1, + None, + residual, + post, + comb, + pre, + norm_weight, + eps, + world_size=world_size, + ) + + +@cache_once +def _quant_module(world_size, top_k, cluster_size, weight_dtype): + args = make_cpp_args( + world_size, + 5120, + top_k, + cluster_size, + is_arch_support_pdl(), + weight_dtype, + True, + True, + ) + return load_jit( + "moe_finalize_all_reduce_mhc_quant", + *args, + cuda_files=["distributed/all_reduce_fusion.cuh"], + cuda_wrappers=[("run", f"MoeFinalizeAllReduceKernel<{args}>::run_mhc_quant")], + ) + + +@register_custom_op( + mutates_args=["out", "mhc_out", "normalized", "quantized", "scales"] +) +def _run_quant( + world_size: int, + top_k: int, + cluster_size: int, + out: torch.Tensor, + mhc_out: torch.Tensor, + normalized: torch.Tensor, + quantized: torch.Tensor, + scales: torch.Tensor, + gemm2: torch.Tensor, + idx: torch.Tensor, + weights: torch.Tensor, + shared: Optional[torch.Tensor], + residual: torch.Tensor, + post: torch.Tensor, + comb: torch.Tensor, + pre: torch.Tensor, + norm_weight: torch.Tensor, + eps: float, +) -> None: + comm = get_registered_comm(world_size) + assert comm is not None + _quant_module(world_size, top_k, cluster_size, weights.dtype).run( + comm, + out, + gemm2, + idx, + weights, + shared, + mhc_out, + residual, + post, + comb, + pre, + norm_weight, + eps, + normalized, + quantized, + scales, + ) + + +def moe_finalize_all_reduce_mhc_quant( + gemm2, + idx, + weights, + top_k, + shared, + residual, + post, + comb, + pre, + norm_weight, + eps, + *, + world_size, + cluster_size=None, +): + rows = weights.shape[0] + assert 0 < rows <= 8 + out = torch.empty((rows, 5120), dtype=torch.bfloat16, device=gemm2.device) + mhc_out = torch.empty_like(residual) + normalized = torch.empty_like(out) + quantized = torch.empty_like(out, dtype=torch.float8_e4m3fn) + scales = torch.empty(160 * 128, device=gemm2.device, dtype=torch.uint8) + _run_quant( + world_size, + top_k, + cluster_size or default_cluster_size(5120), + out, + mhc_out, + normalized, + quantized, + scales, + gemm2, + idx, + weights, + shared, + residual, + post, + comb, + pre, + norm_weight, + eps, + ) + return out, mhc_out, normalized, quantized, scales diff --git a/python/sglang/kernels/ops/communication/nvlink_comm.py b/python/sglang/kernels/ops/communication/nvlink_comm.py new file mode 100644 index 000000000000..00efd9385c32 --- /dev/null +++ b/python/sglang/kernels/ops/communication/nvlink_comm.py @@ -0,0 +1,282 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, Final, List, NamedTuple + +import torch + +from sglang.kernels.jit.utils import ( + cache_once, + is_arch_support_pdl, + load_jit, + make_cpp_args, +) + +if TYPE_CHECKING: + from tvm_ffi import Module + + from sglang.kernels.ops.communication.all_reduce import Communicator + + +SUPPORTED_OPS: Final = ["all_reduce", "all_gather", "reduce_scatter"] + + +class Partition(NamedTuple): + num_prefix_tokens: int # excluive prefix sum of tokens + num_local_tokens: int # number of tokens in this rank + + +def get_token_partion(num_tokens: int, comm: Communicator) -> Partition: + rank = comm.rank + world_size = comm.world_size + avg_tokens = num_tokens // world_size + rem_tokens = num_tokens % world_size + return Partition( + num_prefix_tokens=avg_tokens * rank + min(rank, rem_tokens), + num_local_tokens=avg_tokens + (1 if rank < rem_tokens else 0), + ) + + +class CopyEngineFlags(NamedTuple): + flags: torch.Tensor # [2 * world_size,] on symm-mem + flags_ptr: int + flags_mc_ptr: int + + +def make_ce_flags( + group, + world_size: int, + *, + flags: torch.Tensor | None = None, +) -> CopyEngineFlags: + """Allocate the flag array the copy-engine barrier waits on.""" + from torch._C._distributed_c10d import _SymmetricMemory + + if flags is None: + flags = _SymmetricMemory.empty_strided_p2p( + (2 * world_size,), + [1], + torch.int32, + torch.device("cuda", torch.cuda.current_device()), + group.group_name, + ) + flags.zero_() + mc_ptr = get_multicast_ptr(flags) + torch.cuda.synchronize() + return CopyEngineFlags(flags, flags.data_ptr(), mc_ptr) + + +def get_multicast_ptr(tensor: torch.Tensor) -> int: + """Multicast alias of a symmetric-memory tensor. Collective on first call. + + torch caches the handle per allocation, so this stays cheap on repeat; a + cache of our own keyed by address would go stale the moment an allocation is + freed and the address reused. + """ + from torch._C._distributed_c10d import _SymmetricMemory + + ptr = _SymmetricMemory.rendezvous(tensor).multicast_ptr + assert ptr != 0, "tensor has no multicast alias; was it allocated p2p?" + return ptr + + +@cache_once +def _jit_misc_module() -> Module: + args = make_cpp_args(is_arch_support_pdl()) + return load_jit( + "nvl_comm_misc", + *args, + cuda_files=["distributed/nvlink_comm.cuh"], + cuda_wrappers=[ + ("barrier", f"nvlink_barrier<{args}>"), + ("all_gather_copy_engine", f"all_gather_copy_engine<{args}>"), + ("all_gather_copy_engine_unicast", "all_gather_copy_engine_unicast"), + ], + ) + + +@cache_once +def _jit_pull_module(dtype: torch.dtype, num_unroll: int) -> Module: + args = make_cpp_args(dtype, is_arch_support_pdl()) + return load_jit( + "nvl_comm_pull", + *args, + f"unroll{num_unroll}", + cuda_files=["distributed/nvlink_comm.cuh"], + cuda_wrappers=[ + (n, f"NVLinkComm<{args}>::{n}_pull<{num_unroll}>") for n in SUPPORTED_OPS + ], + ) + + +@cache_once +def _jit_push_module(dtype: torch.dtype, world_size: int) -> Module: + args = make_cpp_args(dtype, is_arch_support_pdl()) + return load_jit( + "nvl_comm_push", + *args, + f"world{world_size}", + cuda_files=["distributed/nvlink_comm.cuh"], + cuda_wrappers=[ + (n, f"NVLinkComm<{args}>::{n}_push<{world_size}>") for n in SUPPORTED_OPS + ], + ) + + +# `residual` on any of these is folded into the reduction rather than costing a +# separate pass. It may be shaped like this rank's shard or like the whole +# tensor; in the latter case this rank's slice is taken, so a ragged split needs +# no view on the caller's side. +def all_reduce_push( + comm: Communicator, + input: torch.Tensor, + output: torch.Tensor, + residual: torch.Tensor | None = None, +) -> None: + _jit_push_module(input.dtype, comm.world_size).all_reduce( + comm, input, output, residual + ) + + +def all_gather_push( + comm: Communicator, + input: torch.Tensor, + output: torch.Tensor, + residual: torch.Tensor | None = None, +) -> None: + _jit_push_module(input.dtype, comm.world_size).all_gather( + comm, input, output, residual + ) + + +def reduce_scatter_push( + comm: Communicator, + input: torch.Tensor, + output: torch.Tensor, + residual: torch.Tensor | None = None, +) -> None: + _jit_push_module(input.dtype, comm.world_size).reduce_scatter( + comm, input, output, residual + ) + + +def all_reduce_pull( + comm: Communicator, + input: torch.Tensor, + output: torch.Tensor, + residual: torch.Tensor | None = None, + *, + in_mc_ptr: int = 0, + out_mc_ptr: int = 0, + num_unroll=4, + num_blocks_hint: int = 0, +) -> None: + _jit_pull_module(input.dtype, num_unroll).all_reduce( + comm, + input, + output, + residual, + in_mc_ptr or get_multicast_ptr(input), + out_mc_ptr or get_multicast_ptr(output), + num_blocks_hint, + ) + + +def all_gather_pull( + comm: Communicator, + input: torch.Tensor, + output: torch.Tensor, + residual: torch.Tensor | None = None, + *, + out_mc_ptr: int = 0, + num_unroll=4, + num_blocks_hint: int = 0, +) -> None: + _jit_pull_module(input.dtype, num_unroll).all_gather( + comm, + input, + output, + residual, + out_mc_ptr or get_multicast_ptr(output), + num_blocks_hint, + ) + + +def reduce_scatter_pull( + comm: Communicator, + input: torch.Tensor, + output: torch.Tensor, + residual: torch.Tensor | None = None, + *, + in_mc_ptr: int = 0, + num_unroll=4, + num_blocks_hint: int = 0, +) -> None: + _jit_pull_module(input.dtype, num_unroll).reduce_scatter( + comm, + input, + output, + residual, + in_mc_ptr or get_multicast_ptr(input), + num_blocks_hint, + ) + + +def all_gather_copy_engine_multicast( + comm: Communicator, + input: torch.Tensor, + output: torch.Tensor, + *, + out_mc_ptr: int = 0, +) -> None: + _jit_misc_module().all_gather_copy_engine( + comm, + input, + output, + out_mc_ptr or get_multicast_ptr(output), + ) + + +def all_gather_copy_engine_unicast( + comm: Communicator, + input: torch.Tensor, + output: torch.Tensor, + *, + peer_out_ptrs: List[int] | None = None, + stream: int | None = None, + ce_flags: CopyEngineFlags, +) -> None: + """All-gather that launches no kernel at all. + + The copy engine moves the payload -- one peer-to-peer copy per rank, started + at this rank so the links are not all driven in the same order -- and the two + barriers around it are stream memory ops. `output` must be symmetric memory, + since this rank writes its shard straight into every peer's copy; `input` is + read locally and can be an ordinary tensor. `group` is only needed on the + first call for a given communicator, to allocate the barrier flags. + """ + from torch._C._distributed_c10d import _SymmetricMemory + + if peer_out_ptrs is None: + base = _SymmetricMemory.rendezvous(output).buffer_ptrs + # `output` may be a view into the middle of the allocation, and the peer + # pointers address its base; the same shift applies on every rank. + shift = output.data_ptr() - int(base[comm.rank]) + peer_out_ptrs = [int(p) + shift for p in base] + if stream is None: + stream = torch.cuda.current_stream().cuda_stream + _jit_misc_module().all_gather_copy_engine_unicast( + comm, input, output, peer_out_ptrs, *ce_flags[1:], stream + ) + + +def barrier(comm: Communicator, *, stream: int | None = None) -> None: + """Multicast barrier across the plane's ranks. + + The stream is passed explicitly because this call has no tensor argument, + and tvm-ffi only publishes the framework stream when one is present. Left + to the environment, the launch lands on a stale stream and a graph capture + silently drops it. + """ + if stream is None: + stream = torch.cuda.current_stream().cuda_stream + _jit_misc_module().barrier(comm, stream) diff --git a/test/registered/kernel/communication/test_moe_finalize_all_reduce.py b/test/registered/kernel/communication/test_moe_finalize_all_reduce.py new file mode 100644 index 000000000000..75cd9ffe49e4 --- /dev/null +++ b/test/registered/kernel/communication/test_moe_finalize_all_reduce.py @@ -0,0 +1,312 @@ +"""Fused deferred-MoE finalize + push all-reduce (``moe_finalize_all_reduce``) +against a torch reference, for bf16 and fp32 routing weights. + +The routed-MoE runners hand this kernel FlashInfer's ``do_finalize=False`` +triple. With unpacked ``(topk_ids, topk_weights)`` routing the weights arrive +in fp32, with packed routing in bf16; both must reproduce the unfused path's +numerics (fp32 accumulation, bf16 rounding at the routed combine, after the +``+ shared`` add and at the all-reduce output). +""" + +from __future__ import annotations + +import atexit +import logging +import os + +import pytest +import torch +import torch.distributed as dist + +import sglang.srt.distributed.parallel_state as ps +from sglang.kernels.jit.utils import cache_once, get_ci_test_range +from sglang.kernels.ops.communication import all_reduce_fusion +from sglang.kernels.ops.communication.mp import register_comm_cleanup +from sglang.srt.distributed.device_communicators.custom_all_reduce_v2 import ( + CustomAllReduceV2, +) +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kernels.utils import multigpu_pytest_main + +register_cuda_ci(est_time=180, stage="base-b-kernel-unit", runner_config="4-gpu-b200") + +HIDDEN = 5120 # DeepSeek-V4 hidden size, the width the fused path is used at +TOP_K = 6 +MB = 1024 * 1024 +NUM_TOKENS = get_ci_test_range([1, 2, 8, 64], [1, 8, 64]) +WEIGHT_DTYPES = [torch.bfloat16, torch.float32] + + +def _precompile(num_gpus): + for ws in num_gpus: + for dt in WEIGHT_DTYPES: + all_reduce_fusion.compile_moe_finalize_all_reduce( + ws, HIDDEN, TOP_K, weight_dtype=dt + ) + + +@cache_once +def _init_world(): + local_rank = int(os.environ["LOCAL_RANK"]) + world_size = int(os.environ["WORLD_SIZE"]) + torch.cuda.set_device(local_rank) + dist.init_process_group(backend="gloo") + ps._WORLD = coord = ps.init_world_group( + ranks=list(range(world_size)), + local_rank=local_rank, + backend="nccl", + ) + atexit.register(dist.destroy_process_group) + logging.disable(logging.INFO) + torch.cuda.set_stream(torch.cuda.Stream()) + return coord.cpu_group + + +@cache_once +def _init_nccl_group(): + _init_world() + local_rank = int(os.environ["LOCAL_RANK"]) + group = dist.new_group(backend="nccl", device_id=torch.device(f"cuda:{local_rank}")) + assert isinstance(group, dist.ProcessGroup) + return group + + +def _device() -> torch.device: + return torch.device(f"cuda:{int(os.environ['LOCAL_RANK'])}") + + +@cache_once +def _init_comm() -> CustomAllReduceV2: + cpu_group = _init_world() + comm = CustomAllReduceV2( + cpu_group, _device(), max_pull_size=1 * MB, max_push_size=2 * MB + ) + if comm.disabled: + raise RuntimeError("moe_finalize_all_reduce requires CustomAllReduceV2") + all_reduce_fusion.register_comm(comm.obj) + register_comm_cleanup(comm) + return comm + + +def _make_inputs(num_tokens: int, weight_dtype: torch.dtype, exact: bool, seed: int): + """Per-rank permuted GEMM2 rows, routing slots and shared-expert output. + + ``exact`` keeps every value a small dyadic number so the fp32 sums and the + bf16 roundings are lossless and the kernel can be checked bit-exactly. + """ + g = torch.Generator().manual_seed(seed * 7919 + dist.get_rank()) + num_slots = num_tokens * TOP_K + num_rows = num_slots + 8 # a few padded rows no slot points at + if exact: + gemm2 = torch.randint(-8, 9, (num_rows, HIDDEN), generator=g).to(torch.bfloat16) + weights = torch.randint(0, 8, (num_tokens, TOP_K), generator=g) / 8.0 + shared = torch.randint(-8, 9, (num_tokens, HIDDEN), generator=g).to( + torch.bfloat16 + ) + else: + gemm2 = torch.randn(num_rows, HIDDEN, generator=g).to(torch.bfloat16) + weights = torch.rand(num_tokens, TOP_K, generator=g) * 1.5 + shared = torch.randn(num_tokens, HIDDEN, generator=g).to(torch.bfloat16) + idx = torch.randperm(num_rows, generator=g)[:num_slots].to(torch.int32) + # EP: slots routed to an expert another rank owns carry -1 and contribute nothing. + idx[torch.rand(num_slots, generator=g) < 0.1] = -1 + dev = _device() + return gemm2.to(dev), idx.to(dev), weights.to(weight_dtype).to(dev), shared.to(dev) + + +def _local_ref(gemm2, idx, weights, shared): + """Unfused numerics: fp32 accumulate, bf16 at the combine and after + shared.""" + rows = idx.view(weights.shape).long() + valid = (rows >= 0).float() + gathered = gemm2[rows.clamp(min=0)].float() # [T, top_k, H] + w = weights.float() * valid + routed = (gathered * w.unsqueeze(-1)).sum(dim=1).to(torch.bfloat16) + if shared is None: + return routed + return (routed.float() + shared.float()).to(torch.bfloat16) + + +def _all_reduce_ref(local: torch.Tensor) -> torch.Tensor: + """fp32-accumulating bf16 all-reduce in rank order.""" + group = _init_nccl_group() + gathered = [torch.empty_like(local) for _ in range(dist.get_world_size(group))] + dist.all_gather(gathered, local, group=group) + acc = torch.zeros(local.shape, dtype=torch.float32, device=local.device) + for x in gathered: + acc += x.float() + return acc.to(torch.bfloat16) + + +def _fused(comm, gemm2, idx, weights, shared): + out = all_reduce_fusion.moe_finalize_all_reduce( + gemm2, + idx, + weights, + TOP_K, + shared, + world_size=comm.world_size, + hidden_dim=HIDDEN, + ) + torch.cuda.synchronize() + return out + + +@pytest.mark.parametrize("num_tokens", NUM_TOKENS) +@pytest.mark.parametrize("weight_dtype", WEIGHT_DTYPES, ids=["bf16", "fp32"]) +@pytest.mark.parametrize("use_shared", [False, True]) +@torch.inference_mode() +def test_moe_finalize_all_reduce_exact(num_tokens, weight_dtype, use_shared): + comm = _init_comm() + gemm2, idx, weights, shared = _make_inputs( + num_tokens, weight_dtype, exact=True, seed=num_tokens + ) + shared = shared if use_shared else None + ref = _all_reduce_ref(_local_ref(gemm2, idx, weights, shared)) + out = _fused(comm, gemm2, idx, weights, shared) + torch.testing.assert_close(out, ref, atol=0, rtol=0) + + +@pytest.mark.parametrize("num_tokens", NUM_TOKENS) +@pytest.mark.parametrize("use_shared", [False, True]) +@torch.inference_mode() +def test_moe_finalize_all_reduce_fp32_weights(num_tokens, use_shared): + """fp32 weights are consumed at fp32: the kernel tracks an fp32 reference + within bf16 output tolerance, and is not the bf16-rounded-weight result.""" + comm = _init_comm() + gemm2, idx, weights, shared = _make_inputs( + num_tokens, torch.float32, exact=False, seed=100 + num_tokens + ) + shared = shared if use_shared else None + ref = _all_reduce_ref(_local_ref(gemm2, idx, weights, shared)) + out = _fused(comm, gemm2, idx, weights, shared) + # The kernel's sequential fmaf and torch's sum round differently in fp32, + # so a rank-local combine (magnitude up to ~16) can flip one bf16 ulp + # before the cross-rank sum: allow that, nothing more. + torch.testing.assert_close(out, ref, atol=0.125, rtol=0.01) + out_bf16_weights = _fused(comm, gemm2, idx, weights.to(torch.bfloat16), shared) + assert not torch.equal(out, out_bf16_weights) + + +@pytest.mark.parametrize("num_tokens", [1, 5, 6, 8]) +@pytest.mark.parametrize("weight_dtype", WEIGHT_DTYPES, ids=["bf16", "fp32"]) +@pytest.mark.parametrize("epilogue", ["post", "norm", "quant"]) +@pytest.mark.parametrize("use_shared", [False, True]) +@pytest.mark.parametrize("seed", [0, 13]) +@torch.inference_mode() +def test_mhc_epilogue_graph(num_tokens, weight_dtype, epilogue, use_shared, seed): + from sglang.kernels.ops.communication import all_reduce_mhc + from sglang.kernels.ops.layernorm.hc_combine_norm import hc_combine_norm + from sglang.kernels.ops.layernorm.mhc_post_split_h import mhc_post_split_h + from sglang.srt.layers.quantization.fp8_utils import flashinfer_mxfp8_quantize + + comm = _init_comm() + gemm2, idx, weights, shared = _make_inputs( + num_tokens, weight_dtype, exact=False, seed=31 + ) + shared = shared if use_shared else None + top_k = TOP_K + # The attention epilogue uses the same reduction with one contribution. + if epilogue == "norm" and not use_shared: + top_k = 1 + gemm2 = gemm2[:num_tokens].contiguous() + idx = torch.arange(num_tokens, device=_device(), dtype=torch.int32) + weights = torch.ones(num_tokens, 1, device=_device(), dtype=weight_dtype) + residual = torch.randn( + num_tokens, 4, HIDDEN, device=_device(), dtype=torch.bfloat16 + ) + post = torch.randn(num_tokens, 4, device=_device()) + comb = torch.randn(num_tokens, 4, 4, device=_device()) + pre = torch.rand(num_tokens, 4, device=_device()) + nw = torch.randn(HIDDEN, device=_device(), dtype=torch.bfloat16) + torch.manual_seed(seed * 7919 + dist.get_rank()) + residual.normal_() + post.normal_() + comb.normal_() + pre.uniform_() + nw.normal_() + kernel = { + "post": all_reduce_mhc.moe_finalize_all_reduce_mhc, + "norm": all_reduce_mhc.moe_finalize_all_reduce_mhc_norm, + "quant": all_reduce_mhc.moe_finalize_all_reduce_mhc_quant, + }[epilogue] + + def chain(): + old = all_reduce_fusion.moe_finalize_all_reduce( + gemm2, + idx, + weights, + top_k, + shared, + world_size=comm.world_size, + hidden_dim=HIDDEN, + ) + ref_post = mhc_post_split_h(old, residual, post, comb) + ref_norm = hc_combine_norm(ref_post.flatten(1), pre, nw, 1e-6) + args = [gemm2, idx, weights, top_k, shared, residual, post, comb] + if epilogue != "post": + args.extend((pre, nw, 1e-6)) + outputs = kernel(*args, world_size=comm.world_size) + # Exercise the shared counters between generic push and row-cluster AR. + comm.custom_all_reduce(outputs[0]) + return old, ref_post, ref_norm, outputs + + chain() + torch.cuda.synchronize() + graph = torch.cuda.CUDAGraph() + with comm.capture(), torch.cuda.graph(graph): + old, ref_post, ref_norm, outputs = chain() + + def check_outputs(): + reduced, actual_post = outputs[:2] + assert torch.equal(old.view(torch.int16), reduced.view(torch.int16)) + assert torch.equal(ref_post.view(torch.int16), actual_post.view(torch.int16)) + if epilogue == "post": + return + actual_norm = outputs[2] + # The cluster and Triton RMS reductions have different addition orders. + torch.testing.assert_close(actual_norm, ref_norm, rtol=0.008, atol=0.0001) + # Different reduction trees can land on opposite sides of BF16 ties. + # Bound each value, rather than a data-dependent bit-identical fraction. + ulp = ( + actual_norm.view(torch.int16).int() - ref_norm.view(torch.int16).int() + ).abs() + ulp.masked_fill_(actual_norm == ref_norm, 0) # Treat signed zeros equally. + assert ulp.max() <= 1 + if epilogue == "quant": + q, sf = outputs[3:] + for backend in ("cuda", "cute-dsl"): + ref_q, ref_sf = flashinfer_mxfp8_quantize( + actual_norm, True, 32, backend + ) + assert torch.equal( + q.view(torch.uint8).flatten(), ref_q.view(torch.uint8).flatten() + ) + assert torch.equal(sf, ref_sf.flatten()) + + for replay, magnitude in enumerate((1e-3, 1.0, 1e3, 0.0, 1.0, 0.0, 1e-2)): + residual.normal_().mul_(magnitude) + gemm2.normal_().mul_(magnitude) + if epilogue == "quant": + outputs[4].fill_(0xAB) # Also require rewriting the padded scale rows. + if dist.get_rank() == replay % dist.get_world_size(): + torch.cuda._sleep(100000) + graph.replay() + torch.cuda.synchronize() + error = None + try: + check_outputs() + except AssertionError as exc: + error = f"rank={dist.get_rank()}, replay={replay}, scale={magnitude}: {exc}" + errors = [None] * dist.get_world_size() + # A failed assertion must stop all ranks before the next collective. + dist.all_gather_object(errors, error) + assert not any(errors), "\n".join(e for e in errors if e) + + +if __name__ == "__main__": + multigpu_pytest_main( + __name__, + __file__, + num_gpus=(4,), + pre_launch_fn=_precompile, + ) diff --git a/test/registered/kernel/communication/test_nvlink_comm.py b/test/registered/kernel/communication/test_nvlink_comm.py new file mode 100644 index 000000000000..b0ec65c3b067 --- /dev/null +++ b/test/registered/kernel/communication/test_nvlink_comm.py @@ -0,0 +1,202 @@ +"""Correctness of the NVLink collectives (``nvlink_comm``) against NCCL. + +All-reduce, all-gather and reduce-scatter on the push and pull planes of a +``CustomAllReduceV2`` communicator, with and without the folded residual, plus +the two copy-engine all-gathers, over token counts that cover the ragged split +(7), the remainder loop alone (1) and the bandwidth band (1024). The residual +is the same on every rank, as it is in a TP layer: the pull all-reduce folds +it in over each rank's token slice. + +Usage:: + + python test/registered/kernels/ops/communication/test_nvlink_comm.py --num-gpu 4 +""" + +from __future__ import annotations + +import atexit +import os +from typing import Dict, Tuple + +import pytest +import torch +import torch.distributed as dist + +import sglang.srt.distributed.parallel_state as ps +from sglang.kernels.jit.utils import cache_once +from sglang.kernels.ops.communication import nvlink_comm as nvl +from sglang.kernels.ops.communication.mp import register_comm_cleanup +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kernels.utils import multigpu_pytest_main + +register_cuda_ci(est_time=90, stage="base-b-kernel-unit", runner_config="4-gpu-b200") + +HIDDEN = 7168 +DTYPE = torch.bfloat16 +PUSH_SLOT_MB = 32 +PULL_MB = 4 + + +def _device() -> torch.device: + return torch.device("cuda", int(os.environ["LOCAL_RANK"])) + + +@cache_once +def _init_world(): + local_rank = int(os.environ["LOCAL_RANK"]) + world_size = int(os.environ["WORLD_SIZE"]) + torch.cuda.set_device(local_rank) + dist.init_process_group(backend="gloo") + ps._WORLD = coord = ps.init_world_group( + ranks=list(range(world_size)), local_rank=local_rank, backend="nccl" + ) + atexit.register(dist.destroy_process_group) + nccl_group = dist.new_group(backend="nccl", device_id=_device()) + return coord.cpu_group, nccl_group + + +@cache_once +def _init_comm(): + from sglang.srt.distributed.device_communicators.custom_all_reduce_v2 import ( + CustomAllReduceV2, + ) + + cpu_group, _ = _init_world() + comm = CustomAllReduceV2( + cpu_group, + _device(), + max_push_size=PUSH_SLOT_MB << 20, + max_pull_size=PULL_MB << 20, + ) + if comm.disabled: + pytest.skip("CustomAllReduceV2 is disabled on this system") + if not comm.has_multicast: + pytest.skip("the nvlink collectives need a multicast plane") + register_comm_cleanup(comm) + return comm + + +_SYMM: Dict[Tuple[int, int], torch.Tensor] = {} + + +def _symm(shape: Tuple[int, int]) -> torch.Tensor: + """Symmetric memory with a multicast alias; one allocation per shape, since + the allocation is collective and never returned.""" + from torch._C._distributed_c10d import _SymmetricMemory + + if shape not in _SYMM: + cpu_group, _ = _init_world() + t = _SymmetricMemory.empty_strided_p2p( + (shape[0] * shape[1],), [1], DTYPE, _device(), cpu_group.group_name + ) + _SymmetricMemory.rendezvous(t) + _SYMM[shape] = t.view(shape) + return _SYMM[shape] + + +def _shapes(op: str, tokens: int, world_size: int): + if op == "all_gather": + return (tokens, HIDDEN), (tokens * world_size, HIDDEN) + if op == "reduce_scatter": + return (tokens * world_size, HIDDEN), (tokens, HIDDEN) + return (tokens, HIDDEN), (tokens, HIDDEN) + + +def _reference(op, x, residual, nccl_group, world_size): + """fp32 NCCL reference with the kernels' residual placement: the gather adds + it to this rank's shard before gathering, the reductions to the output.""" + x = x.float() + res = residual.float() if residual is not None else 0 + if op == "all_reduce": + y = x.clone() + dist.all_reduce(y, group=nccl_group) + return y + res + if op == "all_gather": + x = (x + res).contiguous() + out = torch.empty( + (x.shape[0] * world_size, HIDDEN), dtype=torch.float32, device=x.device + ) + dist.all_gather_into_tensor(out, x, group=nccl_group) + return out + out = torch.empty( + (x.shape[0] // world_size, HIDDEN), dtype=torch.float32, device=x.device + ) + dist.reduce_scatter_tensor(out, x.contiguous(), group=nccl_group) + return out + res + + +_FNS = { + ("all_reduce", "push"): nvl.all_reduce_push, + ("all_gather", "push"): nvl.all_gather_push, + ("reduce_scatter", "push"): nvl.reduce_scatter_push, + ("all_reduce", "pull"): nvl.all_reduce_pull, + ("all_gather", "pull"): nvl.all_gather_pull, + ("reduce_scatter", "pull"): nvl.reduce_scatter_pull, +} + + +@pytest.mark.parametrize("tokens", [1, 7, 128, 1024]) +@pytest.mark.parametrize("residual", [False, True]) +@pytest.mark.parametrize("plane", ["push", "pull"]) +@pytest.mark.parametrize("op", nvl.SUPPORTED_OPS) +def test_collective(op: str, plane: str, residual: bool, tokens: int) -> None: + cpu_group, nccl_group = _init_world() + comm = _init_comm() + world_size = dist.get_world_size(cpu_group) + rank = dist.get_rank(cpu_group) + device = _device() + in_shape, out_shape = _shapes(op, tokens, world_size) + sym_in, sym_out = _symm(in_shape), _symm(out_shape) + gen = torch.Generator(device=device).manual_seed(1000 * tokens + rank) + sym_in.copy_(torch.randn(in_shape, dtype=DTYPE, device=device, generator=gen)) + sym_out.zero_() + res = None + if residual: + shared = torch.Generator(device=device).manual_seed(7 * tokens) + res_shape = in_shape if op == "all_gather" else out_shape + res = torch.randn(res_shape, dtype=DTYPE, device=device, generator=shared) + ref = _reference(op, sym_in, res, nccl_group, world_size) + dist.barrier(nccl_group) + torch.cuda.synchronize() + _FNS[(op, plane)](comm.obj, sym_in, sym_out, res) + torch.cuda.synchronize() + dist.barrier(nccl_group) + if op == "all_gather": + # a pure copy (plus one bf16 add with the residual): bit-exact + torch.testing.assert_close( + sym_out.float(), ref.to(DTYPE).float(), atol=0, rtol=0 + ) + else: + # bf16 sums in a different order than NCCL's fp32 tree + torch.testing.assert_close(sym_out.float(), ref, atol=0.1, rtol=0.02) + + +@pytest.mark.parametrize("tokens", [1, 7, 128]) +@pytest.mark.parametrize("variant", ["multicast", "unicast"]) +def test_copy_engine_all_gather(variant: str, tokens: int) -> None: + cpu_group, nccl_group = _init_world() + comm = _init_comm() + world_size = dist.get_world_size(cpu_group) + device = _device() + in_shape, out_shape = _shapes("all_gather", tokens, world_size) + sym_in, sym_out = _symm(in_shape), _symm(out_shape) + gen = torch.Generator(device=device).manual_seed( + 50 * tokens + dist.get_rank(cpu_group) + ) + sym_in.copy_(torch.randn(in_shape, dtype=DTYPE, device=device, generator=gen)) + sym_out.zero_() + ref = _reference("all_gather", sym_in, None, nccl_group, world_size) + dist.barrier(nccl_group) + torch.cuda.synchronize() + if variant == "multicast": + nvl.all_gather_copy_engine_multicast(comm.obj, sym_in, sym_out) + else: + flags = nvl.make_ce_flags(cpu_group, world_size) + nvl.all_gather_copy_engine_unicast(comm.obj, sym_in, sym_out, ce_flags=flags) + torch.cuda.synchronize() + dist.barrier(nccl_group) + torch.testing.assert_close(sym_out.float(), ref, atol=0, rtol=0) + + +if __name__ == "__main__": + multigpu_pytest_main(__name__, __file__, num_gpus=(4, 8)) From 5873fba21cfd57935cca3aa06980d44c84143bf4 Mon Sep 17 00:00:00 2001 From: hnyls2002 Date: Tue, 15 Sep 2026 15:53:40 -0700 Subject: [PATCH 05/30] dsv4.1: extract vocab gather and sharded greedy selection --- .../ops/speculative/dspark/sharded_greedy.py | 108 ++++++++ .../device_communicators/vocab_gather.py | 239 ++++++++++++++++++ .../kernel/communication/test_vocab_gather.py | 141 +++++++++++ .../speculative/test_dspark_sharded_greedy.py | 102 ++++++++ 4 files changed, 590 insertions(+) create mode 100644 python/sglang/kernels/ops/speculative/dspark/sharded_greedy.py create mode 100644 python/sglang/srt/distributed/device_communicators/vocab_gather.py create mode 100644 test/registered/kernel/communication/test_vocab_gather.py create mode 100644 test/registered/kernel/speculative/test_dspark_sharded_greedy.py diff --git a/python/sglang/kernels/ops/speculative/dspark/sharded_greedy.py b/python/sglang/kernels/ops/speculative/dspark/sharded_greedy.py new file mode 100644 index 000000000000..a5d63fd4eb8e --- /dev/null +++ b/python/sglang/kernels/ops/speculative/dspark/sharded_greedy.py @@ -0,0 +1,108 @@ +"""Exact greedy selection from TP-local base logits and rounded Markov bias.""" + +import torch +import triton +import triton.language as tl + + +@triton.jit +def _partial( + B, + X, + P, + BS: tl.constexpr, + XS: tl.constexpr, + WIDTH: tl.constexpr, + OFFSET: tl.constexpr, + PARTS: tl.constexpr, + BLOCK: tl.constexpr, +): + row, part = tl.program_id(0), tl.program_id(1) + i = part * BLOCK + tl.arange(0, BLOCK) + valid = i < WIDTH + bias = tl.load(B + row * BS + i, valid, 0).to(tl.float32) + base = tl.load(X + row * XS + i, valid, 0).to(tl.float32) + value = base + bias + nan = valid & (value != value) + has_nan = tl.sum(nan.to(tl.int32), 0) > 0 + maximum = tl.max(tl.where(valid & ~nan, value, -float("inf")), 0) + wins = tl.where(has_nan, nan, valid & (value == maximum)) + idx = tl.min(tl.where(wins, i + OFFSET, 2147483647), 0) + maximum = tl.where(has_nan, float("nan"), maximum) + tl.store(P + (row * PARTS + part) * 2, maximum) + tl.store(P + (row * PARTS + part) * 2 + 1, idx.to(tl.float32, bitcast=True)) + + +@triton.jit +def _finish( + P, + OUT, + BS: tl.constexpr, + PARTS: tl.constexpr, + WORLD: tl.constexpr, + BLOCK: tl.constexpr, +): + row = tl.program_id(0) + i = tl.arange(0, BLOCK) + valid = i < PARTS * WORLD + rank, part = i // PARTS, i % PARTS + offset = ((rank * BS + row) * PARTS + part) * 2 + value = tl.load(P + offset, valid, -float("inf")) + idx = tl.load(P + offset + 1, valid, 0).to(tl.int32, bitcast=True) + # The NVLink push transport changes +0 to -0 as its Lamport sentinel. + # Indices are nonnegative int32 bits, so strip that sign bit to recover ID 0. + idx &= 2147483647 + valid &= idx != 2147483647 + nan = valid & (value != value) + has_nan = tl.sum(nan.to(tl.int32), 0) > 0 + maximum = tl.max(tl.where(valid & ~nan, value, -float("inf")), 0) + wins = tl.where(has_nan, nan, valid & (value == maximum)) + result = tl.min(tl.where(wins, idx, 2147483647), 0) + tl.store(OUT + row, result.to(tl.int64)) + + +def sharded_greedy_step(bias, base_local, *, group, vocab_start, gather=None): + """Match argmax of rank-ordered BuildStepLocal/all_gather, excluding padding. + + ``bias`` is the original GEMM's already-rounded result. Communication carries + eight (value, global-index-bits) pairs per row; indices are transported as + bits and are never converted numerically to float. + """ + assert bias.ndim == base_local.ndim == 2 + assert bias.shape[0] == base_local.shape[0] + assert bias.shape[1] <= base_local.shape[1] + assert bias.stride(1) == base_local.stride(1) == 1 + rows, width = bias.shape + block = 4096 + parts = triton.cdiv(base_local.shape[1], block) + assert parts > 0 + partial = torch.empty((rows, parts, 2), device=bias.device, dtype=torch.float32) + _partial[(rows, parts)]( + bias, + base_local, + partial, + bias.stride(0), + base_local.stride(0), + width, + vocab_start, + parts, + block, + num_warps=4, + ) + # The padded partition width fixes the transport shape on all ranks; + # WIDTH masks real entries, including a completely empty final shard. + if gather is not None: + gathered = gather(partial.view(rows, parts * 2)) + else: + gathered = group.all_gather(partial, dim=0) if group.world_size > 1 else partial + result = torch.empty(rows, device=bias.device, dtype=torch.int64) + _finish[(rows,)]( + gathered, + result, + rows, + parts, + group.world_size, + triton.next_power_of_2(parts * group.world_size), + num_warps=4, + ) + return result diff --git a/python/sglang/srt/distributed/device_communicators/vocab_gather.py b/python/sglang/srt/distributed/device_communicators/vocab_gather.py new file mode 100644 index 000000000000..8a6ca10dcb81 --- /dev/null +++ b/python/sglang/srt/distributed/device_communicators/vocab_gather.py @@ -0,0 +1,239 @@ +"""All-gather of a vocab-parallel row block across a TP group. + +Every implementation takes this rank's ``[rows, local_width]`` slice and returns +``[rows, world_size * local_width]`` with the ranks' slices side by side, the +layout ``GroupCoordinator.all_gather(dim=-1)`` produces. Callers pick one with +``make_vocab_gather`` at init and call it unconditionally afterwards; which +transport runs is the implementation's business, including any fallback. +""" + +from __future__ import annotations + +import logging +from abc import ABC, abstractmethod +from typing import Optional, Tuple + +import torch + +logger = logging.getLogger(__name__) + + +class VocabGather(ABC): + """``[rows, local] -> [rows, world_size * local]``, ranks side by side.""" + + @abstractmethod + def __call__(self, local: torch.Tensor) -> torch.Tensor: ... + + @abstractmethod + def gather_stacked(self, local: torch.Tensor) -> torch.Tensor: + """Gather compact row blocks as [world_size * rows, local_width].""" + ... + + +class LocalVocabGather(VocabGather): + """A group of one: the slice is the whole row.""" + + def __call__(self, local: torch.Tensor) -> torch.Tensor: + return local + + def gather_stacked(self, local: torch.Tensor) -> torch.Tensor: + return local + + +class NcclVocabGather(VocabGather): + """The group coordinator's all_gather along the last dim (NCCL ring).""" + + def __init__(self, group) -> None: + self.group = group + + def __call__(self, local: torch.Tensor) -> torch.Tensor: + return self.group.all_gather(local, dim=-1) + + def gather_stacked(self, local: torch.Tensor) -> torch.Tensor: + return self.group.all_gather(local, dim=0) + + +def _alloc_symm( + group, shape: Tuple[int, int], dtype: torch.dtype +) -> Tuple[torch.Tensor, int]: + """A symmetric-memory tensor on ``group`` and its multicast alias (0 when + the group has none). Collective: every rank of the group must call it, in + the same order, outside CUDA-graph capture.""" + from torch._C._distributed_c10d import _SymmetricMemory + + # a GroupCoordinator names the allocation by its cpu_group, as + # CustomAllReduceV2 does; a torch process group names it itself + pg = getattr(group, "cpu_group", group) + buf = _SymmetricMemory.empty_strided_p2p( + (shape[0] * shape[1],), + [1], + dtype, + torch.device("cuda", torch.cuda.current_device()), + pg.group_name, + ) + mc_ptr = int(_SymmetricMemory.rendezvous(buf).multicast_ptr) + return buf.view(shape), mc_ptr + + +class NVLinkVocabGather(VocabGather): + """The NVLink collectives on CustomAllReduceV2's multicast plane. + + Both kernels gather along the row axis, so the ranks come back stacked and + are transposed into place. A slice that fits one slot of the push plane + takes the push kernel into a fresh tensor; a larger one that fits + ``pull_out`` takes the pull kernel into that symmetric-memory output, which + is reused every call, so the result is copied out of it; anything else goes + to ``fallback``, the NCCL ring. + + ``pull_out`` (``[world_size * symm_rows, local_width]``) is allocated here: + the allocation is collective and captured graphs keep its address, and with + CUDA graphs on the capture warm-up reaches it at the largest batch anyway. + """ + + def __init__( + self, + *, + ca_comm, + group, + local_width: int, + dtype: torch.dtype, + symm_rows: int, + fallback: VocabGather, + ) -> None: + self.comm = ca_comm.obj + self.world_size = int(group.world_size) + self.slot_bytes = int(ca_comm.max_push_size) + self.fallback = fallback + self.pull_out: Optional[torch.Tensor] = None + self.pull_mc_ptr = 0 + if symm_rows > 0 and self.comm.pull is not None: + out, mc_ptr = _alloc_symm( + group, (self.world_size * symm_rows, local_width), dtype + ) + if mc_ptr != 0: + self.pull_out, self.pull_mc_ptr = out, mc_ptr + logger.info( + "NVLink vocab gather: pull output %s (%d MB)", + tuple(out.shape), + out.numel() * out.element_size() >> 20, + ) + else: + logger.warning("NVLink vocab gather: no multicast alias, pull path off") + + def __call__(self, local: torch.Tensor) -> torch.Tensor: + rows = local.shape[0] + if local.nbytes <= self.slot_bytes: + return self._push(local) + total_rows = self.world_size * rows + if self.pull_out is not None and total_rows <= self.pull_out.shape[0]: + return self._pull(local, self.pull_out[:total_rows]) + return self.fallback(local) + + def gather_stacked(self, local: torch.Tensor) -> torch.Tensor: + # Compact argmax partials need rank-major output and no symmetric pull + # buffer. Unaligned rows and payloads past the push slot use NCCL. + if ( + local.is_contiguous() + and local.shape[1] * local.element_size() % 16 == 0 + and local.nbytes <= self.slot_bytes + ): + return self._push_stacked(local) + return self.fallback.gather_stacked(local) + + def _push(self, local: torch.Tensor) -> torch.Tensor: + return self._unstack(self._push_stacked(local)) + + def _push_stacked(self, local: torch.Tensor) -> torch.Tensor: + from sglang.kernels.ops.communication import nvlink_comm + + rows, width = local.shape + gathered = torch.empty( + (self.world_size * rows, width), dtype=local.dtype, device=local.device + ) + nvlink_comm.all_gather_push(self.comm, local, gathered) + return gathered + + def _pull(self, local: torch.Tensor, out: torch.Tensor) -> torch.Tensor: + from sglang.kernels.ops.communication import nvlink_comm + + nvlink_comm.all_gather_pull(self.comm, local, out, out_mc_ptr=self.pull_mc_ptr) + full = self._unstack(out) + # the transpose copies except at one row, where it would alias the + # shared buffer that the next call overwrites + return full.clone() if full.data_ptr() == out.data_ptr() else full + + def _unstack(self, gathered: torch.Tensor) -> torch.Tensor: + """``[world_size * rows, width]`` stacked by rank -> ``[rows, world_size * width]``.""" + rows = gathered.shape[0] // self.world_size + width = gathered.shape[1] + if rows == 1: + return gathered.view(1, self.world_size * width) + return ( + gathered.view(self.world_size, rows, width) + .transpose(0, 1) + .reshape(rows, self.world_size * width) + ) + + +def _nvlink_ca_comm(group, *, local_width: int, dtype: torch.dtype): + """The group's CustomAllReduceV2 when it can carry this gather on its + multicast plane, else None.""" + ca_comm = getattr(group, "ca_comm", None) + if ca_comm is None or getattr(ca_comm, "disabled", True): + return None + comm = getattr(ca_comm, "obj", None) + if comm is None or not getattr(ca_comm, "has_multicast", False): + return None + if comm.push is None or comm.world_size != group.world_size: + return None + # the kernels move 16-byte vectors along the row + if local_width % (128 // torch.finfo(dtype).bits) != 0: + return None + return ca_comm + + +def _default_symm_rows() -> int: + """One row per request in the largest batch the server runs (the + scheduler's max running requests, else the decode graph's max batch); 0 + when the server config is not published (offline use).""" + try: + from sglang.srt.runtime_context import get_exec, get_schedule + + return int( + get_schedule().max_running_requests + or get_exec().graph.cuda_graph_config.decode.max_bs + or 0 + ) + except Exception: + return 0 + + +def make_vocab_gather( + group, + *, + local_width: int, + dtype: torch.dtype = torch.float32, + prefer_nvlink: bool = True, + symm_rows: Optional[int] = None, +) -> VocabGather: + """The gather for ``group``: local for a group of one, NVLink when the + group's custom all-reduce has a multicast plane (and ``prefer_nvlink``), + the NCCL ring otherwise. ``symm_rows`` is the row capacity of the NVLink + gather's symmetric-memory output for slices past the push slot; None sizes + it for the server's largest batch, 0 leaves those slices to NCCL.""" + if group is None or group.world_size == 1: + return LocalVocabGather() + nccl = NcclVocabGather(group) + if not prefer_nvlink: + return nccl + ca_comm = _nvlink_ca_comm(group, local_width=local_width, dtype=dtype) + if ca_comm is None: + return nccl + return NVLinkVocabGather( + ca_comm=ca_comm, + group=group, + local_width=local_width, + dtype=dtype, + symm_rows=_default_symm_rows() if symm_rows is None else symm_rows, + fallback=nccl, + ) diff --git a/test/registered/kernel/communication/test_vocab_gather.py b/test/registered/kernel/communication/test_vocab_gather.py new file mode 100644 index 000000000000..7eb9055fd930 --- /dev/null +++ b/test/registered/kernel/communication/test_vocab_gather.py @@ -0,0 +1,141 @@ +"""``VocabGather``: the TP vocab-parallel row gather behind the DSpark draft head. + +On a TP group built the way the server builds it (custom all-reduce on, so +``ca_comm`` is a CustomAllReduceV2), ``make_vocab_gather`` must pick the NVLink +gather, and every path of that gather (push kernel, pull kernel into the +symmetric-memory output, NCCL past its capacity) must match the NCCL +``all_gather(dim=-1)`` reference, hand back a result that does not alias the +shared output, and replay correctly inside a CUDA graph. + +Usage:: + + python test/registered/kernels/ops/communication/test_vocab_gather.py --num-gpu 4 +""" + +from __future__ import annotations + +import atexit +import os + +import pytest +import torch +import torch.distributed as dist + +import sglang.srt.distributed.parallel_state as ps +from sglang.kernels.jit.utils import cache_once +from sglang.srt.distributed.device_communicators.vocab_gather import ( + NcclVocabGather, + NVLinkVocabGather, + make_vocab_gather, +) +from sglang.srt.distributed.parallel_state import GroupCoordinator +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kernels.utils import multigpu_pytest_main + +register_cuda_ci(est_time=60, stage="base-b-kernel-unit", runner_config="4-gpu-b200") + +LOCAL_WIDTH = 32320 # DeepSeek-V4.1's 129280-entry vocab over TP4 +SYMM_ROWS = 256 + + +def _device() -> torch.device: + return torch.device("cuda", int(os.environ["LOCAL_RANK"])) + + +@cache_once +def _tp_group() -> GroupCoordinator: + local_rank = int(os.environ["LOCAL_RANK"]) + world_size = int(os.environ["WORLD_SIZE"]) + torch.cuda.set_device(local_rank) + dist.init_process_group(backend="gloo") + ps._WORLD = ps.init_world_group( + ranks=list(range(world_size)), local_rank=local_rank, backend="nccl" + ) + atexit.register(dist.destroy_process_group) + return GroupCoordinator( + group_ranks=[list(range(world_size))], + local_rank=local_rank, + torch_distributed_backend="nccl", + use_pynccl=False, + use_pymscclpp=False, + use_custom_allreduce=True, + use_torch_symm_mem_all_reduce=False, + use_hpu_communicator=False, + use_xpu_communicator=False, + use_npu_communicator=False, + group_name="vocab_gather_test", + ) + + +@cache_once +def _gathers(): + tp = _tp_group() + nvlink = make_vocab_gather(tp, local_width=LOCAL_WIDTH, symm_rows=SYMM_ROWS) + if not isinstance(nvlink, NVLinkVocabGather): + pytest.skip("this TP group has no multicast plane") + nccl = make_vocab_gather(tp, local_width=LOCAL_WIDTH, prefer_nvlink=False) + return nvlink, nccl + + +def _rows(rows: int, seed: int) -> torch.Tensor: + gen = torch.Generator(device=_device()).manual_seed(seed + dist.get_rank()) + x = torch.randn(rows, LOCAL_WIDTH, device=_device(), generator=gen) + x[:, ::97] = 0.0 # exact zeros: the push kernel's Lamport sentinel path + return x + + +def _sync() -> None: + torch.cuda.synchronize() + dist.barrier() + + +def _check(nvlink: NVLinkVocabGather, nccl: NcclVocabGather, x: torch.Tensor) -> None: + ref = nccl(x) + got = nvlink(x) + _sync() + assert got.shape == ref.shape + assert bool((got == ref).all()) # -0.0 == 0.0, so the sentinel flip is invisible + if nvlink.pull_out is not None: + lo = nvlink.pull_out.data_ptr() + hi = lo + nvlink.pull_out.numel() * nvlink.pull_out.element_size() + assert not (lo <= got.data_ptr() < hi), "result aliases the shared pull output" + + +@pytest.mark.parametrize("rows", [1, 3, 6, 25, 64, 200, SYMM_ROWS + 1]) +def test_matches_nccl(rows: int) -> None: + nvlink, nccl = _gathers() + _check(nvlink, nccl, _rows(rows, seed=rows)) + + +def test_consecutive_pull_results_survive() -> None: + nvlink, nccl = _gathers() + a, b = _rows(64, seed=1), _rows(64, seed=2) + ra, rb = nvlink(a), nvlink(b) + _sync() + assert bool((ra == nccl(a)).all()) and bool((rb == nccl(b)).all()) + + +@pytest.mark.parametrize("rows", [1, 64]) +def test_cuda_graph_replay(rows: int) -> None: + nvlink, nccl = _gathers() + x = _rows(rows, seed=100 + rows) + nvlink(x) # eager once: JIT load + _sync() + graph = torch.cuda.CUDAGraph() + stream = torch.cuda.Stream() + with torch.cuda.stream(stream): + _sync() + with torch.cuda.graph(graph, stream=stream): + out = nvlink(x) + _sync() + for it in range(3): + x.copy_(_rows(rows, seed=200 + it)) + ref = nccl(x) + _sync() + graph.replay() + _sync() + assert bool((out == ref).all()) + + +if __name__ == "__main__": + multigpu_pytest_main(__name__, __file__, num_gpus=(4, 8)) diff --git a/test/registered/kernel/speculative/test_dspark_sharded_greedy.py b/test/registered/kernel/speculative/test_dspark_sharded_greedy.py new file mode 100644 index 000000000000..2310014bebb8 --- /dev/null +++ b/test/registered/kernel/speculative/test_dspark_sharded_greedy.py @@ -0,0 +1,102 @@ +"""Four-rank selection equivalence, including graph replay and padded shards.""" + +import atexit +import os + +import pytest +import torch +import torch.distributed as dist + +import sglang.srt.distributed.parallel_state as ps +from sglang.kernels.jit.utils import cache_once +from sglang.kernels.ops.speculative.dspark.sharded_greedy import sharded_greedy_step +from sglang.srt.distributed.device_communicators.vocab_gather import ( + NVLinkVocabGather, + make_vocab_gather, +) +from sglang.srt.distributed.parallel_state import GroupCoordinator +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kernels.utils import multigpu_pytest_main + +register_cuda_ci(est_time=120, stage="base-b-kernel-unit", runner_config="4-gpu-gb300") + + +@cache_once +def group(): + rank, world = int(os.environ["LOCAL_RANK"]), int(os.environ["WORLD_SIZE"]) + torch.cuda.set_device(rank) + dist.init_process_group(backend="gloo") + ps._WORLD = ps.init_world_group(list(range(world)), rank, backend="nccl") + atexit.register(dist.destroy_process_group) + torch.cuda.set_stream(torch.cuda.Stream()) + return GroupCoordinator( + group_ranks=[list(range(world))], + local_rank=rank, + torch_distributed_backend="nccl", + use_pynccl=False, + use_pymscclpp=False, + use_custom_allreduce=True, + use_torch_symm_mem_all_reduce=False, + use_hpu_communicator=False, + use_xpu_communicator=False, + use_npu_communicator=False, + group_name="sharded_greedy_test", + ) + + +@pytest.mark.parametrize("nvlink", [False, True]) +@pytest.mark.parametrize("m", [1, 4, 64]) +@pytest.mark.parametrize("width,last", [(32320, 32320), (8192, 17), (8, 0)]) +@pytest.mark.parametrize("case", ["random", "tie", "nan", "inf"]) +def test_sharded_selection_graph(m, width, last, case, nvlink): + g = group() + transport = make_vocab_gather( + g, local_width=width, prefer_nvlink=nvlink, symm_rows=0 + ) + if nvlink and not isinstance(transport, NVLinkVocabGather): + pytest.skip("this TP group has no multicast plane") + rank = g.rank_in_group + real = last if rank == g.world_size - 1 else width + # A slice of a block's logits has a non-contiguous row stride. + storage = torch.randn(m, 5, width, device="cuda") + base = storage[:, 2] + bias = torch.randn(m, real, device="cuda", dtype=torch.bfloat16) + + def candidate(): + return sharded_greedy_step( + bias, + base, + group=g, + vocab_start=rank * width, + gather=transport.gather_stacked, + ) + + for _ in range(3): + candidate() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + out = candidate() + for replay in range(4): + storage.normal_() + bias.normal_() + if case == "tie": + storage.zero_() + bias.zero_() + if case == "nan" and real: + base[:, replay % real] = float("nan") + if case == "inf": + storage.fill_(-float("inf")) + if replay % 2 and real: + base[:, replay % real] = float("inf") + graph.replay() + # Independent reference: gather complete, correctly padded FP32 logits. + local = torch.full((m, width), -float("inf"), device="cuda") + local[:, :real] = base[:, :real] + bias.float() + full = g.all_gather(local, dim=-1) + ref = full[:, : (g.world_size - 1) * width + last].argmax(-1) + torch.cuda.synchronize() + assert torch.equal(out, ref) + + +if __name__ == "__main__": + multigpu_pytest_main(__name__, __file__, num_gpus=(4,)) From 8c02375ad3d48d4838a07f5d79590ddef0c4c6f2 Mon Sep 17 00:00:00 2001 From: hnyls2002 Date: Tue, 15 Sep 2026 17:20:48 -0700 Subject: [PATCH 06/30] fix sharded greedy test suite --- .../registered/kernel/speculative/test_dspark_sharded_greedy.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/test/registered/kernel/speculative/test_dspark_sharded_greedy.py b/test/registered/kernel/speculative/test_dspark_sharded_greedy.py index 2310014bebb8..63ee7cd0ff4d 100644 --- a/test/registered/kernel/speculative/test_dspark_sharded_greedy.py +++ b/test/registered/kernel/speculative/test_dspark_sharded_greedy.py @@ -18,7 +18,7 @@ from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.kernels.utils import multigpu_pytest_main -register_cuda_ci(est_time=120, stage="base-b-kernel-unit", runner_config="4-gpu-gb300") +register_cuda_ci(est_time=120, stage="base-b-kernel-unit", runner_config="4-gpu-b200") @cache_once From b297ac27e625147df25fa0063a84b2ca548dde51 Mon Sep 17 00:00:00 2001 From: hnyls2002 Date: Tue, 15 Sep 2026 17:22:55 -0700 Subject: [PATCH 07/30] clarify sharded greedy docstring; split test cases; fix stale usage path --- .../ops/speculative/dspark/sharded_greedy.py | 7 +- .../kernel/communication/test_vocab_gather.py | 2 +- .../speculative/test_dspark_sharded_greedy.py | 70 +++++++++++++++---- 3 files changed, 63 insertions(+), 16 deletions(-) diff --git a/python/sglang/kernels/ops/speculative/dspark/sharded_greedy.py b/python/sglang/kernels/ops/speculative/dspark/sharded_greedy.py index a5d63fd4eb8e..d065f169a903 100644 --- a/python/sglang/kernels/ops/speculative/dspark/sharded_greedy.py +++ b/python/sglang/kernels/ops/speculative/dspark/sharded_greedy.py @@ -62,7 +62,12 @@ def _finish( def sharded_greedy_step(bias, base_local, *, group, vocab_start, gather=None): - """Match argmax of rank-ordered BuildStepLocal/all_gather, excluding padding. + """Fused BuildStepLocal + vocab gather + argmax, without materializing logits. + + Equivalent to argmax of rank-ordered ``build_step_local``/all_gather over the + sharded vocab, excluding padding, but each rank reduces its own shard first so + the transport carries a few partial (value, index) pairs per row instead of + the full local logits. ``bias`` is the original GEMM's already-rounded result. Communication carries eight (value, global-index-bits) pairs per row; indices are transported as diff --git a/test/registered/kernel/communication/test_vocab_gather.py b/test/registered/kernel/communication/test_vocab_gather.py index 7eb9055fd930..bb7ddaeaeb0d 100644 --- a/test/registered/kernel/communication/test_vocab_gather.py +++ b/test/registered/kernel/communication/test_vocab_gather.py @@ -9,7 +9,7 @@ Usage:: - python test/registered/kernels/ops/communication/test_vocab_gather.py --num-gpu 4 + python test/registered/kernel/communication/test_vocab_gather.py --num-gpu 4 """ from __future__ import annotations diff --git a/test/registered/kernel/speculative/test_dspark_sharded_greedy.py b/test/registered/kernel/speculative/test_dspark_sharded_greedy.py index 63ee7cd0ff4d..2ea62b48f378 100644 --- a/test/registered/kernel/speculative/test_dspark_sharded_greedy.py +++ b/test/registered/kernel/speculative/test_dspark_sharded_greedy.py @@ -44,11 +44,13 @@ def group(): ) -@pytest.mark.parametrize("nvlink", [False, True]) -@pytest.mark.parametrize("m", [1, 4, 64]) -@pytest.mark.parametrize("width,last", [(32320, 32320), (8192, 17), (8, 0)]) -@pytest.mark.parametrize("case", ["random", "tie", "nan", "inf"]) -def test_sharded_selection_graph(m, width, last, case, nvlink): +_NVLINK = pytest.mark.parametrize("nvlink", [False, True]) +_ROWS = pytest.mark.parametrize("m", [1, 4, 64]) +_SHARDS = pytest.mark.parametrize("width,last", [(32320, 32320), (8192, 17), (8, 0)]) + + +def _replay_and_check(m, width, last, nvlink, perturb): + """Capture the 4-rank selection graph, then replay it under `perturb`.""" g = group() transport = make_vocab_gather( g, local_width=width, prefer_nvlink=nvlink, symm_rows=0 @@ -79,15 +81,7 @@ def candidate(): for replay in range(4): storage.normal_() bias.normal_() - if case == "tie": - storage.zero_() - bias.zero_() - if case == "nan" and real: - base[:, replay % real] = float("nan") - if case == "inf": - storage.fill_(-float("inf")) - if replay % 2 and real: - base[:, replay % real] = float("inf") + perturb(replay, storage, base, bias, real) graph.replay() # Independent reference: gather complete, correctly padded FP32 logits. local = torch.full((m, width), -float("inf"), device="cuda") @@ -98,5 +92,53 @@ def candidate(): assert torch.equal(out, ref) +def _random(replay, storage, base, bias, real): + pass + + +def _tie(replay, storage, base, bias, real): + storage.zero_() + bias.zero_() + + +def _nan(replay, storage, base, bias, real): + if real: + base[:, replay % real] = float("nan") + + +def _inf(replay, storage, base, bias, real): + storage.fill_(-float("inf")) + if replay % 2 and real: + base[:, replay % real] = float("inf") + + +@_NVLINK +@_ROWS +@_SHARDS +def test_random_logits(m, width, last, nvlink): + _replay_and_check(m, width, last, nvlink, _random) + + +@_NVLINK +@_ROWS +@_SHARDS +def test_ties_resolve_to_the_lowest_index(m, width, last, nvlink): + _replay_and_check(m, width, last, nvlink, _tie) + + +@_NVLINK +@_ROWS +@_SHARDS +def test_nan_propagates(m, width, last, nvlink): + _replay_and_check(m, width, last, nvlink, _nan) + + +@_NVLINK +@_ROWS +@_SHARDS +def test_infinities(m, width, last, nvlink): + _replay_and_check(m, width, last, nvlink, _inf) + + if __name__ == "__main__": multigpu_pytest_main(__name__, __file__, num_gpus=(4,)) From d3d9a85d02e4449a70d816830f8600046e991141 Mon Sep 17 00:00:00 2001 From: BBuf <1182563586@qq.com> Date: Wed, 16 Sep 2026 12:21:40 +0800 Subject: [PATCH 08/30] fix communication guards and cover ragged collectives --- .../jit/csrc/distributed/nvlink_comm.cuh | 4 +- .../ops/communication/all_reduce_fusion.py | 10 + .../ops/communication/all_reduce_mhc.py | 3 + .../kernels/ops/communication/nvlink_comm.py | 8 +- .../test_moe_finalize_all_reduce.py | 312 ------------------ .../kernel/communication/test_nvlink_comm.py | 76 ++++- .../kernel/communication/test_vocab_gather.py | 141 -------- .../speculative/test_dspark_sharded_greedy.py | 144 -------- 8 files changed, 89 insertions(+), 609 deletions(-) delete mode 100644 test/registered/kernel/communication/test_moe_finalize_all_reduce.py delete mode 100644 test/registered/kernel/communication/test_vocab_gather.py delete mode 100644 test/registered/kernel/speculative/test_dspark_sharded_greedy.py diff --git a/python/sglang/kernels/jit/csrc/distributed/nvlink_comm.cuh b/python/sglang/kernels/jit/csrc/distributed/nvlink_comm.cuh index 1254f4fce1b7..36517ff2a71e 100644 --- a/python/sglang/kernels/jit/csrc/distributed/nvlink_comm.cuh +++ b/python/sglang/kernels/jit/csrc/distributed/nvlink_comm.cuh @@ -397,10 +397,10 @@ PULL_KERNEL void nvlink_pull_kernel(const __grid_constant__ NVLinkCommPullParams template __global__ void nvlink_barrier_kernel(Semaphore* sem_local, Semaphore* sem_mc, uint32_t world_size) { using device::distributed::McBarrier; - device::PDLWaitPrimary(); + device::PDLWaitPrimary(); const auto barrier = McBarrier{sem_local, sem_mc, world_size, 1}; barrier.arrive_relaxed(0); - device::PDLTriggerSecondary(); + device::PDLTriggerSecondary(); } /// Block size for the push kernel: the smallest that still spreads the work diff --git a/python/sglang/kernels/ops/communication/all_reduce_fusion.py b/python/sglang/kernels/ops/communication/all_reduce_fusion.py index b2db0551168d..ab5db2490b01 100644 --- a/python/sglang/kernels/ops/communication/all_reduce_fusion.py +++ b/python/sglang/kernels/ops/communication/all_reduce_fusion.py @@ -35,7 +35,9 @@ from sglang.kernels.jit.utils import ( cache_once, + get_jit_cuda_arch, is_arch_support_pdl, + is_hip_runtime, load_jit, make_cpp_args, ) @@ -119,6 +121,13 @@ def fits_push_slot(max_push_size: int, num_tokens: int, hidden_dim: int) -> bool # shared-add and norm variants are compiled into it and picked at call time. +def _require_cluster_launch_arch() -> None: + if is_hip_runtime() or get_jit_cuda_arch().major < 9: + raise RuntimeError( + "fused all-reduce cluster kernels require CUDA SM90 or newer" + ) + + @cache_once def _jit_module( world_size: int, @@ -127,6 +136,7 @@ def _jit_module( cluster_size: int, weight_dtype: torch.dtype, ) -> Module: + _require_cluster_launch_arch() assert cluster_size in valid_cluster_sizes(hidden_dim), ( f"cluster_size={cluster_size} is not valid for hidden_dim={hidden_dim}; " f"choose from {valid_cluster_sizes(hidden_dim)}" diff --git a/python/sglang/kernels/ops/communication/all_reduce_mhc.py b/python/sglang/kernels/ops/communication/all_reduce_mhc.py index d17acb961bca..8bf7da2947ca 100644 --- a/python/sglang/kernels/ops/communication/all_reduce_mhc.py +++ b/python/sglang/kernels/ops/communication/all_reduce_mhc.py @@ -11,6 +11,7 @@ make_cpp_args, ) from sglang.kernels.ops.communication.all_reduce_fusion import ( + _require_cluster_launch_arch, default_cluster_size, get_registered_comm, ) @@ -19,6 +20,7 @@ @cache_once def _module(world_size, top_k, cluster_size, weight_dtype): + _require_cluster_launch_arch() args = make_cpp_args( world_size, 5120, top_k, cluster_size, is_arch_support_pdl(), weight_dtype, True ) @@ -209,6 +211,7 @@ def all_reduce_mhc_norm(x, residual, post, comb, pre, norm_weight, eps, *, world @cache_once def _quant_module(world_size, top_k, cluster_size, weight_dtype): + _require_cluster_launch_arch() args = make_cpp_args( world_size, 5120, diff --git a/python/sglang/kernels/ops/communication/nvlink_comm.py b/python/sglang/kernels/ops/communication/nvlink_comm.py index 00efd9385c32..42f834746c4b 100644 --- a/python/sglang/kernels/ops/communication/nvlink_comm.py +++ b/python/sglang/kernels/ops/communication/nvlink_comm.py @@ -21,11 +21,11 @@ class Partition(NamedTuple): - num_prefix_tokens: int # excluive prefix sum of tokens + num_prefix_tokens: int # exclusive prefix sum of tokens num_local_tokens: int # number of tokens in this rank -def get_token_partion(num_tokens: int, comm: Communicator) -> Partition: +def get_token_partition(num_tokens: int, comm: Communicator) -> Partition: rank = comm.rank world_size = comm.world_size avg_tokens = num_tokens // world_size @@ -251,8 +251,8 @@ def all_gather_copy_engine_unicast( at this rank so the links are not all driven in the same order -- and the two barriers around it are stream memory ops. `output` must be symmetric memory, since this rank writes its shard straight into every peer's copy; `input` is - read locally and can be an ordinary tensor. `group` is only needed on the - first call for a given communicator, to allocate the barrier flags. + read locally and can be an ordinary tensor. Allocate ``ce_flags`` once per + communicator with :func:`make_ce_flags` and reuse it across calls. """ from torch._C._distributed_c10d import _SymmetricMemory diff --git a/test/registered/kernel/communication/test_moe_finalize_all_reduce.py b/test/registered/kernel/communication/test_moe_finalize_all_reduce.py deleted file mode 100644 index 75cd9ffe49e4..000000000000 --- a/test/registered/kernel/communication/test_moe_finalize_all_reduce.py +++ /dev/null @@ -1,312 +0,0 @@ -"""Fused deferred-MoE finalize + push all-reduce (``moe_finalize_all_reduce``) -against a torch reference, for bf16 and fp32 routing weights. - -The routed-MoE runners hand this kernel FlashInfer's ``do_finalize=False`` -triple. With unpacked ``(topk_ids, topk_weights)`` routing the weights arrive -in fp32, with packed routing in bf16; both must reproduce the unfused path's -numerics (fp32 accumulation, bf16 rounding at the routed combine, after the -``+ shared`` add and at the all-reduce output). -""" - -from __future__ import annotations - -import atexit -import logging -import os - -import pytest -import torch -import torch.distributed as dist - -import sglang.srt.distributed.parallel_state as ps -from sglang.kernels.jit.utils import cache_once, get_ci_test_range -from sglang.kernels.ops.communication import all_reduce_fusion -from sglang.kernels.ops.communication.mp import register_comm_cleanup -from sglang.srt.distributed.device_communicators.custom_all_reduce_v2 import ( - CustomAllReduceV2, -) -from sglang.test.ci.ci_register import register_cuda_ci -from sglang.test.kernels.utils import multigpu_pytest_main - -register_cuda_ci(est_time=180, stage="base-b-kernel-unit", runner_config="4-gpu-b200") - -HIDDEN = 5120 # DeepSeek-V4 hidden size, the width the fused path is used at -TOP_K = 6 -MB = 1024 * 1024 -NUM_TOKENS = get_ci_test_range([1, 2, 8, 64], [1, 8, 64]) -WEIGHT_DTYPES = [torch.bfloat16, torch.float32] - - -def _precompile(num_gpus): - for ws in num_gpus: - for dt in WEIGHT_DTYPES: - all_reduce_fusion.compile_moe_finalize_all_reduce( - ws, HIDDEN, TOP_K, weight_dtype=dt - ) - - -@cache_once -def _init_world(): - local_rank = int(os.environ["LOCAL_RANK"]) - world_size = int(os.environ["WORLD_SIZE"]) - torch.cuda.set_device(local_rank) - dist.init_process_group(backend="gloo") - ps._WORLD = coord = ps.init_world_group( - ranks=list(range(world_size)), - local_rank=local_rank, - backend="nccl", - ) - atexit.register(dist.destroy_process_group) - logging.disable(logging.INFO) - torch.cuda.set_stream(torch.cuda.Stream()) - return coord.cpu_group - - -@cache_once -def _init_nccl_group(): - _init_world() - local_rank = int(os.environ["LOCAL_RANK"]) - group = dist.new_group(backend="nccl", device_id=torch.device(f"cuda:{local_rank}")) - assert isinstance(group, dist.ProcessGroup) - return group - - -def _device() -> torch.device: - return torch.device(f"cuda:{int(os.environ['LOCAL_RANK'])}") - - -@cache_once -def _init_comm() -> CustomAllReduceV2: - cpu_group = _init_world() - comm = CustomAllReduceV2( - cpu_group, _device(), max_pull_size=1 * MB, max_push_size=2 * MB - ) - if comm.disabled: - raise RuntimeError("moe_finalize_all_reduce requires CustomAllReduceV2") - all_reduce_fusion.register_comm(comm.obj) - register_comm_cleanup(comm) - return comm - - -def _make_inputs(num_tokens: int, weight_dtype: torch.dtype, exact: bool, seed: int): - """Per-rank permuted GEMM2 rows, routing slots and shared-expert output. - - ``exact`` keeps every value a small dyadic number so the fp32 sums and the - bf16 roundings are lossless and the kernel can be checked bit-exactly. - """ - g = torch.Generator().manual_seed(seed * 7919 + dist.get_rank()) - num_slots = num_tokens * TOP_K - num_rows = num_slots + 8 # a few padded rows no slot points at - if exact: - gemm2 = torch.randint(-8, 9, (num_rows, HIDDEN), generator=g).to(torch.bfloat16) - weights = torch.randint(0, 8, (num_tokens, TOP_K), generator=g) / 8.0 - shared = torch.randint(-8, 9, (num_tokens, HIDDEN), generator=g).to( - torch.bfloat16 - ) - else: - gemm2 = torch.randn(num_rows, HIDDEN, generator=g).to(torch.bfloat16) - weights = torch.rand(num_tokens, TOP_K, generator=g) * 1.5 - shared = torch.randn(num_tokens, HIDDEN, generator=g).to(torch.bfloat16) - idx = torch.randperm(num_rows, generator=g)[:num_slots].to(torch.int32) - # EP: slots routed to an expert another rank owns carry -1 and contribute nothing. - idx[torch.rand(num_slots, generator=g) < 0.1] = -1 - dev = _device() - return gemm2.to(dev), idx.to(dev), weights.to(weight_dtype).to(dev), shared.to(dev) - - -def _local_ref(gemm2, idx, weights, shared): - """Unfused numerics: fp32 accumulate, bf16 at the combine and after + shared.""" - rows = idx.view(weights.shape).long() - valid = (rows >= 0).float() - gathered = gemm2[rows.clamp(min=0)].float() # [T, top_k, H] - w = weights.float() * valid - routed = (gathered * w.unsqueeze(-1)).sum(dim=1).to(torch.bfloat16) - if shared is None: - return routed - return (routed.float() + shared.float()).to(torch.bfloat16) - - -def _all_reduce_ref(local: torch.Tensor) -> torch.Tensor: - """fp32-accumulating bf16 all-reduce in rank order.""" - group = _init_nccl_group() - gathered = [torch.empty_like(local) for _ in range(dist.get_world_size(group))] - dist.all_gather(gathered, local, group=group) - acc = torch.zeros(local.shape, dtype=torch.float32, device=local.device) - for x in gathered: - acc += x.float() - return acc.to(torch.bfloat16) - - -def _fused(comm, gemm2, idx, weights, shared): - out = all_reduce_fusion.moe_finalize_all_reduce( - gemm2, - idx, - weights, - TOP_K, - shared, - world_size=comm.world_size, - hidden_dim=HIDDEN, - ) - torch.cuda.synchronize() - return out - - -@pytest.mark.parametrize("num_tokens", NUM_TOKENS) -@pytest.mark.parametrize("weight_dtype", WEIGHT_DTYPES, ids=["bf16", "fp32"]) -@pytest.mark.parametrize("use_shared", [False, True]) -@torch.inference_mode() -def test_moe_finalize_all_reduce_exact(num_tokens, weight_dtype, use_shared): - comm = _init_comm() - gemm2, idx, weights, shared = _make_inputs( - num_tokens, weight_dtype, exact=True, seed=num_tokens - ) - shared = shared if use_shared else None - ref = _all_reduce_ref(_local_ref(gemm2, idx, weights, shared)) - out = _fused(comm, gemm2, idx, weights, shared) - torch.testing.assert_close(out, ref, atol=0, rtol=0) - - -@pytest.mark.parametrize("num_tokens", NUM_TOKENS) -@pytest.mark.parametrize("use_shared", [False, True]) -@torch.inference_mode() -def test_moe_finalize_all_reduce_fp32_weights(num_tokens, use_shared): - """fp32 weights are consumed at fp32: the kernel tracks an fp32 reference - within bf16 output tolerance, and is not the bf16-rounded-weight result.""" - comm = _init_comm() - gemm2, idx, weights, shared = _make_inputs( - num_tokens, torch.float32, exact=False, seed=100 + num_tokens - ) - shared = shared if use_shared else None - ref = _all_reduce_ref(_local_ref(gemm2, idx, weights, shared)) - out = _fused(comm, gemm2, idx, weights, shared) - # The kernel's sequential fmaf and torch's sum round differently in fp32, - # so a rank-local combine (magnitude up to ~16) can flip one bf16 ulp - # before the cross-rank sum: allow that, nothing more. - torch.testing.assert_close(out, ref, atol=0.125, rtol=0.01) - out_bf16_weights = _fused(comm, gemm2, idx, weights.to(torch.bfloat16), shared) - assert not torch.equal(out, out_bf16_weights) - - -@pytest.mark.parametrize("num_tokens", [1, 5, 6, 8]) -@pytest.mark.parametrize("weight_dtype", WEIGHT_DTYPES, ids=["bf16", "fp32"]) -@pytest.mark.parametrize("epilogue", ["post", "norm", "quant"]) -@pytest.mark.parametrize("use_shared", [False, True]) -@pytest.mark.parametrize("seed", [0, 13]) -@torch.inference_mode() -def test_mhc_epilogue_graph(num_tokens, weight_dtype, epilogue, use_shared, seed): - from sglang.kernels.ops.communication import all_reduce_mhc - from sglang.kernels.ops.layernorm.hc_combine_norm import hc_combine_norm - from sglang.kernels.ops.layernorm.mhc_post_split_h import mhc_post_split_h - from sglang.srt.layers.quantization.fp8_utils import flashinfer_mxfp8_quantize - - comm = _init_comm() - gemm2, idx, weights, shared = _make_inputs( - num_tokens, weight_dtype, exact=False, seed=31 - ) - shared = shared if use_shared else None - top_k = TOP_K - # The attention epilogue uses the same reduction with one contribution. - if epilogue == "norm" and not use_shared: - top_k = 1 - gemm2 = gemm2[:num_tokens].contiguous() - idx = torch.arange(num_tokens, device=_device(), dtype=torch.int32) - weights = torch.ones(num_tokens, 1, device=_device(), dtype=weight_dtype) - residual = torch.randn( - num_tokens, 4, HIDDEN, device=_device(), dtype=torch.bfloat16 - ) - post = torch.randn(num_tokens, 4, device=_device()) - comb = torch.randn(num_tokens, 4, 4, device=_device()) - pre = torch.rand(num_tokens, 4, device=_device()) - nw = torch.randn(HIDDEN, device=_device(), dtype=torch.bfloat16) - torch.manual_seed(seed * 7919 + dist.get_rank()) - residual.normal_() - post.normal_() - comb.normal_() - pre.uniform_() - nw.normal_() - kernel = { - "post": all_reduce_mhc.moe_finalize_all_reduce_mhc, - "norm": all_reduce_mhc.moe_finalize_all_reduce_mhc_norm, - "quant": all_reduce_mhc.moe_finalize_all_reduce_mhc_quant, - }[epilogue] - - def chain(): - old = all_reduce_fusion.moe_finalize_all_reduce( - gemm2, - idx, - weights, - top_k, - shared, - world_size=comm.world_size, - hidden_dim=HIDDEN, - ) - ref_post = mhc_post_split_h(old, residual, post, comb) - ref_norm = hc_combine_norm(ref_post.flatten(1), pre, nw, 1e-6) - args = [gemm2, idx, weights, top_k, shared, residual, post, comb] - if epilogue != "post": - args.extend((pre, nw, 1e-6)) - outputs = kernel(*args, world_size=comm.world_size) - # Exercise the shared counters between generic push and row-cluster AR. - comm.custom_all_reduce(outputs[0]) - return old, ref_post, ref_norm, outputs - - chain() - torch.cuda.synchronize() - graph = torch.cuda.CUDAGraph() - with comm.capture(), torch.cuda.graph(graph): - old, ref_post, ref_norm, outputs = chain() - - def check_outputs(): - reduced, actual_post = outputs[:2] - assert torch.equal(old.view(torch.int16), reduced.view(torch.int16)) - assert torch.equal(ref_post.view(torch.int16), actual_post.view(torch.int16)) - if epilogue == "post": - return - actual_norm = outputs[2] - # The cluster and Triton RMS reductions have different addition orders. - torch.testing.assert_close(actual_norm, ref_norm, rtol=0.008, atol=0.0001) - # Different reduction trees can land on opposite sides of BF16 ties. - # Bound each value, rather than a data-dependent bit-identical fraction. - ulp = ( - actual_norm.view(torch.int16).int() - ref_norm.view(torch.int16).int() - ).abs() - ulp.masked_fill_(actual_norm == ref_norm, 0) # Treat signed zeros equally. - assert ulp.max() <= 1 - if epilogue == "quant": - q, sf = outputs[3:] - for backend in ("cuda", "cute-dsl"): - ref_q, ref_sf = flashinfer_mxfp8_quantize( - actual_norm, True, 32, backend - ) - assert torch.equal( - q.view(torch.uint8).flatten(), ref_q.view(torch.uint8).flatten() - ) - assert torch.equal(sf, ref_sf.flatten()) - - for replay, magnitude in enumerate((1e-3, 1.0, 1e3, 0.0, 1.0, 0.0, 1e-2)): - residual.normal_().mul_(magnitude) - gemm2.normal_().mul_(magnitude) - if epilogue == "quant": - outputs[4].fill_(0xAB) # Also require rewriting the padded scale rows. - if dist.get_rank() == replay % dist.get_world_size(): - torch.cuda._sleep(100000) - graph.replay() - torch.cuda.synchronize() - error = None - try: - check_outputs() - except AssertionError as exc: - error = f"rank={dist.get_rank()}, replay={replay}, scale={magnitude}: {exc}" - errors = [None] * dist.get_world_size() - # A failed assertion must stop all ranks before the next collective. - dist.all_gather_object(errors, error) - assert not any(errors), "\n".join(e for e in errors if e) - - -if __name__ == "__main__": - multigpu_pytest_main( - __name__, - __file__, - num_gpus=(4,), - pre_launch_fn=_precompile, - ) diff --git a/test/registered/kernel/communication/test_nvlink_comm.py b/test/registered/kernel/communication/test_nvlink_comm.py index b0ec65c3b067..7fb7018b0931 100644 --- a/test/registered/kernel/communication/test_nvlink_comm.py +++ b/test/registered/kernel/communication/test_nvlink_comm.py @@ -135,10 +135,17 @@ def _reference(op, x, residual, nccl_group, world_size): } -@pytest.mark.parametrize("tokens", [1, 7, 128, 1024]) -@pytest.mark.parametrize("residual", [False, True]) -@pytest.mark.parametrize("plane", ["push", "pull"]) -@pytest.mark.parametrize("op", nvl.SUPPORTED_OPS) +@pytest.mark.parametrize( + "op,plane,residual,tokens", + [ + ("all_reduce", "push", False, 1), + ("all_reduce", "pull", True, 7), + ("all_gather", "push", True, 7), + ("all_gather", "pull", False, 1024), + ("reduce_scatter", "push", False, 7), + ("reduce_scatter", "pull", True, 1024), + ], +) def test_collective(op: str, plane: str, residual: bool, tokens: int) -> None: cpu_group, nccl_group = _init_world() comm = _init_comm() @@ -171,13 +178,70 @@ def test_collective(op: str, plane: str, residual: bool, tokens: int) -> None: torch.testing.assert_close(sym_out.float(), ref, atol=0.1, rtol=0.02) -@pytest.mark.parametrize("tokens", [1, 7, 128]) +@pytest.mark.parametrize("op", ["all_gather", "reduce_scatter"]) +@pytest.mark.parametrize("plane", ["push", "pull"]) +def test_ragged_collective_with_empty_ranks(op: str, plane: str) -> None: + """Use fewer total rows than ranks so AG/RS exercise zero-row shards.""" + cpu_group, nccl_group = _init_world() + comm = _init_comm() + world_size = dist.get_world_size(cpu_group) + rank = dist.get_rank(cpu_group) + device = _device() + total_tokens = world_size - 1 + partition = nvl.get_token_partition(total_tokens, comm.obj) + routes = [ + total_tokens // world_size + (peer < total_tokens % world_size) + for peer in range(world_size) + ] + max_local_tokens = max(routes) + + if op == "all_gather": + input_storage = _symm((max_local_tokens, HIDDEN)) + sym_in = input_storage[: partition.num_local_tokens] + sym_out = _symm((total_tokens, HIDDEN)) + else: + sym_in = _symm((total_tokens, HIDDEN)) + output_storage = _symm((max_local_tokens, HIDDEN)) + sym_out = output_storage[: partition.num_local_tokens] + + gen = torch.Generator(device=device).manual_seed(9000 + rank) + sym_in.copy_(torch.randn(sym_in.shape, dtype=DTYPE, device=device, generator=gen)) + sym_out.zero_() + + if op == "all_gather": + padded = torch.zeros( + (max_local_tokens, HIDDEN), dtype=torch.float32, device=device + ) + padded[: partition.num_local_tokens].copy_(sym_in.float()) + gathered = [torch.empty_like(padded) for _ in range(world_size)] + dist.all_gather(gathered, padded, group=nccl_group) + ref = torch.cat([chunk[:count] for chunk, count in zip(gathered, routes)]) + else: + reduced = sym_in.float().clone() + dist.all_reduce(reduced, group=nccl_group) + begin = partition.num_prefix_tokens + ref = reduced[begin : begin + partition.num_local_tokens] + + dist.barrier(nccl_group) + torch.cuda.synchronize() + _FNS[(op, plane)](comm.obj, sym_in, sym_out) + torch.cuda.synchronize() + dist.barrier(nccl_group) + if op == "all_gather": + torch.testing.assert_close( + sym_out.float(), ref.to(DTYPE).float(), atol=0, rtol=0 + ) + else: + torch.testing.assert_close(sym_out.float(), ref, atol=0.1, rtol=0.02) + + @pytest.mark.parametrize("variant", ["multicast", "unicast"]) -def test_copy_engine_all_gather(variant: str, tokens: int) -> None: +def test_copy_engine_all_gather(variant: str) -> None: cpu_group, nccl_group = _init_world() comm = _init_comm() world_size = dist.get_world_size(cpu_group) device = _device() + tokens = 7 in_shape, out_shape = _shapes("all_gather", tokens, world_size) sym_in, sym_out = _symm(in_shape), _symm(out_shape) gen = torch.Generator(device=device).manual_seed( diff --git a/test/registered/kernel/communication/test_vocab_gather.py b/test/registered/kernel/communication/test_vocab_gather.py deleted file mode 100644 index bb7ddaeaeb0d..000000000000 --- a/test/registered/kernel/communication/test_vocab_gather.py +++ /dev/null @@ -1,141 +0,0 @@ -"""``VocabGather``: the TP vocab-parallel row gather behind the DSpark draft head. - -On a TP group built the way the server builds it (custom all-reduce on, so -``ca_comm`` is a CustomAllReduceV2), ``make_vocab_gather`` must pick the NVLink -gather, and every path of that gather (push kernel, pull kernel into the -symmetric-memory output, NCCL past its capacity) must match the NCCL -``all_gather(dim=-1)`` reference, hand back a result that does not alias the -shared output, and replay correctly inside a CUDA graph. - -Usage:: - - python test/registered/kernel/communication/test_vocab_gather.py --num-gpu 4 -""" - -from __future__ import annotations - -import atexit -import os - -import pytest -import torch -import torch.distributed as dist - -import sglang.srt.distributed.parallel_state as ps -from sglang.kernels.jit.utils import cache_once -from sglang.srt.distributed.device_communicators.vocab_gather import ( - NcclVocabGather, - NVLinkVocabGather, - make_vocab_gather, -) -from sglang.srt.distributed.parallel_state import GroupCoordinator -from sglang.test.ci.ci_register import register_cuda_ci -from sglang.test.kernels.utils import multigpu_pytest_main - -register_cuda_ci(est_time=60, stage="base-b-kernel-unit", runner_config="4-gpu-b200") - -LOCAL_WIDTH = 32320 # DeepSeek-V4.1's 129280-entry vocab over TP4 -SYMM_ROWS = 256 - - -def _device() -> torch.device: - return torch.device("cuda", int(os.environ["LOCAL_RANK"])) - - -@cache_once -def _tp_group() -> GroupCoordinator: - local_rank = int(os.environ["LOCAL_RANK"]) - world_size = int(os.environ["WORLD_SIZE"]) - torch.cuda.set_device(local_rank) - dist.init_process_group(backend="gloo") - ps._WORLD = ps.init_world_group( - ranks=list(range(world_size)), local_rank=local_rank, backend="nccl" - ) - atexit.register(dist.destroy_process_group) - return GroupCoordinator( - group_ranks=[list(range(world_size))], - local_rank=local_rank, - torch_distributed_backend="nccl", - use_pynccl=False, - use_pymscclpp=False, - use_custom_allreduce=True, - use_torch_symm_mem_all_reduce=False, - use_hpu_communicator=False, - use_xpu_communicator=False, - use_npu_communicator=False, - group_name="vocab_gather_test", - ) - - -@cache_once -def _gathers(): - tp = _tp_group() - nvlink = make_vocab_gather(tp, local_width=LOCAL_WIDTH, symm_rows=SYMM_ROWS) - if not isinstance(nvlink, NVLinkVocabGather): - pytest.skip("this TP group has no multicast plane") - nccl = make_vocab_gather(tp, local_width=LOCAL_WIDTH, prefer_nvlink=False) - return nvlink, nccl - - -def _rows(rows: int, seed: int) -> torch.Tensor: - gen = torch.Generator(device=_device()).manual_seed(seed + dist.get_rank()) - x = torch.randn(rows, LOCAL_WIDTH, device=_device(), generator=gen) - x[:, ::97] = 0.0 # exact zeros: the push kernel's Lamport sentinel path - return x - - -def _sync() -> None: - torch.cuda.synchronize() - dist.barrier() - - -def _check(nvlink: NVLinkVocabGather, nccl: NcclVocabGather, x: torch.Tensor) -> None: - ref = nccl(x) - got = nvlink(x) - _sync() - assert got.shape == ref.shape - assert bool((got == ref).all()) # -0.0 == 0.0, so the sentinel flip is invisible - if nvlink.pull_out is not None: - lo = nvlink.pull_out.data_ptr() - hi = lo + nvlink.pull_out.numel() * nvlink.pull_out.element_size() - assert not (lo <= got.data_ptr() < hi), "result aliases the shared pull output" - - -@pytest.mark.parametrize("rows", [1, 3, 6, 25, 64, 200, SYMM_ROWS + 1]) -def test_matches_nccl(rows: int) -> None: - nvlink, nccl = _gathers() - _check(nvlink, nccl, _rows(rows, seed=rows)) - - -def test_consecutive_pull_results_survive() -> None: - nvlink, nccl = _gathers() - a, b = _rows(64, seed=1), _rows(64, seed=2) - ra, rb = nvlink(a), nvlink(b) - _sync() - assert bool((ra == nccl(a)).all()) and bool((rb == nccl(b)).all()) - - -@pytest.mark.parametrize("rows", [1, 64]) -def test_cuda_graph_replay(rows: int) -> None: - nvlink, nccl = _gathers() - x = _rows(rows, seed=100 + rows) - nvlink(x) # eager once: JIT load - _sync() - graph = torch.cuda.CUDAGraph() - stream = torch.cuda.Stream() - with torch.cuda.stream(stream): - _sync() - with torch.cuda.graph(graph, stream=stream): - out = nvlink(x) - _sync() - for it in range(3): - x.copy_(_rows(rows, seed=200 + it)) - ref = nccl(x) - _sync() - graph.replay() - _sync() - assert bool((out == ref).all()) - - -if __name__ == "__main__": - multigpu_pytest_main(__name__, __file__, num_gpus=(4, 8)) diff --git a/test/registered/kernel/speculative/test_dspark_sharded_greedy.py b/test/registered/kernel/speculative/test_dspark_sharded_greedy.py deleted file mode 100644 index 2ea62b48f378..000000000000 --- a/test/registered/kernel/speculative/test_dspark_sharded_greedy.py +++ /dev/null @@ -1,144 +0,0 @@ -"""Four-rank selection equivalence, including graph replay and padded shards.""" - -import atexit -import os - -import pytest -import torch -import torch.distributed as dist - -import sglang.srt.distributed.parallel_state as ps -from sglang.kernels.jit.utils import cache_once -from sglang.kernels.ops.speculative.dspark.sharded_greedy import sharded_greedy_step -from sglang.srt.distributed.device_communicators.vocab_gather import ( - NVLinkVocabGather, - make_vocab_gather, -) -from sglang.srt.distributed.parallel_state import GroupCoordinator -from sglang.test.ci.ci_register import register_cuda_ci -from sglang.test.kernels.utils import multigpu_pytest_main - -register_cuda_ci(est_time=120, stage="base-b-kernel-unit", runner_config="4-gpu-b200") - - -@cache_once -def group(): - rank, world = int(os.environ["LOCAL_RANK"]), int(os.environ["WORLD_SIZE"]) - torch.cuda.set_device(rank) - dist.init_process_group(backend="gloo") - ps._WORLD = ps.init_world_group(list(range(world)), rank, backend="nccl") - atexit.register(dist.destroy_process_group) - torch.cuda.set_stream(torch.cuda.Stream()) - return GroupCoordinator( - group_ranks=[list(range(world))], - local_rank=rank, - torch_distributed_backend="nccl", - use_pynccl=False, - use_pymscclpp=False, - use_custom_allreduce=True, - use_torch_symm_mem_all_reduce=False, - use_hpu_communicator=False, - use_xpu_communicator=False, - use_npu_communicator=False, - group_name="sharded_greedy_test", - ) - - -_NVLINK = pytest.mark.parametrize("nvlink", [False, True]) -_ROWS = pytest.mark.parametrize("m", [1, 4, 64]) -_SHARDS = pytest.mark.parametrize("width,last", [(32320, 32320), (8192, 17), (8, 0)]) - - -def _replay_and_check(m, width, last, nvlink, perturb): - """Capture the 4-rank selection graph, then replay it under `perturb`.""" - g = group() - transport = make_vocab_gather( - g, local_width=width, prefer_nvlink=nvlink, symm_rows=0 - ) - if nvlink and not isinstance(transport, NVLinkVocabGather): - pytest.skip("this TP group has no multicast plane") - rank = g.rank_in_group - real = last if rank == g.world_size - 1 else width - # A slice of a block's logits has a non-contiguous row stride. - storage = torch.randn(m, 5, width, device="cuda") - base = storage[:, 2] - bias = torch.randn(m, real, device="cuda", dtype=torch.bfloat16) - - def candidate(): - return sharded_greedy_step( - bias, - base, - group=g, - vocab_start=rank * width, - gather=transport.gather_stacked, - ) - - for _ in range(3): - candidate() - graph = torch.cuda.CUDAGraph() - with torch.cuda.graph(graph): - out = candidate() - for replay in range(4): - storage.normal_() - bias.normal_() - perturb(replay, storage, base, bias, real) - graph.replay() - # Independent reference: gather complete, correctly padded FP32 logits. - local = torch.full((m, width), -float("inf"), device="cuda") - local[:, :real] = base[:, :real] + bias.float() - full = g.all_gather(local, dim=-1) - ref = full[:, : (g.world_size - 1) * width + last].argmax(-1) - torch.cuda.synchronize() - assert torch.equal(out, ref) - - -def _random(replay, storage, base, bias, real): - pass - - -def _tie(replay, storage, base, bias, real): - storage.zero_() - bias.zero_() - - -def _nan(replay, storage, base, bias, real): - if real: - base[:, replay % real] = float("nan") - - -def _inf(replay, storage, base, bias, real): - storage.fill_(-float("inf")) - if replay % 2 and real: - base[:, replay % real] = float("inf") - - -@_NVLINK -@_ROWS -@_SHARDS -def test_random_logits(m, width, last, nvlink): - _replay_and_check(m, width, last, nvlink, _random) - - -@_NVLINK -@_ROWS -@_SHARDS -def test_ties_resolve_to_the_lowest_index(m, width, last, nvlink): - _replay_and_check(m, width, last, nvlink, _tie) - - -@_NVLINK -@_ROWS -@_SHARDS -def test_nan_propagates(m, width, last, nvlink): - _replay_and_check(m, width, last, nvlink, _nan) - - -@_NVLINK -@_ROWS -@_SHARDS -def test_infinities(m, width, last, nvlink): - _replay_and_check(m, width, last, nvlink, _inf) - - -if __name__ == "__main__": - multigpu_pytest_main(__name__, __file__, num_gpus=(4,)) From 5a8166b6cd2d8c558557fde7e2778c19deb8cc87 Mon Sep 17 00:00:00 2001 From: BBuf <1182563586@qq.com> Date: Wed, 16 Sep 2026 13:25:46 +0800 Subject: [PATCH 09/30] remove standalone nvlink communication test --- .../kernel/communication/test_nvlink_comm.py | 266 ------------------ 1 file changed, 266 deletions(-) delete mode 100644 test/registered/kernel/communication/test_nvlink_comm.py diff --git a/test/registered/kernel/communication/test_nvlink_comm.py b/test/registered/kernel/communication/test_nvlink_comm.py deleted file mode 100644 index 7fb7018b0931..000000000000 --- a/test/registered/kernel/communication/test_nvlink_comm.py +++ /dev/null @@ -1,266 +0,0 @@ -"""Correctness of the NVLink collectives (``nvlink_comm``) against NCCL. - -All-reduce, all-gather and reduce-scatter on the push and pull planes of a -``CustomAllReduceV2`` communicator, with and without the folded residual, plus -the two copy-engine all-gathers, over token counts that cover the ragged split -(7), the remainder loop alone (1) and the bandwidth band (1024). The residual -is the same on every rank, as it is in a TP layer: the pull all-reduce folds -it in over each rank's token slice. - -Usage:: - - python test/registered/kernels/ops/communication/test_nvlink_comm.py --num-gpu 4 -""" - -from __future__ import annotations - -import atexit -import os -from typing import Dict, Tuple - -import pytest -import torch -import torch.distributed as dist - -import sglang.srt.distributed.parallel_state as ps -from sglang.kernels.jit.utils import cache_once -from sglang.kernels.ops.communication import nvlink_comm as nvl -from sglang.kernels.ops.communication.mp import register_comm_cleanup -from sglang.test.ci.ci_register import register_cuda_ci -from sglang.test.kernels.utils import multigpu_pytest_main - -register_cuda_ci(est_time=90, stage="base-b-kernel-unit", runner_config="4-gpu-b200") - -HIDDEN = 7168 -DTYPE = torch.bfloat16 -PUSH_SLOT_MB = 32 -PULL_MB = 4 - - -def _device() -> torch.device: - return torch.device("cuda", int(os.environ["LOCAL_RANK"])) - - -@cache_once -def _init_world(): - local_rank = int(os.environ["LOCAL_RANK"]) - world_size = int(os.environ["WORLD_SIZE"]) - torch.cuda.set_device(local_rank) - dist.init_process_group(backend="gloo") - ps._WORLD = coord = ps.init_world_group( - ranks=list(range(world_size)), local_rank=local_rank, backend="nccl" - ) - atexit.register(dist.destroy_process_group) - nccl_group = dist.new_group(backend="nccl", device_id=_device()) - return coord.cpu_group, nccl_group - - -@cache_once -def _init_comm(): - from sglang.srt.distributed.device_communicators.custom_all_reduce_v2 import ( - CustomAllReduceV2, - ) - - cpu_group, _ = _init_world() - comm = CustomAllReduceV2( - cpu_group, - _device(), - max_push_size=PUSH_SLOT_MB << 20, - max_pull_size=PULL_MB << 20, - ) - if comm.disabled: - pytest.skip("CustomAllReduceV2 is disabled on this system") - if not comm.has_multicast: - pytest.skip("the nvlink collectives need a multicast plane") - register_comm_cleanup(comm) - return comm - - -_SYMM: Dict[Tuple[int, int], torch.Tensor] = {} - - -def _symm(shape: Tuple[int, int]) -> torch.Tensor: - """Symmetric memory with a multicast alias; one allocation per shape, since - the allocation is collective and never returned.""" - from torch._C._distributed_c10d import _SymmetricMemory - - if shape not in _SYMM: - cpu_group, _ = _init_world() - t = _SymmetricMemory.empty_strided_p2p( - (shape[0] * shape[1],), [1], DTYPE, _device(), cpu_group.group_name - ) - _SymmetricMemory.rendezvous(t) - _SYMM[shape] = t.view(shape) - return _SYMM[shape] - - -def _shapes(op: str, tokens: int, world_size: int): - if op == "all_gather": - return (tokens, HIDDEN), (tokens * world_size, HIDDEN) - if op == "reduce_scatter": - return (tokens * world_size, HIDDEN), (tokens, HIDDEN) - return (tokens, HIDDEN), (tokens, HIDDEN) - - -def _reference(op, x, residual, nccl_group, world_size): - """fp32 NCCL reference with the kernels' residual placement: the gather adds - it to this rank's shard before gathering, the reductions to the output.""" - x = x.float() - res = residual.float() if residual is not None else 0 - if op == "all_reduce": - y = x.clone() - dist.all_reduce(y, group=nccl_group) - return y + res - if op == "all_gather": - x = (x + res).contiguous() - out = torch.empty( - (x.shape[0] * world_size, HIDDEN), dtype=torch.float32, device=x.device - ) - dist.all_gather_into_tensor(out, x, group=nccl_group) - return out - out = torch.empty( - (x.shape[0] // world_size, HIDDEN), dtype=torch.float32, device=x.device - ) - dist.reduce_scatter_tensor(out, x.contiguous(), group=nccl_group) - return out + res - - -_FNS = { - ("all_reduce", "push"): nvl.all_reduce_push, - ("all_gather", "push"): nvl.all_gather_push, - ("reduce_scatter", "push"): nvl.reduce_scatter_push, - ("all_reduce", "pull"): nvl.all_reduce_pull, - ("all_gather", "pull"): nvl.all_gather_pull, - ("reduce_scatter", "pull"): nvl.reduce_scatter_pull, -} - - -@pytest.mark.parametrize( - "op,plane,residual,tokens", - [ - ("all_reduce", "push", False, 1), - ("all_reduce", "pull", True, 7), - ("all_gather", "push", True, 7), - ("all_gather", "pull", False, 1024), - ("reduce_scatter", "push", False, 7), - ("reduce_scatter", "pull", True, 1024), - ], -) -def test_collective(op: str, plane: str, residual: bool, tokens: int) -> None: - cpu_group, nccl_group = _init_world() - comm = _init_comm() - world_size = dist.get_world_size(cpu_group) - rank = dist.get_rank(cpu_group) - device = _device() - in_shape, out_shape = _shapes(op, tokens, world_size) - sym_in, sym_out = _symm(in_shape), _symm(out_shape) - gen = torch.Generator(device=device).manual_seed(1000 * tokens + rank) - sym_in.copy_(torch.randn(in_shape, dtype=DTYPE, device=device, generator=gen)) - sym_out.zero_() - res = None - if residual: - shared = torch.Generator(device=device).manual_seed(7 * tokens) - res_shape = in_shape if op == "all_gather" else out_shape - res = torch.randn(res_shape, dtype=DTYPE, device=device, generator=shared) - ref = _reference(op, sym_in, res, nccl_group, world_size) - dist.barrier(nccl_group) - torch.cuda.synchronize() - _FNS[(op, plane)](comm.obj, sym_in, sym_out, res) - torch.cuda.synchronize() - dist.barrier(nccl_group) - if op == "all_gather": - # a pure copy (plus one bf16 add with the residual): bit-exact - torch.testing.assert_close( - sym_out.float(), ref.to(DTYPE).float(), atol=0, rtol=0 - ) - else: - # bf16 sums in a different order than NCCL's fp32 tree - torch.testing.assert_close(sym_out.float(), ref, atol=0.1, rtol=0.02) - - -@pytest.mark.parametrize("op", ["all_gather", "reduce_scatter"]) -@pytest.mark.parametrize("plane", ["push", "pull"]) -def test_ragged_collective_with_empty_ranks(op: str, plane: str) -> None: - """Use fewer total rows than ranks so AG/RS exercise zero-row shards.""" - cpu_group, nccl_group = _init_world() - comm = _init_comm() - world_size = dist.get_world_size(cpu_group) - rank = dist.get_rank(cpu_group) - device = _device() - total_tokens = world_size - 1 - partition = nvl.get_token_partition(total_tokens, comm.obj) - routes = [ - total_tokens // world_size + (peer < total_tokens % world_size) - for peer in range(world_size) - ] - max_local_tokens = max(routes) - - if op == "all_gather": - input_storage = _symm((max_local_tokens, HIDDEN)) - sym_in = input_storage[: partition.num_local_tokens] - sym_out = _symm((total_tokens, HIDDEN)) - else: - sym_in = _symm((total_tokens, HIDDEN)) - output_storage = _symm((max_local_tokens, HIDDEN)) - sym_out = output_storage[: partition.num_local_tokens] - - gen = torch.Generator(device=device).manual_seed(9000 + rank) - sym_in.copy_(torch.randn(sym_in.shape, dtype=DTYPE, device=device, generator=gen)) - sym_out.zero_() - - if op == "all_gather": - padded = torch.zeros( - (max_local_tokens, HIDDEN), dtype=torch.float32, device=device - ) - padded[: partition.num_local_tokens].copy_(sym_in.float()) - gathered = [torch.empty_like(padded) for _ in range(world_size)] - dist.all_gather(gathered, padded, group=nccl_group) - ref = torch.cat([chunk[:count] for chunk, count in zip(gathered, routes)]) - else: - reduced = sym_in.float().clone() - dist.all_reduce(reduced, group=nccl_group) - begin = partition.num_prefix_tokens - ref = reduced[begin : begin + partition.num_local_tokens] - - dist.barrier(nccl_group) - torch.cuda.synchronize() - _FNS[(op, plane)](comm.obj, sym_in, sym_out) - torch.cuda.synchronize() - dist.barrier(nccl_group) - if op == "all_gather": - torch.testing.assert_close( - sym_out.float(), ref.to(DTYPE).float(), atol=0, rtol=0 - ) - else: - torch.testing.assert_close(sym_out.float(), ref, atol=0.1, rtol=0.02) - - -@pytest.mark.parametrize("variant", ["multicast", "unicast"]) -def test_copy_engine_all_gather(variant: str) -> None: - cpu_group, nccl_group = _init_world() - comm = _init_comm() - world_size = dist.get_world_size(cpu_group) - device = _device() - tokens = 7 - in_shape, out_shape = _shapes("all_gather", tokens, world_size) - sym_in, sym_out = _symm(in_shape), _symm(out_shape) - gen = torch.Generator(device=device).manual_seed( - 50 * tokens + dist.get_rank(cpu_group) - ) - sym_in.copy_(torch.randn(in_shape, dtype=DTYPE, device=device, generator=gen)) - sym_out.zero_() - ref = _reference("all_gather", sym_in, None, nccl_group, world_size) - dist.barrier(nccl_group) - torch.cuda.synchronize() - if variant == "multicast": - nvl.all_gather_copy_engine_multicast(comm.obj, sym_in, sym_out) - else: - flags = nvl.make_ce_flags(cpu_group, world_size) - nvl.all_gather_copy_engine_unicast(comm.obj, sym_in, sym_out, ce_flags=flags) - torch.cuda.synchronize() - dist.barrier(nccl_group) - torch.testing.assert_close(sym_out.float(), ref, atol=0, rtol=0) - - -if __name__ == "__main__": - multigpu_pytest_main(__name__, __file__, num_gpus=(4, 8)) From 3c5130febbd9af1a0e4b1c6eb98705673df2d0d4 Mon Sep 17 00:00:00 2001 From: hnyls2002 Date: Tue, 15 Sep 2026 15:53:40 -0700 Subject: [PATCH 10/30] dsv4.1: extract compression and metadata kernels --- .../kernels/jit/csrc/deepseek_v4/c1.cuh | 322 +++++++++++++ .../kernels/jit/csrc/deepseek_v4/c2.cuh | 454 ++++++++++++++++++ .../sgl_kernel/deepseek_v4/fp4_utils.cuh | 64 +++ .../sgl_kernel/deepseek_v4/kv_layout.cuh | 220 +++++++++ .../sglang/kernels/ops/attention/dsv4/c1.py | 103 ++++ .../sglang/kernels/ops/attention/dsv4/c2.py | 179 +++++++ .../kernels/ops/attention/dsv4/kv_layout.py | 73 +++ .../ops/attention/dsv4/metadata_kernel.py | 60 +++ .../ops/attention/dsv41_small_metadata.py | 129 +++++ .../kernel/attention/dsv4/test_c2_verify.py | 391 +++++++++++++++ .../attention/dsv4/test_v41_kv_store.py | 433 +++++++++++++++++ .../attention/test_dsv41_small_metadata.py | 56 +++ 12 files changed, 2484 insertions(+) create mode 100644 python/sglang/kernels/jit/csrc/deepseek_v4/c1.cuh create mode 100644 python/sglang/kernels/jit/csrc/deepseek_v4/c2.cuh create mode 100644 python/sglang/kernels/jit/include/sgl_kernel/deepseek_v4/fp4_utils.cuh create mode 100644 python/sglang/kernels/jit/include/sgl_kernel/deepseek_v4/kv_layout.cuh create mode 100644 python/sglang/kernels/ops/attention/dsv4/c1.py create mode 100644 python/sglang/kernels/ops/attention/dsv4/c2.py create mode 100644 python/sglang/kernels/ops/attention/dsv4/kv_layout.py create mode 100644 python/sglang/kernels/ops/attention/dsv41_small_metadata.py create mode 100644 test/registered/kernel/attention/dsv4/test_c2_verify.py create mode 100644 test/registered/kernel/attention/dsv4/test_v41_kv_store.py create mode 100644 test/registered/kernel/attention/test_dsv41_small_metadata.py diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/c1.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/c1.cuh new file mode 100644 index 000000000000..6dac3e3c1f47 --- /dev/null +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/c1.cuh @@ -0,0 +1,322 @@ +#include +#include + +#include +#include +#include +#include +#include + +#include +#include +#include + +#include + +#include +#include + +namespace sglang { + +/// \brief Ratio-1 decode compressor: RMSNorm and the whole main-KV write. +/// +/// The ratio-1 compressor pools nothing -- `project` is a single bf16 `wkv` +/// GEMM and the latent it produces stands for the token itself -- so the +/// kernel's input is the GEMM output and its RoPE position is `positions`, not +/// `positions - 1`. `kv_output` is the pre-RoPE latent, for the index-K +/// branch's `wk` projection. +struct C1Params { + const bf16_t* __restrict__ kv_input; // [num_tokens, kHeadDim] bf16 + bf16_t* __restrict__ kv_output; // [num_tokens, kHeadDim] bf16, pre-RoPE + const bf16_t* __restrict__ norm_weight; // [kHeadDim] bf16 + const float* __restrict__ freqs_cis; // [max_pos, kRopeDim] fp32, real/imag interleaved + const void* __restrict__ positions; // [num_tokens] PosT + const void* __restrict__ out_loc; // [num_tokens] LocT compressed slot; 0 marks a padded row + uint8_t* __restrict__ kvcache; // [npages, kPageBytes] uint8 + float eps; +}; + +/// Elements per thread; 256 threads per token was measured fastest on B200 decode batches. +/// At head_dim 512, (512 - 64) / 2 = 224 threads keeps the nope/rope split warp-aligned; +/// the fp8 amax reduction requires every lane in its full-warp mask to participate. +constexpr uint32_t kC1VecSize = 2; + +/// \brief RMSNorm + RoPE tail + fp4 fake-quant + the FlashMLA store. +/// +/// One CTA per token, `kHeadDim / kC1VecSize` threads over the row. +/// +/// The three reductions have three different widths and are not +/// interchangeable: the RMSNorm statistic spans the row, an fp8 store scale +/// spans 64 elements, an fp4 block spans 16. All asserted below. +/// +/// kLayout is the cache's page format. V4 (584 B/token) and V41 (528 B/token, +/// fp8 with per-32 scales) store the fake-quantized value; V41_FP4 (288 B/token) +/// stores the e2m1 codes and their e4m3 scales directly, so the fp4 rounding +/// happens once and no fp8 rounding follows it. +template < + int64_t kHeadDim, + int64_t kRopeDim, + int32_t kPageBits, + typename PosT, + typename LocT, + deepseek_v4::KVLayout kLayout, + bool kUsePDL> +__global__ +__launch_bounds__(kHeadDim / kC1VecSize) void flash_c1_decode_kernel(const __grid_constant__ C1Params params) { + using namespace device; + using deepseek_v4::KVLayout; + using deepseek_v4::fp8::cast_to_ue8m0; + using deepseek_v4::fp8::inv_scale_ue8m0; + using deepseek_v4::fp8::pack_fp8; + + /// Threads over one token, and the leading ones of those that carry the fp8 + /// nope part; the rest carry the bf16 RoPE tail. + constexpr uint32_t kVecSize = kC1VecSize; + constexpr uint32_t kRowLanes = kHeadDim / kVecSize; + constexpr uint32_t kNopeLanes = (kHeadDim - kRopeDim) / kVecSize; + constexpr uint32_t kRowWarps = kRowLanes / kWarpThreads; + constexpr uint32_t kFp8Lanes = 64 / kVecSize; + constexpr uint32_t kFp4Lanes = deepseek_v4::fp4::kCompressedKVBlockSize / kVecSize; + using Paged = deepseek_v4::PagedKV; + static_assert(kHeadDim == 512 && kRopeDim == 64, "the FlashMLA layouts require (512, 64)"); + static_assert(kHeadDim % kVecSize == 0 && kVecSize % 2 == 0); + static_assert(kRowLanes % kWarpThreads == 0, "a token owns a whole number of warps"); + static_assert(kNopeLanes % kFp8Lanes == 0, "the nope part must end on an fp8 scale block"); + static_assert( + (kHeadDim - kRopeDim) % deepseek_v4::fp4::kCompressedKVBlockSize == 0, + "no fp4 block may straddle the nope/rope seam"); + static_assert(kFp8Lanes <= kWarpThreads && kFp4Lanes <= kWarpThreads); + + using bf16_vec_t = AlignedVector; + using fp8_vec_t = AlignedVector; + using freq_vec_t = AlignedVector; + + const uint32_t tx = threadIdx.x; + const uint32_t row = blockIdx.x; + + // `out_loc` and `positions` are step metadata, independent of the PDL producer. + // Slots fit in int32; padded rows are suppressed at the cache store. + const auto out_loc = static_cast(static_cast(params.out_loc)[row]); + const auto position = static_cast(static_cast(params.positions)[row]); + PDLWaitPrimary(); + + float data[kVecSize]; + bf16_vec_t latent; + { + bf16_vec_t input, weight; + input.load(params.kv_input + row * kHeadDim, tx); + weight.load(params.norm_weight, tx); + + // `project` already returns bf16 at ratio 1, so `finish`'s `.to(bfloat16)` + // is a no-op and the statistic is taken over the loaded values as they are. + float local_sqrsum = 0.0f; +#pragma unroll + for (uint32_t j = 0; j < kVecSize / 2; ++j) { + const auto [x, y] = cast(input[j]); + local_sqrsum += x * x; + local_sqrsum += y * y; + data[j * 2 + 0] = x; + data[j * 2 + 1] = y; + } + + __shared__ float s_warp_sum[kRowWarps]; + s_warp_sum[tx / kWarpThreads] = warp::reduce_sum(local_sqrsum); + __syncthreads(); + float sqrsum = 0.0f; +#pragma unroll + for (uint32_t i = 0; i < kRowWarps; ++i) { + sqrsum += s_warp_sum[i]; + } + constexpr float kInvHeadDim = 1.0f / static_cast(kHeadDim); + const auto norm_factor = math::rsqrt(sqrsum * kInvHeadDim + params.eps); + +#pragma unroll + for (uint32_t j = 0; j < kVecSize / 2; ++j) { + const auto [wx, wy] = cast(weight[j]); + const auto x = data[j * 2 + 0] * norm_factor * wx; + const auto y = data[j * 2 + 1] * norm_factor * wy; + latent[j] = cast(fp32x2_t{x, y}); + } + } + + latent.store(params.kv_output + row * kHeadDim, tx); + PDLTriggerSecondary(); + + // Match finish()'s bf16 rounding before the main-KV RoPE. +#pragma unroll + for (uint32_t j = 0; j < kVecSize / 2; ++j) { + const auto [x, y] = cast(latent[j]); + data[j * 2 + 0] = x; + data[j * 2 + 1] = y; + } + + if (tx >= kNopeLanes) { + // `rope_tail` ends in `.to(x.dtype)`, so the rotated value is rounded back + // to bf16 before the fake-quant widens it again. + freq_vec_t freq; + freq.load(params.freqs_cis + position * kRopeDim, tx - kNopeLanes); +#pragma unroll + for (uint32_t j = 0; j < kVecSize / 2; ++j) { + const auto k = j * 2; + const auto x_real = data[k + 0]; + const auto x_imag = data[k + 1]; + const auto f_real = freq[k + 0]; + const auto f_imag = freq[k + 1]; + const auto rotated = + cast(fp32x2_t{x_real * f_real - x_imag * f_imag, x_real * f_imag + x_imag * f_real}); + const auto [r0, r1] = cast(rotated); + data[k + 0] = r0; + data[k + 1] = r1; + } + } + + if constexpr (kLayout == KVLayout::V41_FP4) { + // The fp4 cache takes the rotated bf16 value as is: its row quantizer is the + // fake quantization, minus the dequantization. + if (out_loc <= 0) return; + const auto kv_row = Paged::row(params.kvcache, out_loc); + return deepseek_v4::v41::store_row(kv_row.data, kv_row.scale, tx, data); + } + + // FP4/E4M3 fake-quant over 16 elements, i.e. kFp4Lanes threads. + { + float amax = fabsf(data[0]); +#pragma unroll + for (uint32_t i = 1; i < kVecSize; ++i) { + amax = fmaxf(amax, fabsf(data[i])); + } + amax = warp::reduce_max(amax); + const auto scale = deepseek_v4::fp4::compressed_kv_scale(amax); +#pragma unroll + for (uint32_t i = 0; i < kVecSize / 2; ++i) { + const auto [x, y] = deepseek_v4::fp4::fake_quant_compressed_kv_x2({data[i * 2 + 0], data[i * 2 + 1]}, scale); + data[i * 2 + 0] = x; + data[i * 2 + 1] = y; + } + } + + // A padded CUDA-graph row carries `out_loc == 0`, the reserved dummy slot, + // and must publish nothing: at ratio 1 the compressed slot *is* the FULL + // slot, so there is nothing to divide and no other marker to read. + if (out_loc <= 0) return; + const auto kv_row = Paged::row(params.kvcache, out_loc); + + if constexpr (kLayout == KVLayout::V41) { + // fp8 with one ue8m0 scale per 32 elements over the whole row, RoPE included. + return deepseek_v4::v41::store_row(kv_row.data, kv_row.scale, tx, data); + } + + const auto value_ptr = kv_row.data; + + if (tx >= kNopeLanes) { + bf16_vec_t rope_out; +#pragma unroll + for (uint32_t j = 0; j < kVecSize / 2; ++j) { + rope_out[j] = cast(fp32x2_t{data[j * 2 + 0], data[j * 2 + 1]}); + } + rope_out.store(value_ptr + (kHeadDim - kRopeDim), tx - kNopeLanes); + } else { + // fp8 e4m3 with one ue8m0 scale per 64 elements. + float abs_max = fabsf(data[0]); +#pragma unroll + for (uint32_t i = 1; i < kVecSize; ++i) { + abs_max = fmaxf(abs_max, fabsf(data[i])); + } + abs_max = warp::reduce_max(abs_max); + const auto scale_ue8m0 = cast_to_ue8m0(fmaxf(1e-4f, abs_max) / math::FP8_E4M3_MAX); + const auto inv_scale = inv_scale_ue8m0(scale_ue8m0); + fp8_vec_t nope_out; +#pragma unroll + for (uint32_t i = 0; i < kVecSize / 2; ++i) { + nope_out[i] = pack_fp8(data[i * 2 + 0] * inv_scale, data[i * 2 + 1] * inv_scale); + } + nope_out.store(value_ptr, tx); + kv_row.scale[tx / kFp8Lanes] = scale_ue8m0; + } +} + +/// \brief Host side of `flash_c1_decode_kernel`. +template +struct FlashC1DecodeKernel { + static constexpr int32_t kPageBits = std::bit_width(kPageSize) - 1; + static constexpr int64_t kPageBytes = deepseek_v4::kv_page_bytes(kPageSize); + static constexpr uint32_t kBlockSize = kHeadDim / kC1VecSize; + + static_assert(std::has_single_bit(kPageSize), "the page/slot split needs a power-of-two page"); + static_assert(kBlockSize % device::kWarpThreads == 0 && kBlockSize <= 1024); + static_assert(kLayout != deepseek_v4::KVLayout::V4 || kPageBytes == host::div_ceil(584ll * kPageSize, 576) * 576); + + template + static constexpr auto kernel = flash_c1_decode_kernel; + + /// \brief The (`positions`, `out_loc`) dtype pair, resolved at run time. + static auto select(const bool pos_i32, const bool loc_i32) { + if (pos_i32) return loc_i32 ? kernel : kernel; + return loc_i32 ? kernel : kernel; + } + + /// \brief RMSNorm + RoPE + fp4 fake-quant + the FlashMLA store, one launch. + /// + /// \param kv_input `[num_tokens, kHeadDim]` bf16, the `wkv` projection. + /// \param kv_output `[num_tokens, kHeadDim]` bf16, the pre-RoPE latent. + /// \param norm_weight `[kHeadDim]` bf16. + /// \param freqs_cis `[max_pos, kRopeDim]` fp32, real/imag interleaved. + /// \param positions `[num_tokens]` int32 or int64, indexed as-is. + /// \param out_loc `[num_tokens]` int32 or int64, the compressed slot; `0` is a padded row. + /// \param kvcache `[npages, kPageBytes]` uint8, or the pool's fp8 view of it. + static void run_decode_fusion( + const tvm::ffi::TensorView kv_input, + const tvm::ffi::TensorView kv_output, + const tvm::ffi::TensorView norm_weight, + const tvm::ffi::TensorView freqs_cis, + const tvm::ffi::TensorView positions, + const tvm::ffi::TensorView out_loc, + const tvm::ffi::TensorView kvcache, + const float eps) { + using namespace host; + + auto N = SymbolicSize{"num_tokens"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + + TensorMatcher({N, kHeadDim}) // + .with_dtype() + .with_device(device_) + .verify(kv_input) + .verify(kv_output); + TensorMatcher({kHeadDim}).with_dtype().with_device(device_).verify(norm_weight); + // Real/imag interleaved, so the trailing dim is kRopeDim, not kRopeDim / 2. + TensorMatcher({-1, kRopeDim}).with_dtype().with_device(device_).verify(freqs_cis); + // The scheduler's `out_cache_loc` (which `c1_out_loc` aliases at ratio 1) + // is int64; the unit tests hand int32. Both are indexed as-is. + auto pos_dtype = SymbolicDType{}; + auto loc_dtype = SymbolicDType{}; + TensorMatcher({N}).with_dtype(pos_dtype).with_device(device_).verify(positions); + TensorMatcher({N}).with_dtype(loc_dtype).with_device(device_).verify(out_loc); + // The pool allocates the buffer as uint8 and hands it out viewed as its fp8 + // dtype (`get_extra_key_buffer`); both are one byte per element. + TensorMatcher({-1, kPageBytes}).with_dtype().with_device(device_).verify(kvcache); + + const auto num_tokens = static_cast(N.unwrap()); + if (num_tokens == 0) return; + + const auto params = C1Params{ + .kv_input = static_cast(kv_input.data_ptr()), + .kv_output = static_cast(kv_output.data_ptr()), + .norm_weight = static_cast(norm_weight.data_ptr()), + .freqs_cis = static_cast(freqs_cis.data_ptr()), + .positions = positions.data_ptr(), + .out_loc = out_loc.data_ptr(), + .kvcache = static_cast(kvcache.data_ptr()), + .eps = eps, + }; + const auto k = select(pos_dtype.is_type(), loc_dtype.is_type()); + LaunchKernel(num_tokens, kBlockSize, device_.unwrap()) // + .enable_pdl(kUsePDL)(k, params); + } +}; + +// ensure that C++ wrapper can work +using enum deepseek_v4::KVLayout; + +} // namespace sglang diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/c2.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/c2.cuh new file mode 100644 index 000000000000..f16c308885cf --- /dev/null +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/c2.cuh @@ -0,0 +1,454 @@ +#include +#include + +#include +#include +#include +#include +#include + +#include +#include +#include + +#include + +#include +#include +#include + +namespace sglang { + +/// \brief Ratio-2 decode compressor: pair-pool, RMSNorm and the main-KV write. +/// +/// `kv_input` and `kv_state` rows are `2 * kHeadDim` floats, kv then score. +/// `kv_output` is the pre-RoPE latent, for the index-K branch's `wk`. +struct C2Params { + const float* __restrict__ kv_input; // [num_tokens, 2 * kHeadDim] fp32 + /// `CompressStatePool`'s flat `KVAndScore` buffer, `[size, 2 * kHeadDim]` + /// fp32 with kv in the low half and score in the high half. A request's + /// pending pair lives at `req * ring_size + pos % ring_size`. + float* __restrict__ kv_state; + bf16_t* __restrict__ kv_output; // [num_tokens, kHeadDim] bf16, pre-RoPE + const bf16_t* __restrict__ norm_weight; // [kHeadDim] bf16 + const float* __restrict__ freqs_cis; // [max_pos, kRopeDim] fp32, real/imag interleaved + const void* __restrict__ positions; // [num_tokens] PosT + const int64_t* __restrict__ req; // [num_tokens], req_pool_idx per token + const void* __restrict__ raw_out_loc; // [num_tokens] LocT, the FULL slot; 0 marks a padded row + uint8_t* __restrict__ kvcache; // [npages, kPageBytes] uint8 + /// Positions per request slot in the pair-state ring. + uint32_t ring_size; + float eps; +}; + +/// Elements per thread; 256 threads per token was measured fastest on B200 decode batches. +/// The launcher and launch bounds share this value. At head_dim 512, the nope/rope split +/// is (512 - 64) / 2 = 224 threads, keeping the fp8 amax reduction's full-warp mask valid. +constexpr uint32_t kC2VecSize = 2; + +/// \brief grid = num_tokens, block = kHeadDim / kC2VecSize. +/// +/// An odd position completes a group with its even predecessor; an even one +/// parks itself in the state. +/// +/// Target-verify runs that same schedule with `draft_len` consecutive positions +/// per request instead of one, so a row's partner is usually the row before it +/// in `kv_input` rather than the ring. That is the whole difference, and a 2D +/// grid answers it without arithmetic: `blockIdx.x` is the position inside the +/// block, `blockIdx.y` the request. +/// +/// The three reductions below have three different widths and are not +/// interchangeable: the RMSNorm statistic spans the row, an fp8 store scale +/// spans 64 elements, an fp4 block spans 16. All asserted. +/// +/// kLayout is the cache's page format. V4 (584 B/token) and V41 (528 B/token, +/// fp8 with per-32 scales) store the fake-quantized value; V41_FP4 (288 B/token) +/// stores the e2m1 codes and their e4m3 scales directly, so the fp4 rounding +/// happens once and no fp8 rounding follows it. +template < + bool kStore, + bool kVerify, + int64_t kHeadDim, + int64_t kRopeDim, + int32_t kPageBits, + typename PosT, + typename LocT, + deepseek_v4::KVLayout kLayout, + bool kUsePDL> +__global__ __launch_bounds__(kHeadDim / kC2VecSize) void flash_c2_decode_kernel(const C2Params params) { + using namespace device; + using deepseek_v4::KVLayout; + using deepseek_v4::fp8::cast_to_ue8m0; + using deepseek_v4::fp8::inv_scale_ue8m0; + using deepseek_v4::fp8::pack_fp8; + + constexpr uint32_t kVecSize = kC2VecSize; + constexpr uint32_t kCTASize = kHeadDim / kVecSize; + constexpr int64_t kStride = kHeadDim * 2; + /// Threads covering the fp8 nope part; the rest carry the bf16 RoPE tail. + constexpr uint32_t kNopeThreads = (kHeadDim - kRopeDim) / kVecSize; + constexpr uint32_t kFp8Lanes = 64 / kVecSize; + constexpr uint32_t kFp4Lanes = deepseek_v4::fp4::kCompressedKVBlockSize / kVecSize; + using Paged = deepseek_v4::PagedKV; + + static_assert(kHeadDim == (kVecSize * kCTASize)); + static_assert(kCTASize % kWarpThreads == 0); + static_assert(kNopeThreads % kFp8Lanes == 0, "the nope part must end on an fp8 scale block"); + static_assert(kWarpThreads % kFp8Lanes == 0 && kWarpThreads % kFp4Lanes == 0); + static_assert(kHeadDim == 512 && kRopeDim == 64, "the FlashMLA layouts require (512, 64)"); + using fp32_vec_t = AlignedVector; + using bf16_vec_t = AlignedVector; + + const auto tx = threadIdx.x; + // Verify gives each request a CTA column; decode a flat grid of one row each. + const auto row = kVerify ? blockIdx.y * gridDim.x + blockIdx.x : blockIdx.x; + // Slots fit in int32 whatever width the scheduler hands them in. + const auto raw_out_loc = static_cast(static_cast(params.raw_out_loc)[row]); + const auto pos = static_cast(params.positions)[row]; + // A completing row reads the slot left by `pos - 1`; + // a pending row writes its own slot, so reads and writes stay disjoint. + const auto rid = params.req[row]; + PDLWaitPrimary(); + + fp32_vec_t kv_new, score_new; + kv_new.load(params.kv_input + row * kStride, tx); + score_new.load(params.kv_input + row * kStride, tx + kCTASize); + + const auto ring = static_cast(rid) * params.ring_size; + const auto read_row = ring + (pos - 1 + params.ring_size) % params.ring_size; + const auto write_row = ring + pos % params.ring_size; + + fp32_vec_t kv_old, score_old; + const float* partner = params.kv_state + read_row * kStride; + if constexpr (kVerify) { + if (blockIdx.x != 0) partner = params.kv_input + static_cast(row - 1) * kStride; + } + kv_old.load(partner, tx); + score_old.load(partner, tx + kCTASize); + + if ((pos & 1) == 0) { + // padded case + if (raw_out_loc == 0) return PDLTriggerSecondary(); + kv_new.store(params.kv_state + write_row * kStride, tx); + score_new.store(params.kv_state + write_row * kStride, tx + kCTASize); + return PDLTriggerSecondary(); + } + + constexpr uint32_t kNumWarps = kCTASize / kWarpThreads; + __shared__ float s_warp_sum[kNumWarps]; + fp32_vec_t staged, freq; + bf16_vec_t weight, out; + weight.load(params.norm_weight, tx); + if constexpr (kStore) { + if (tx >= kNopeThreads) freq.load(params.freqs_cis + (pos - 1) * kRopeDim, tx - kNopeThreads); + } + + // With two scores `exp(-|s0 - s1|)` is the whole softmax: one exp, argument + // always <= 0, so no max-subtraction pass and no overflow. +#pragma unroll + for (uint32_t i = 0; i < kVecSize; ++i) { + const auto delta = score_old[i] - score_new[i]; + const auto scale = expf(-fabsf(delta)); + const auto scale_0 = delta > 0 ? 1.0f : scale; + const auto scale_1 = delta > 0 ? scale : 1.0f; + staged[i] = (kv_old[i] * scale_0 + kv_new[i] * scale_1) / (1.0f + scale); + } + + // `finish` casts to bf16 before the norm, so the sum of squares has to see + // the rounded values. + float local_sqrsum = 0.0f; +#pragma unroll + for (uint32_t i = 0; i < kVecSize / 2; ++i) { + const auto packed = fp32x2_t{staged[i * 2 + 0], staged[i * 2 + 1]}; + const auto [x, y] = cast(cast(packed)); + local_sqrsum += x * x; + local_sqrsum += y * y; + staged[i * 2 + 0] = x; + staged[i * 2 + 1] = y; + } + const auto warp_sum = warp::reduce_sum(local_sqrsum); + s_warp_sum[tx / kWarpThreads] = warp_sum; + __syncthreads(); + + float sqrsum = 0.0f; +#pragma unroll + for (uint32_t i = 0; i < kNumWarps; ++i) { + sqrsum += s_warp_sum[i]; + } + constexpr float kInvScale = 1.0f / static_cast(kHeadDim); + const auto norm_factor = math::rsqrt(sqrsum * kInvScale + params.eps); + +#pragma unroll + for (uint32_t i = 0; i < kVecSize / 2; ++i) { + const auto [wx, wy] = cast(weight[i]); + const auto x = staged[i * 2 + 0] * norm_factor * wx; + const auto y = staged[i * 2 + 1] * norm_factor * wy; + out[i] = cast(fp32x2_t{x, y}); + } + // The pre-RoPE latent, for the index-K branch's `wk` projection. Published + // before the trigger because that GEMM is the successor that reads it. + out.store(params.kv_output, static_cast(row) * kCTASize + tx); + PDLTriggerSecondary(); + + if constexpr (kStore) { + // ---- main-KV branch: RoPE tail, fp4 fake-quant, 584-byte store ---- + // Match finish()'s bf16 rounding before RoPE. +#pragma unroll + for (uint32_t i = 0; i < kVecSize / 2; ++i) { + const auto [x, y] = cast(out[i]); + staged[i * 2 + 0] = x; + staged[i * 2 + 1] = y; + } + + if (tx >= kNopeThreads) { + // Match rope_tail()'s bf16 rounding before fake quantization. + // Only odd positions reach here; the latent represents `pos - 1`. + freq.load(params.freqs_cis + (pos - 1) * kRopeDim, tx - kNopeThreads); +#pragma unroll + for (uint32_t i = 0; i < kVecSize / 2; ++i) { + const auto x_real = staged[i * 2 + 0]; + const auto x_imag = staged[i * 2 + 1]; + const auto f_real = x_real * freq[i * 2 + 0] - x_imag * freq[i * 2 + 1]; + const auto f_imag = x_real * freq[i * 2 + 1] + x_imag * freq[i * 2 + 0]; + const auto rotated = cast(fp32x2_t{f_real, f_imag}); + const auto [r0, r1] = cast(rotated); + staged[i * 2 + 0] = r0; + staged[i * 2 + 1] = r1; + } + } + + if constexpr (kLayout == KVLayout::V41_FP4) { + // The fp4 cache takes the rotated bf16 value as is: its row quantizer is the + // fake quantization, minus the dequantization. + if (raw_out_loc == 0) return; + const int32_t out_loc = raw_out_loc >> 1; + const auto kv_row = Paged::row(params.kvcache, out_loc); + return deepseek_v4::v41::store_row(kv_row.data, kv_row.scale, tx, staged); + } + + // FP4/E4M3 fake-quant over 16 elements, i.e. kFp4Lanes threads. + { + float amax = fabsf(staged[0]); +#pragma unroll + for (uint32_t i = 1; i < kVecSize; ++i) { + amax = fmaxf(amax, fabsf(staged[i])); + } + amax = warp::reduce_max(amax); + const auto scale = deepseek_v4::fp4::compressed_kv_scale(amax); +#pragma unroll + for (uint32_t i = 0; i < kVecSize / 2; ++i) { + const auto [x, y] = + deepseek_v4::fp4::fake_quant_compressed_kv_x2({staged[i * 2 + 0], staged[i * 2 + 1]}, scale); + staged[i * 2 + 0] = x; + staged[i * 2 + 1] = y; + } + } + + // padded case + if (raw_out_loc == 0) return; + // `raw_out_loc / ratio`; ratio 2 makes it a shift. + const int32_t out_loc = raw_out_loc >> 1; + const auto kv_row = Paged::row(params.kvcache, out_loc); + + if constexpr (kLayout == KVLayout::V41) { + // fp8 with one ue8m0 scale per 32 elements over the whole row, RoPE included. + return deepseek_v4::v41::store_row(kv_row.data, kv_row.scale, tx, staged); + } + + const auto value_ptr = kv_row.data; + + if (tx >= kNopeThreads) { + bf16_vec_t rope_out; +#pragma unroll + for (uint32_t i = 0; i < kVecSize / 2; ++i) { + rope_out[i] = cast(fp32x2_t{staged[i * 2 + 0], staged[i * 2 + 1]}); + } + rope_out.store(value_ptr + (kHeadDim - kRopeDim), tx - kNopeThreads); + } else { + // fp8 e4m3 with one ue8m0 scale per 64 elements. + auto abs_max = fabsf(staged[0]); +#pragma unroll + for (uint32_t i = 1; i < kVecSize; ++i) { + abs_max = fmaxf(abs_max, fabsf(staged[i])); + } + abs_max = warp::reduce_max(abs_max); + const auto scale_ue8m0 = cast_to_ue8m0(fmaxf(1e-4f, abs_max) / math::FP8_E4M3_MAX); + const auto inv_scale = inv_scale_ue8m0(scale_ue8m0); +#pragma unroll + for (uint32_t i = 0; i < kVecSize / 2; ++i) { + reinterpret_cast(value_ptr)[tx * (kVecSize / 2) + i] = + pack_fp8(staged[i * 2 + 0] * inv_scale, staged[i * 2 + 1] * inv_scale); + } + kv_row.scale[tx / kFp8Lanes] = scale_ue8m0; + } + } +} + +template +struct FlashC2DecodeKernel { + static constexpr uint32_t kBlockSize = kHeadDim / kC2VecSize; + static constexpr int32_t kPageBits = std::bit_width(kPageSize) - 1; + static constexpr int64_t kPageBytes = deepseek_v4::kv_page_bytes(kPageSize); + static_assert(kLayout != deepseek_v4::KVLayout::V4 || kPageBytes == host::div_ceil(584ll * kPageSize, 576) * 576); + template + static constexpr auto kernel = + flash_c2_decode_kernel; + + /// \brief The (`positions`, `raw_out_loc`) dtype pair, resolved at run time. + template + static auto select(const bool pos_i32, const bool loc_i32) { + if (pos_i32) return loc_i32 ? kernel : kernel; + return loc_i32 ? kernel : kernel; + } + + // The sum of squares is reduced through a fixed-size shared array, so the CTA + // has to be a whole number of warps. + static_assert(kHeadDim % (4 * device::kWarpThreads) == 0, "head_dim must be a multiple of 128"); + static_assert(std::has_single_bit(kPageSize), "the page/slot split needs a power-of-two page"); + + /// \brief Pool + norm only. The main-KV write stays with the caller. + static void run_decode( + const tvm::ffi::TensorView kv_input, + const tvm::ffi::TensorView kv_state, + const tvm::ffi::TensorView kv_output, + const tvm::ffi::TensorView norm_weight, + const tvm::ffi::TensorView positions, + const tvm::ffi::TensorView req, + const tvm::ffi::TensorView raw_out_loc, + const float eps, + const int64_t ring_size) { + launch( + kv_input, + kv_state, + kv_output, + norm_weight, + positions, + req, + raw_out_loc, + eps, + ring_size, + std::nullopt, + std::nullopt, + /*draft_len=*/1); + } + + /// \brief `run_decode_fusion` for a target-verify block. + /// + /// `draft_len` consecutive positions per request, request-major, which the + /// grid reproduces as `draft_len x batch`. + static void run_decode_fusion( + const tvm::ffi::TensorView kv_input, + const tvm::ffi::TensorView kv_state, + const tvm::ffi::TensorView kv_output, + const tvm::ffi::TensorView norm_weight, + const tvm::ffi::TensorView positions, + const tvm::ffi::TensorView req, + const tvm::ffi::TensorView raw_out_loc, + const float eps, + const tvm::ffi::TensorView freqs_cis, + const tvm::ffi::TensorView kvcache, + const int64_t ring_size, + const int64_t draft_len) { + launch( + kv_input, + kv_state, + kv_output, + norm_weight, + positions, + req, + raw_out_loc, + eps, + ring_size, + freqs_cis, + kvcache, + draft_len); + } + + private: + using MaybeTensor = std::optional; + + static void launch( + const tvm::ffi::TensorView kv_input, + const tvm::ffi::TensorView kv_state, + const tvm::ffi::TensorView kv_output, + const tvm::ffi::TensorView norm_weight, + const tvm::ffi::TensorView positions, + const tvm::ffi::TensorView req, + const tvm::ffi::TensorView raw_out_loc, + const float eps, + const int64_t ring_size, + const MaybeTensor freqs_cis, + const MaybeTensor kvcache, + const int64_t draft_len) { + using namespace host; + + auto N = SymbolicSize{"num_tokens"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + + TensorMatcher({N, kHeadDim * 2}).with_dtype().with_device(device_).verify(kv_input); + TensorMatcher({-1, kHeadDim * 2}).with_dtype().with_device(device_).verify(kv_state); + // Only rows that complete a group are written. + TensorMatcher({N, kHeadDim}).with_dtype().with_device(device_).verify(kv_output); + TensorMatcher({kHeadDim}).with_dtype().with_device(device_).verify(norm_weight); + // Metadata retains its original dtypes: the scheduler uses int64 locations, + // while callers may also supply int32 locations and positions. + auto pos_dtype = SymbolicDType{}; + auto loc_dtype = SymbolicDType{}; + TensorMatcher({N}).with_dtype(pos_dtype).with_device(device_).verify(positions); + TensorMatcher({N}).with_dtype().with_device(device_).verify(req); + TensorMatcher({N}).with_dtype(loc_dtype).with_device(device_).verify(raw_out_loc); + + const auto store = freqs_cis.has_value(); + if (store) { + // Real/imag interleaved, so the trailing dim is kRopeDim, not kRopeDim / 2. + TensorMatcher({-1, kRopeDim}).with_dtype().with_device(device_).verify(*freqs_cis); + // The pool allocates the buffer as uint8 and hands it out viewed as its + // fp8 dtype (`get_extra_key_buffer`); both are one byte per element. + TensorMatcher({-1, kPageBytes}).with_dtype().with_device(device_).verify(*kvcache); + } + + const auto num_tokens = static_cast(N.unwrap()); + if (num_tokens == 0) return; + const auto is_verify = draft_len > 1; + CHECK_HOST(ring_size > 0 && draft_len >= 1); + CHECK_HOST(!is_verify || num_tokens % draft_len == 0); + CHECK_HOST(!is_verify || ring_size > draft_len) + << "the pair-state ring (" << ring_size << ") must be wider than the draft length (" << draft_len << ")"; + const auto params = C2Params{ + .kv_input = static_cast(kv_input.data_ptr()), + .kv_state = static_cast(kv_state.data_ptr()), + .kv_output = static_cast(kv_output.data_ptr()), + .norm_weight = static_cast(norm_weight.data_ptr()), + .freqs_cis = store ? static_cast(freqs_cis->data_ptr()) : nullptr, + .positions = positions.data_ptr(), + .req = static_cast(req.data_ptr()), + .raw_out_loc = raw_out_loc.data_ptr(), + .kvcache = store ? static_cast(kvcache->data_ptr()) : nullptr, + .ring_size = static_cast(ring_size), + .eps = eps, + }; + // `LaunchKernel` is move-only, so each arm builds its own. + const auto pos_i32 = pos_dtype.is_type(); + const auto loc_i32 = loc_dtype.is_type(); + if (is_verify) { + const auto block = static_cast(draft_len); + const auto k = select(pos_i32, loc_i32); + LaunchKernel(dim3{block, num_tokens / block}, kBlockSize, device_.unwrap()) // + .enable_pdl(kUsePDL)(k, params); + } else if (store) { + const auto k = select(pos_i32, loc_i32); + LaunchKernel(num_tokens, kBlockSize, device_.unwrap()) // + .enable_pdl(kUsePDL)(k, params); + } else { + const auto k = select(pos_i32, loc_i32); + LaunchKernel(num_tokens, kBlockSize, device_.unwrap()) // + .enable_pdl(kUsePDL)(k, params); + } + } +}; + +// ensure that C++ wrapper can work +using enum deepseek_v4::KVLayout; + +} // namespace sglang diff --git a/python/sglang/kernels/jit/include/sgl_kernel/deepseek_v4/fp4_utils.cuh b/python/sglang/kernels/jit/include/sgl_kernel/deepseek_v4/fp4_utils.cuh new file mode 100644 index 000000000000..cb444c17d995 --- /dev/null +++ b/python/sglang/kernels/jit/include/sgl_kernel/deepseek_v4/fp4_utils.cuh @@ -0,0 +1,64 @@ +#pragma once + +#include +#include + +#include + +#ifndef USE_ROCM +#include +#endif + +// FP4 (e2m1) helpers: per-32 UE8M0 for the indexer, per-16 E4M3 for compressed KV. + +namespace sglang { + +namespace deepseek_v4::fp4 { + +/// Largest finite e2m1 value. +constexpr float kMax = 6.0f; +/// `6 * 2^-126`, the amax floor `torch_quant.fake_quant_fp4` clamps to. +constexpr float kAmaxFloor = 6.0f * 1.1754943508222875e-38f; +/// Elements sharing one ue8m0 scale. +constexpr uint32_t kBlockSize = 32; +/// Compressed-KV elements sharing one E4M3 scale. +constexpr uint32_t kCompressedKVBlockSize = 16; + +/// \brief Round amax / 6 to a positive finite E4M3 scale, ties to even. +SGL_DEVICE float compressed_kv_scale(float amax) { + const auto raw = fminf(fmaxf(amax * (1.0f / kMax), 0x1p-9f), 448.0f); + return static_cast(__nv_fp8_e4m3(raw)); +} + +/// \brief Quantize compressed KV with its E4M3 scale and return dequantized values. +SGL_DEVICE fp32x2_t fake_quant_compressed_kv_x2(fp32x2_t x, float scale) { + const fp32x2_t scaled{__fdiv_rn(x.x, scale) + 0.0f, __fdiv_rn(x.y, scale) + 0.0f}; + const auto code = __nv_cvt_float2_to_fp4x2(scaled, __NV_E2M1, cudaRoundNearest); + const auto grid = device::cast(fp16x2_t{__nv_cvt_fp4x2_to_halfraw2(code, __NV_E2M1)}); + return {grid.x * scale, grid.y * scale}; +} + +/// \brief Per-block ue8m0 scale and its reciprocal, from the block's absmax. +/// +/// Both come out of one biased exponent, so the reciprocal costs a subtract +/// rather than a division. +SGL_DEVICE fp32x2_t block_scale(float amax) { + const auto exponent = fp8::cast_to_ue8m0(fmaxf(amax, kAmaxFloor) * (1.0f / kMax)); + return {__uint_as_float(static_cast(exponent) << 23), fp8::inv_scale_ue8m0(exponent)}; +} + +/// \brief Round a pair onto the e2m1 grid and back, through `scale`. +/// +/// `cvt.rn.satfinite.e2m1x2.f32` rounds to nearest even and saturates to +-6; +/// every e2m1 value is exact in fp16. Adding `0.0f` during scaling clears negative +/// zero to match `torch.sign(0) == 0` in `torch_quant.round_fp4`. +SGL_DEVICE fp32x2_t fake_quant_x2(fp32x2_t x, float scale, float inv_scale) { + const fp32x2_t scaled{__fmaf_rn(x.x, inv_scale, 0.0f), __fmaf_rn(x.y, inv_scale, 0.0f)}; + const auto code = __nv_cvt_float2_to_fp4x2(scaled, __NV_E2M1, cudaRoundNearest); + const auto grid = device::cast(fp16x2_t{__nv_cvt_fp4x2_to_halfraw2(code, __NV_E2M1)}); + return {grid.x * scale, grid.y * scale}; +} + +} // namespace deepseek_v4::fp4 + +} // namespace sglang diff --git a/python/sglang/kernels/jit/include/sgl_kernel/deepseek_v4/kv_layout.cuh b/python/sglang/kernels/jit/include/sgl_kernel/deepseek_v4/kv_layout.cuh new file mode 100644 index 000000000000..cd34309d5be2 --- /dev/null +++ b/python/sglang/kernels/jit/include/sgl_kernel/deepseek_v4/kv_layout.cuh @@ -0,0 +1,220 @@ +#pragma once + +#include +#include +#include +#include + +#include + +#include +#ifndef USE_ROCM +#include +#include +#endif + +// Paged fp8 / fp4 KV cache layouts read by the d_qk = 512 sparse MLA decode +// kernels. A page block is `page_size` data rows followed by `page_size` scale +// rows, so the scale region starts at byte `page_size * kDataBytes`: +// +// V4 584 B/token: 448 e4m3 + 64 bf16 (RoPE) data, 7 ue8m0 scales + 1 pad, +// one scale per 64 e4m3 values. +// V41 528 B/token: 512 e4m3 data (the RoPE dims are quantized too), +// 16 ue8m0 scales, one per 32 values. +// V41_FP4 288 B/token: 512 e2m1 data packed two per byte (even index in the +// low nibble), 32 e4m3 scales, one per 16 values. +// +// The reader requires the rows of a page to be contiguous and the page stride +// to be a multiple of kPageAlign (its TMA row stride), which is what +// kv_page_bytes pads to. The pure-torch reference of the V4.1 quantizers is +// `sglang.srt.layers.attention.dsv4.torch_quant`. + +namespace sglang { + +namespace deepseek_v4 { + +enum class KVLayout : int32_t { V4 = 0, V41 = 1, V41_FP4 = 2 }; + +template +struct KVLayoutTraits; + +template <> +struct KVLayoutTraits { + static constexpr int64_t kDataBytes = 576; + static constexpr int64_t kScaleBytes = 8; + static constexpr int64_t kTileSize = 64; + static constexpr int64_t kPageAlign = 576; + static constexpr int64_t kBytesPerToken = kDataBytes + kScaleBytes; +}; + +template <> +struct KVLayoutTraits { + static constexpr int64_t kDataBytes = 512; + static constexpr int64_t kScaleBytes = 16; + static constexpr int64_t kTileSize = 32; + static constexpr int64_t kPageAlign = 512; + static constexpr int64_t kBytesPerToken = kDataBytes + kScaleBytes; +}; + +template <> +struct KVLayoutTraits { + static constexpr int64_t kDataBytes = 256; + static constexpr int64_t kScaleBytes = 32; + static constexpr int64_t kTileSize = 16; + static constexpr int64_t kPageAlign = 256; + static constexpr int64_t kBytesPerToken = kDataBytes + kScaleBytes; +}; + +/// Bytes of one page block: `page_size` tokens, padded up to the reader's row stride. +template +constexpr int64_t kv_page_bytes(int64_t page_size) { + using Traits = KVLayoutTraits; + return (page_size * Traits::kBytesPerToken + Traits::kPageAlign - 1) / Traits::kPageAlign * Traits::kPageAlign; +} + +/// Addressing of a paged cache in one layout: a page is `1 << kPageBits` data rows followed +/// by as many scale rows, padded to the layout's kPageAlign. Every member is a compile-time +/// constant or a constant shift / multiply of the token index, so it folds to the same code +/// as the hand-written arithmetic; the index keeps the caller's type `LocT`. +template +struct PagedKV { + using Traits = KVLayoutTraits; + static constexpr int64_t kPageSize = int64_t{1} << kPageBits; + static constexpr int64_t kPageBytes = kv_page_bytes(kPageSize); + /// Byte offset of the scale rows inside a page. + static constexpr int64_t kScaleBase = Traits::kDataBytes << kPageBits; + + template + static constexpr LocT page_of(LocT loc) { + return loc >> kPageBits; + } + template + static constexpr LocT slot_of(LocT loc) { + return loc & (static_cast(kPageSize) - 1); + } + /// Byte offsets of token `loc`'s data row and scale row from the cache base. + template + static constexpr int64_t data_offset(LocT loc) { + return page_of(loc) * kPageBytes + slot_of(loc) * Traits::kDataBytes; + } + template + static constexpr int64_t scale_offset(LocT loc) { + return page_of(loc) * kPageBytes + kScaleBase + slot_of(loc) * Traits::kScaleBytes; + } + + struct Row { + uint8_t* data; + uint8_t* scale; + }; + /// The data row and scale row of token `loc`. + template + SGL_DEVICE static Row row(uint8_t* cache, LocT loc) { + uint8_t* page = cache + page_of(loc) * kPageBytes; + return {page + slot_of(loc) * Traits::kDataBytes, page + kScaleBase + slot_of(loc) * Traits::kScaleBytes}; + } +}; + +#ifndef USE_ROCM + +namespace v41 { + +/// The row helpers below quantize one 512-wide token spread over `512 / kVecSize` threads, +/// thread `tx` holding elements `[kVecSize * tx, kVecSize * (tx + 1))` as fp32. They are +/// warp-collective (sub-warp reductions under the full mask), so every thread of the token +/// must call them together. `kVecSize` is even, a power of two and at most the tile size. +/// NaN / inf inputs are not handled. + +/// Per-thread |max| over the vector. +template +SGL_DEVICE float vec_amax(const float (&v)[kVecSize]) { + float amax = fabsf(v[0]); +#pragma unroll + for (uint32_t i = 1; i < kVecSize; ++i) { + amax = fmaxf(amax, fabsf(v[i])); + } + return amax; +} + +/// V4.1 fp8 row: one ue8m0 scale per 32-element tile. `data_row` is the token's 512 B, +/// `scale_row` its 16 scale bytes. Scale `2^ceil(log2(max(amax / 448, 1e-4)))` stored as the +/// ue8m0 byte, payload `e4m3(x / scale)` rounded to nearest even (|x / scale| <= 448, so it +/// never saturates). +template +SGL_DEVICE void store_row_fp8(uint8_t* data_row, uint8_t* scale_row, uint32_t tx, const float (&v)[kVecSize]) { + using namespace device; + constexpr uint32_t kTileLanes = KVLayoutTraits::kTileSize / kVecSize; + static_assert(kVecSize % 2 == 0 && kTileLanes >= 1 && (kTileLanes & (kTileLanes - 1)) == 0); + + const float amax = warp::reduce_max(vec_amax(v)); + // ceil(log2(max(amax / 448, 1e-4))) straight from the bits of amax: 448 = 1.75 * 2^8, so + // the quotient's exponent is amax's minus 8, plus one when amax's mantissa exceeds 1.75 + // (the quotient is then just above a power of two), floored at 2^-13, the smallest power + // of two >= 1e-4. Exact for every finite amax, without the reference's fp32 division. + const uint32_t bits = __float_as_uint(amax); + const int32_t exponent = + max(static_cast(bits >> 23) - 8 + static_cast((bits & 0x7FFFFFu) > 0x600000u), 114); + // The scale is a power of two, so the multiply by its reciprocal is the exact quotient. + const float inv_scale = fp8::inv_scale_ue8m0(exponent); + AlignedVector out; +#pragma unroll + for (uint32_t i = 0; i < kVecSize / 2; ++i) { + // `cvt.rn.satfinite.e4m3x2` directly: |x / scale| <= 448 needs no clamp. + out[i] = fp8x2_e4m3_t{fp32x2_t{v[2 * i] * inv_scale, v[2 * i + 1] * inv_scale}}; + } + out.store(data_row, tx); + // Every lane of the tile holds the exponent; they all store the same byte. + scale_row[tx / kTileLanes] = static_cast(exponent); +} + +/// V4.1 fp4 row: one e4m3 scale per 16-element tile. `data_row` is the token's 256 B, +/// `scale_row` its 32 scale bytes. Scale `e4m3(clamp(amax / 6, 2^-9, 448))` rounded to +/// nearest even; codes `cvt.rn.satfinite.e2m1x2` of `x / scale` (ties to even, saturating +/// at 6, the sign kept for a value that rounds to zero), the even element in the low nibble. +template +SGL_DEVICE void store_row_fp4(uint8_t* data_row, uint8_t* scale_row, uint32_t tx, const float (&v)[kVecSize]) { + using namespace device; + constexpr uint32_t kTileLanes = KVLayoutTraits::kTileSize / kVecSize; + static_assert(kVecSize % 2 == 0 && kTileLanes >= 1 && (kTileLanes & (kTileLanes - 1)) == 0); + + const float amax = warp::reduce_max(vec_amax(v)); + const __nv_fp8_e4m3 scale_e4m3{fminf(fmaxf(__fdiv_rn(amax, 6.0f), 0x1p-9f), 448.0f)}; + const float scale = static_cast(scale_e4m3); + AlignedVector out; +#pragma unroll + for (uint32_t i = 0; i < kVecSize / 2; ++i) { + // IEEE division by the rounded scale, as the reference divides; a reciprocal multiply + // could land on the other side of an e2m1 tie. + out[i] = static_cast(__nv_cvt_float2_to_fp4x2( + fp32x2_t{__fdiv_rn(v[2 * i], scale), __fdiv_rn(v[2 * i + 1], scale)}, __NV_E2M1, cudaRoundNearest)); + } + out.store(data_row, tx); + // Every lane of the tile holds the scale; they all store the same byte. + scale_row[tx / kTileLanes] = scale_e4m3.__x; +} + +/// Dispatch on the layout for a 512-wide row held as `kVecSize` consecutive fp32 per thread. +/// V4 has no row helper here: its writers keep their nope / RoPE split code. +template +SGL_DEVICE void store_row(uint8_t* data_row, uint8_t* scale_row, uint32_t tx, const float (&v)[kVecSize]) { + static_assert(kLayout != KVLayout::V4, "V4 rows are written by the caller"); + if constexpr (kLayout == KVLayout::V41) { + store_row_fp8(data_row, scale_row, tx, v); + } else { + store_row_fp4(data_row, scale_row, tx, v); + } +} + +/// Same, for a row kept in an `AlignedVector`. +template +SGL_DEVICE void +store_row(uint8_t* data_row, uint8_t* scale_row, uint32_t tx, const device::AlignedVector& v) { + store_row(data_row, scale_row, tx, *reinterpret_cast(v.data())); +} + +} // namespace v41 + +#endif // USE_ROCM + +} // namespace deepseek_v4 + +} // namespace sglang diff --git a/python/sglang/kernels/ops/attention/dsv4/c1.py b/python/sglang/kernels/ops/attention/dsv4/c1.py new file mode 100644 index 000000000000..2feb12318468 --- /dev/null +++ b/python/sglang/kernels/ops/attention/dsv4/c1.py @@ -0,0 +1,103 @@ +"""Fused ratio-1 decode RMSNorm, RoPE, fp4 fake-quant and FlashMLA cache write. + +The input is bf16 with no pooling. RoPE uses the token's own position, and the +compressed slot equals the FULL slot; out_loc == 0 marks graph padding. +The pre-RoPE latent is also returned for the index-key projection. +""" + +from typing import Optional, Union + +import torch + +from sglang.kernels.jit.utils import ( + cache_once, + is_arch_support_pdl, + load_jit, + make_cpp_args, +) + +from .kv_layout import KVLayout +from .utils import make_name + + +@cache_once +def _jit_c1_module(head_dim: int, rope_dim: int, page_size: int, layout: KVLayout): + args = make_cpp_args( + head_dim, + rope_dim, + page_size, + layout.cpp_name, + is_arch_support_pdl(), + ) + return load_jit( + make_name("c1"), + *args, + cuda_files=["deepseek_v4/c1.cuh"], + cuda_wrappers=[ + ("decode_fusion", f"FlashC1DecodeKernel<{args}>::run_decode_fusion"), + ], + ) + + +def c1_decode_norm_rope_store( + kv_input: torch.Tensor, + norm_weight: torch.Tensor, + positions: torch.Tensor, + out_loc: torch.Tensor, + eps: float, + freqs_cis: torch.Tensor, + k_cache: torch.Tensor, + *, + page_size: int, + layout: Union[KVLayout, str] = KVLayout.V4, + out: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """RMSNorm ``kv_input`` and write the main KV slot, in a single launch. + + ``out`` contains the pre-RoPE latent for the index-K branch's ``wk`` projection; + the main KV cache receives the rotated and quantized value. + + :param kv_input: ``[num_tokens, head_dim]`` bf16 -- ``compressor.project(x)`` + at ratio 1, i.e. the raw ``wkv`` output. + :param norm_weight: ``[head_dim]`` bf16. ``DeepseekV41Compressor.norm`` holds + its weight in the model dtype and the multiply is done in + fp32 by promoting it, exactly as the module does. + :param positions: ``[num_tokens]`` int32 or int64, the token's position. + Indexed into ``freqs_cis`` as-is: at ratio 1 the latent + stands for the token itself, so there is no ``- 1`` and no + gather launch. + :param out_loc: ``[num_tokens]`` int32 or int64 ``c1_out_loc``, which at ratio 1 + equals ``raw_out_loc`` (the scheduler's int64 ``out_cache_loc``). + ``0`` marks a padded graph row: it computes + and publishes its latent, which the caller discards, but + writes nothing to the cache. + :param eps: RMSNorm epsilon. + :param freqs_cis: ``[max_pos, rope_dim]`` fp32, real/imag interleaved -- + ``torch.view_as_real(freqs).flatten(-2)``. + :param k_cache: the compressed KV pool buffer for this layer. + :param page_size: slots per page of that pool (``page_size // ratio``, i.e. + the FULL page size at ratio 1). + :param layout: the pool's :class:`KVLayout`. The fp8 layouts (``V4``, + ``V41``) store the fp4 fake-quantized value; ``V41_FP4`` + stores the e2m1 codes themselves, rounding once. + :param out: ``[num_tokens, head_dim]`` bf16 destination for the pre-RoPE + latent. Pass a persistent buffer under CUDA graphs. + :return: ``out``, the pre-RoPE post-norm latent. + """ + num_tokens, head_dim = kv_input.shape + if out is None: + out = kv_input.new_empty((num_tokens, head_dim)) + + layout = KVLayout.parse(layout) + module = _jit_c1_module(head_dim, freqs_cis.shape[-1], page_size, layout) + module.decode_fusion( + kv_input, + out, + norm_weight, + freqs_cis, + positions, + out_loc, + k_cache, + float(eps), + ) + return out diff --git a/python/sglang/kernels/ops/attention/dsv4/c2.py b/python/sglang/kernels/ops/attention/dsv4/c2.py new file mode 100644 index 000000000000..f3fdb5acebc1 --- /dev/null +++ b/python/sglang/kernels/ops/attention/dsv4/c2.py @@ -0,0 +1,179 @@ +"""Fused ratio-2 decode pair-pooling, RMSNorm and optional main-KV write. + +Closed-form softmax and FMA contraction can differ from torch by fp32 ulps; +bf16 rounding boundaries can preserve those differences. The pooling tests +use a tolerance, while state updates and stores from a given latent are bitwise. +Positions and state-ring indices describe the per-request decode schedule. +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Final, Optional, Union + +import torch + +from sglang.kernels.jit.utils import ( + cache_once, + is_arch_support_pdl, + load_jit, + make_cpp_args, +) + +from .kv_layout import KVLayout +from .utils import make_name + +if TYPE_CHECKING: + from tvm_ffi.module import Module + + +@cache_once +def _jit_c2_module( + head_dim: int, + rope_dim: int, + page_size: int, + layout: KVLayout, +) -> Module: + # rope_dim / page_size / layout only shape the store half; the norm-only + # entry point ignores them. + args = make_cpp_args( + head_dim, + rope_dim, + page_size, + layout.cpp_name, + is_arch_support_pdl(), + ) + return load_jit( + make_name("c2"), + *args, + cuda_files=["deepseek_v4/c2.cuh"], + cuda_wrappers=[ + ("decode", f"FlashC2DecodeKernel<{args}>::run_decode"), + ("decode_fusion", f"FlashC2DecodeKernel<{args}>::run_decode_fusion"), + ], + ) + + +def c2_decode_norm( + kv_input: torch.Tensor, + kv_state: torch.Tensor, + norm_weight: torch.Tensor, + positions: torch.Tensor, + req: torch.Tensor, + raw_out_loc: torch.Tensor, + eps: float, + *, + ring_size: int, + out: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """Pair-pool ``kv_input`` against ``kv_state`` and RMSNorm the result. + + :param kv_input: ``[num_tokens, 2 * head_dim]`` fp32, ``| kv | score |``. + :param kv_state: ``CompressStatePool``'s flat ``KVAndScore`` buffer, + ``[size, 2 * head_dim]`` fp32, same ``| kv | score |`` + layout. A request's pending pair lives at + ``req * ring_size + pos % ring_size``, so a completing row + reads what ``pos - 1`` left and a pending one writes its + own slot -- read and write never touch the same row. + :param ring_size: ``CompressStatePool.ring_size``, positions per request. + :param norm_weight: ``[head_dim]`` bf16 -- ``DeepseekV41Compressor.norm`` + holds its weight in the model dtype, and the multiply is + done in fp32 by promoting it, exactly as the module does. + :param positions: ``[num_tokens]`` int32 or int64. An odd position completes a group + with its even predecessor. + :param req: ``[num_tokens]`` int64 ``req_pool_idx``, the ``kv_state`` row + this token pairs through. + :param raw_out_loc: ``[num_tokens]`` int32 or int64, the token's FULL-pool slot + (the scheduler's ``out_cache_loc`` is int64). + ``0`` marks a padded graph row, which the kernel skips + entirely -- reading nothing and writing nothing, so no + spare pair-state row is needed. + :param eps: RMSNorm epsilon. + :param out: ``[num_tokens, head_dim]`` bf16 destination. Pass a persistent + buffer under CUDA graphs. + :return: ``out``, the pre-RoPE post-norm latent. + + .. note:: **Rows at an even position, and padded rows, are not written.** + They complete no group, so the kernel skips their output row rather than + paying for a store the caller discards. Anything already in ``out`` on + those rows survives the call. + """ + num_tokens, fused_dim = kv_input.shape + head_dim = fused_dim // 2 + if out is None: + out = kv_input.new_empty((num_tokens, head_dim), dtype=torch.bfloat16) + + # Norm only: the store half's arguments are irrelevant, fixed to share a build. + _jit_c2_module(head_dim, 64, 128, KVLayout.V4).decode( + kv_input, + kv_state, + out, + norm_weight, + positions, + req, + raw_out_loc, + float(eps), + int(ring_size), + ) + return out + + +def c2_decode_or_verify_norm_rope_store( + kv_input: torch.Tensor, + kv_state: torch.Tensor, + norm_weight: torch.Tensor, + positions: torch.Tensor, + req: torch.Tensor, + raw_out_loc: torch.Tensor, + eps: float, + freqs_cis: torch.Tensor, + k_cache: torch.Tensor, + *, + page_size: int, + ring_size: int, + draft_len: int = 1, + layout: Union[KVLayout, str] = KVLayout.V4, + out: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """``c2_decode_norm`` plus the whole main-KV write, in the same launch. + + ``out`` contains the pre-RoPE latent for the index-K branch's ``wk`` projection. + The cache store uses ``raw_out_loc // 2`` as its slot. + + :param freqs_cis: ``[max_pos, rope_dim]`` fp32, real/imag interleaved -- + ``torch.view_as_real(freqs).flatten(-2)``. Indexed + in-kernel at ``positions - 1``, the position the latent + stands for, so there is no gather launch. + :param k_cache: the compressed KV pool buffer for this layer. + :param page_size: slots per page of that pool (``page_size // ratio``). + :param layout: the pool's :class:`KVLayout`. The fp8 layouts (``V4``, + ``V41``) store the fp4 fake-quantized value; ``V41_FP4`` + stores the e2m1 codes themselves, rounding once. + """ + num_tokens, fused_dim = kv_input.shape + head_dim = fused_dim // 2 + if out is None: + out = kv_input.new_empty((num_tokens, head_dim), dtype=torch.bfloat16) + + layout = KVLayout.parse(layout) + module = _jit_c2_module(head_dim, freqs_cis.shape[-1], page_size, layout) + module.decode_fusion( + kv_input, + kv_state, + out, + norm_weight, + positions, + req, + raw_out_loc, + eps, + freqs_cis, + k_cache, + ring_size, + draft_len, + ) + return out + + +# DO NOT try to modify the alias + +c2_decode_norm_rope_store: Final = c2_decode_or_verify_norm_rope_store +c2_verify_norm_rope_store: Final = c2_decode_or_verify_norm_rope_store diff --git a/python/sglang/kernels/ops/attention/dsv4/kv_layout.py b/python/sglang/kernels/ops/attention/dsv4/kv_layout.py new file mode 100644 index 000000000000..5edd618a629d --- /dev/null +++ b/python/sglang/kernels/ops/attention/dsv4/kv_layout.py @@ -0,0 +1,73 @@ +"""Paged fp8 / fp4 KV cache layouts of the DeepSeek-V4 family sparse MLA decode kernels. + +A page block stores ``page_size`` data rows followed by ``page_size`` scale rows. +The reader selects the format from the bytes per token (the last dim of the +``(num_pages, page_size, 1, bytes_per_token)`` view) and requires the page +stride to be a multiple of its TMA row stride, which :meth:`KVLayout.page_bytes` +pads to. Mirrors ``sgl_kernel/deepseek_v4/kv_layout.cuh``. +""" + +from __future__ import annotations + +import enum +from typing import Union + + +class KVLayout(str, enum.Enum): + # 448 fp8 nope + 64 bf16 rope, 7 ue8m0 scales (+1 pad) per 64 values. + V4 = "v4" + # 512 fp8 (rope quantized too), 16 ue8m0 scales per 32 values. + V41 = "v41" + # 512 e2m1 packed two per byte (even index low nibble), 32 e4m3 scales per 16 values. + V41_FP4 = "v41_fp4" + + @property + def data_bytes(self) -> int: + return {KVLayout.V4: 576, KVLayout.V41: 512, KVLayout.V41_FP4: 256}[self] + + @property + def scale_bytes(self) -> int: + return {KVLayout.V4: 8, KVLayout.V41: 16, KVLayout.V41_FP4: 32}[self] + + @property + def tile_size(self) -> int: + """Values sharing one scale.""" + return {KVLayout.V4: 64, KVLayout.V41: 32, KVLayout.V41_FP4: 16}[self] + + @property + def bytes_per_token(self) -> int: + return self.data_bytes + self.scale_bytes + + @property + def page_align(self) -> int: + """Unit the page stride is padded to: the reader's TMA row stride.""" + return {KVLayout.V4: 576, KVLayout.V41: 512, KVLayout.V41_FP4: 256}[self] + + @property + def is_fp4(self) -> bool: + return self is KVLayout.V41_FP4 + + def page_bytes(self, page_size: int) -> int: + raw = page_size * self.bytes_per_token + return -(-raw // self.page_align) * self.page_align + + def scale_offset(self, page_size: int) -> int: + """Byte offset of the scale rows inside a page.""" + return page_size * self.data_bytes + + @property + def cpp_name(self) -> str: + """The C++ enumerator, for JIT template arguments.""" + return self.name + + @classmethod + def parse(cls, value: Union[str, KVLayout]) -> KVLayout: + if isinstance(value, KVLayout): + return value + return cls(str(value).lower()) + + +def is_valid_kv_layout_pair(kv: KVLayout, extra_kv: KVLayout) -> bool: + """The (main, extra) cache pairs the decode kernel accepts: identical layouts, + or the fp4 extra cache next to a V4.1 fp8 main cache.""" + return extra_kv is kv or (kv is KVLayout.V41 and extra_kv is KVLayout.V41_FP4) diff --git a/python/sglang/kernels/ops/attention/dsv4/metadata_kernel.py b/python/sglang/kernels/ops/attention/dsv4/metadata_kernel.py index 9025187f93f8..5647f08cac14 100644 --- a/python/sglang/kernels/ops/attention/dsv4/metadata_kernel.py +++ b/python/sglang/kernels/ops/attention/dsv4/metadata_kernel.py @@ -5,6 +5,66 @@ import triton.language as tl +@triton.jit +def _fill_all_compressed_indices_kernel( + page_table, + seq_lens, + page_indices, + raw_indices, + PAGE_STRIDE: tl.constexpr, + TOPK: tl.constexpr, + RATIO: tl.constexpr, + PAGE_SIZE: tl.constexpr, + BLOCK: tl.constexpr, +): + row = tl.program_id(0) + positions = tl.arange(0, BLOCK) + length = tl.load(seq_lens + row) + valid = (positions < length) & (positions < TOPK) + slots_per_page = PAGE_SIZE // RATIO + pages = tl.load( + page_table + row * PAGE_STRIDE + positions // slots_per_page, + mask=valid, + other=0, + ) + slots = pages * slots_per_page + positions % slots_per_page + tl.store( + page_indices + row * TOPK + positions, + tl.where(valid, slots, -1), + positions < TOPK, + ) + if raw_indices is not None: + tl.store( + raw_indices + row * TOPK + positions, + tl.where(valid, positions, -1), + positions < TOPK, + ) + + +def fill_all_compressed_indices( + page_table: torch.Tensor, + compressed_seq_lens: torch.Tensor, + page_indices: torch.Tensor, + *, + compress_ratio: int, + page_size: int, + raw_indices: Optional[torch.Tensor] = None, +) -> None: + """Fill all reachable slots; the caller guarantees compressed length <= top-k.""" + topk = page_indices.shape[1] + _fill_all_compressed_indices_kernel[(compressed_seq_lens.numel(),)]( + page_table, + compressed_seq_lens, + page_indices, + raw_indices, + page_table.stride(0), + topk, + compress_ratio, + page_size, + triton.next_power_of_2(topk), + ) + + @triton.jit(do_not_specialize=["bs", "num_write_tokens", "c128_cur_max_seq_len"]) def _init_compressed_attn_metadata_kernel( seq_lens_ptr, diff --git a/python/sglang/kernels/ops/attention/dsv41_small_metadata.py b/python/sglang/kernels/ops/attention/dsv41_small_metadata.py new file mode 100644 index 000000000000..d4783d2128d0 --- /dev/null +++ b/python/sglang/kernels/ops/attention/dsv41_small_metadata.py @@ -0,0 +1,129 @@ +"""Small-batch V4.1 page and compression metadata.""" + +import torch +import triton +import triton.language as tl + +from sglang.kernels.ops.attention.dsv4_attn_metadata_kernels import ( + PageTablePositionsResult, +) + + +@triton.jit +def _small_page_table( + REQ_TO_TOKEN, + REQS, + LENS, + OUT_LENS, + POS, + PAGES, + SWA, + STRIDE: tl.constexpr, + NUM_PAGES: tl.constexpr, + PAGE_SIZE: tl.constexpr, + WINDOW: tl.constexpr, + BLOCK: tl.constexpr, +): + row, tile = tl.program_id(0), tl.program_id(1) + if tile == 0: + length = tl.load(LENS + row).to(tl.int32) + tl.store(OUT_LENS + row, length) + tl.store(POS + row, length - 1) + tl.store(SWA + row, tl.minimum(length, WINDOW)) + req = tl.load(REQS + row).to(tl.int64) + p = tile * BLOCK + tl.arange(0, BLOCK) + slot = tl.load( + REQ_TO_TOKEN + req * STRIDE + p.to(tl.int64) * PAGE_SIZE, + mask=p < NUM_PAGES, + other=0, + ).to(tl.int32) + tl.store(PAGES + row * NUM_PAGES + p, slot // PAGE_SIZE, mask=p < NUM_PAGES) + + +def page_table_positions_small( + *, + req_to_token, + req_pool_indices_repeated, + seq_lens_casual, + max_seq_len, + page_size, + swa_window, +): + assert page_size > 0 and page_size & (page_size - 1) == 0 + rows = seq_lens_casual.numel() + pages = triton.cdiv(max_seq_len, page_size) + kw = dict(device=seq_lens_casual.device, dtype=torch.int32) + lengths, positions, swa = [torch.empty(rows, **kw) for _ in range(3)] + table = torch.empty((rows, pages), **kw) + _small_page_table[(rows, triton.cdiv(pages, 256))]( + req_to_token, + req_pool_indices_repeated, + seq_lens_casual, + lengths, + positions, + table, + swa, + req_to_token.stride(0), + pages, + page_size, + swa_window, + 256, + ) + return PageTablePositionsResult( + seq_lens_casual=lengths, + positions_casual=positions, + page_table=table, + swa_topk_lengths=swa, + ) + + +@triton.jit +def _low_ratio_metadata( + LENS, + LOC, + OUT1, + LEN1, + SPARSE1, + PAGE1, + OUT2, + LEN2, + SPARSE2, + PAGE2, + TOPK: tl.constexpr, + PADDED: tl.constexpr, + BLOCK: tl.constexpr, +): + row = tl.program_id(0) + length = tl.load(LENS + row).to(tl.int32) + loc = tl.load(LOC + row).to(tl.int64) + len1, len2 = tl.maximum(length, 1), tl.maximum(length >> 1, 1) + tl.store(OUT1 + row, loc) + tl.store(OUT2 + row, tl.where((length & 1) == 0, loc >> 1, -1)) + tl.store(LEN1 + row, len1) + tl.store(LEN2 + row, len2) + tl.store(SPARSE1 + row, tl.minimum(len1, TOPK)) + tl.store(SPARSE2 + row, tl.minimum(len2, TOPK)) + cols = tl.arange(0, BLOCK) + tl.store(PAGE1 + row * PADDED + cols, -1, cols < PADDED) + tl.store(PAGE2 + row * PADDED + cols, -1, cols < PADDED) + + +def low_ratio_metadata(seq_lens, out_loc, topk): + assert seq_lens.numel() == out_loc.numel() + rows = seq_lens.numel() + kw = dict(device=seq_lens.device, dtype=torch.int32) + padded = triton.cdiv(topk, 64) * 64 + outputs = [] + for _ in range(2): + outputs.extend( + [ + torch.empty(rows, device=out_loc.device, dtype=torch.int64), + torch.empty(rows, **kw), + torch.empty(rows, **kw), + torch.empty((rows, padded), **kw), + ] + ) + _low_ratio_metadata[(rows,)]( + seq_lens, out_loc, *outputs, topk, padded, triton.next_power_of_2(padded) + ) + return outputs diff --git a/test/registered/kernel/attention/dsv4/test_c2_verify.py b/test/registered/kernel/attention/dsv4/test_c2_verify.py new file mode 100644 index 000000000000..7eb55dc45ef9 --- /dev/null +++ b/test/registered/kernel/attention/dsv4/test_c2_verify.py @@ -0,0 +1,391 @@ +"""Verify compressor: exact decode replay, padding, wrap and rejected prefixes.""" + +import sys + +import pytest +import torch +from torch import nn + +from sglang.kernels.ops.attention.dsv4.c2 import c2_decode_or_verify_norm_rope_store +from sglang.kernels.ops.attention.dsv4.rmsnorm_fp32 import rmsnorm_fp32 +from sglang.srt.model_loader.utils import set_default_torch_dtype +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=60, stage="base-b-kernel-unit", runner_config="4-gpu-b200") +pytestmark = pytest.mark.skipif( + not torch.cuda.is_available() + or torch.version.cuda is None + or torch.cuda.get_device_capability()[0] < 10, + reason="the compressor packs FP4 with SM100 instructions", +) + +EPS = 1e-6 +HEAD_DIMS = (512,) +ROPE_DIM = 64 +RATIO = 2 +PAGE_SIZE = 128 +PAGE_BYTES = -(-584 * PAGE_SIZE // 576) * 576 +DRAFT_LENS = (2, 5, 6, 9) +VERIFY_BATCHES = (1, 3, 8) + + +class RMSNorm(nn.Module): + """fp32 statistics and fp32 weight multiply, cast back at the very end.""" + + def __init__(self, dim: int, eps: float): + super().__init__() + self.eps = eps + self.weight = nn.Parameter(torch.ones(dim)) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + if ( + x.is_cuda + and torch.version.cuda is not None + and x.dtype in (torch.bfloat16, torch.float32) + and self.weight.dtype in (torch.bfloat16, torch.float32) + and x.shape[-1] in (128, 512) + and x.is_contiguous() + and self.weight.is_contiguous() + ): + return rmsnorm_fp32(x, self.weight, self.eps) + dtype = x.dtype + x = x.float() + x = x * torch.rsqrt(x.square().mean(-1, keepdim=True) + self.eps) + return (self.weight * x).to(dtype) + + +def _norm(dim: int, seed: int) -> RMSNorm: + """`DeepseekV41Compressor.norm` as the model holds it: bf16 weight, because + the parameter is created inside `set_default_torch_dtype(model dtype)`.""" + with set_default_torch_dtype(torch.bfloat16): + norm = RMSNorm(dim, EPS).cuda() + assert norm.weight.dtype == torch.bfloat16 + # Not `ones`: a constant weight cannot catch a wrong per-element index. + g = torch.Generator(device="cuda").manual_seed(seed) + with torch.no_grad(): + norm.weight.copy_(torch.randn(dim, generator=g, device="cuda")) + return norm + + +def _freqs(max_pos, seed): + """`layer.freqs_cis` and the fp32 real/imag-interleaved view the kernel + indexes itself, at `positions - 1`.""" + g = torch.Generator(device="cuda").manual_seed(seed) + ang = torch.randn(max_pos, ROPE_DIM // 2, generator=g, device="cuda") + freqs = torch.polar(torch.ones_like(ang), ang) + return freqs, torch.view_as_real(freqs).flatten(-2).contiguous().float() + + +def _cache(max_slot): + """A compressed-pool buffer wide enough for `max_slot`, zeroed so an + untouched slot is recognizable.""" + return torch.zeros( + max_slot // PAGE_SIZE + 2, PAGE_BYTES, dtype=torch.uint8, device="cuda" + ) + + +def _spec_ring_size(draft_len): + """`get_compress_state_ring_size(2, is_speculative=True, draft_len)`.""" + return 1 << (draft_len + 1).bit_length() + + +def _verify_inputs(bs, draft_len, dim, seed, *, starts=None, pad_reqs=0): + """A target-verify batch: `draft_len` consecutive positions per request, + request-major. Block heads alternate parity by default, so some blocks open + by consuming the ring and some by parking into it.""" + g = torch.Generator(device="cuda").manual_seed(seed) + n = bs * draft_len + ring_size = _spec_ring_size(draft_len) + kv_input = torch.randn(n, 2 * dim, generator=g, device="cuda", dtype=torch.float32) + kv_state = torch.randn( + (bs + 1) * ring_size, 2 * dim, generator=g, device="cuda", dtype=torch.float32 + ) + if starts is None: + starts = torch.arange(bs, device="cuda") + 4 + offsets = torch.arange(draft_len, device="cuda") + positions = (starts[:, None] + offsets[None, :]).flatten().to(torch.int32) + req = torch.arange(bs, device="cuda", dtype=torch.int64).repeat_interleave( + draft_len + ) + raw_out_loc = torch.arange(n, device="cuda", dtype=torch.int32) * 2 + 3 + if pad_reqs: + # Graph padding pads whole request slots, aliasing a live + # `req_pool_idx`, and leaves their position buffer at zero. + pad = req >= bs - pad_reqs + raw_out_loc[pad] = 0 + positions[pad] = 0 + req[pad] = 0 + return kv_input, kv_state, positions, req, raw_out_loc, ring_size + + +def _decode_replay( + kv_input, kv_state, norm, positions, req, raw_out_loc, freqs_cis, cache, **kw +): + """The same block, one position per launch -- what one fused verify launch + has to reproduce. Step `j` is a decode step over row `j` of every request, + carrying the pair ring between steps exactly as the served decode path does. + """ + draft_len = kw["draft_len"] + n, dim = positions.shape[0], kv_input.shape[1] // 2 + rows = torch.arange(n, device="cuda").view(-1, draft_len) + out = torch.zeros(n, dim, device="cuda", dtype=torch.bfloat16) + for j in range(draft_len): + idx = rows[:, j] + out[idx] = c2_decode_or_verify_norm_rope_store( + kv_input[idx].contiguous(), + kv_state, + norm.weight.data, + positions[idx].contiguous(), + req[idx].contiguous(), + raw_out_loc[idx].contiguous(), + EPS, + freqs_cis, + cache, + page_size=PAGE_SIZE, + ring_size=kw["ring_size"], + out=torch.zeros(idx.numel(), dim, device="cuda", dtype=torch.bfloat16), + ) + return out + + +def _run_verify(kv_input, kv_state, norm, positions, req, raw_out_loc, **kw): + """One `c2_verify_norm_rope_store` call; returns `(latent, cache)`. `out` is + zeroed so rows the kernel skips compare equal to the replay's.""" + n, dim = positions.shape[0], kv_input.shape[1] // 2 + freqs_cis, cache = kw["freqs_cis"], kw["cache"] + got = c2_decode_or_verify_norm_rope_store( + kv_input, + kv_state, + norm.weight.data, + positions, + req, + raw_out_loc, + EPS, + freqs_cis, + cache, + page_size=PAGE_SIZE, + ring_size=kw["ring_size"], + draft_len=kw["draft_len"], + out=torch.zeros(n, dim, device="cuda", dtype=torch.bfloat16), + ) + return got, cache + + +@pytest.mark.parametrize("pad_reqs", (0, 1)) +@pytest.mark.parametrize("draft_len", DRAFT_LENS) +@pytest.mark.parametrize("bs", VERIFY_BATCHES) +def test_verify_matches_decode_replay(bs, draft_len, pad_reqs): + """The load-bearing property: one verify launch over a block equals + `draft_len` decode launches over the same rows -- same latents, same pair + ring, same cache bytes, bitwise. Verify takes the in-block partner from + `kv_input` where decode takes it from the ring, and that substitution has to + be invisible.""" + if pad_reqs >= bs: + pytest.skip("an all-padded batch has no live block to compare") + dim = HEAD_DIMS[0] + kv_input, kv_state, positions, req, raw_out_loc, ring_size = _verify_inputs( + bs, draft_len, dim, seed=20000 + bs * 97 + draft_len, pad_reqs=pad_reqs + ) + norm = _norm(dim, 20001 + bs + draft_len) + _, freqs_cis = _freqs(int(positions.max().item()) + 2, 20002 + draft_len) + slots_max = int((raw_out_loc // RATIO).max().item()) + cache_v, cache_d = _cache(slots_max), _cache(slots_max) + state_v, state_d = kv_state.clone(), kv_state.clone() + kw = dict(ring_size=ring_size, draft_len=draft_len) + + got, _ = _run_verify( + kv_input, + state_v, + norm, + positions, + req, + raw_out_loc, + freqs_cis=freqs_cis, + cache=cache_v, + **kw, + ) + expected = _decode_replay( + kv_input, state_d, norm, positions, req, raw_out_loc, freqs_cis, cache_d, **kw + ) + + assert cache_d.any(), "the replay stored nothing, so the comparison is empty" + assert torch.equal(got, expected), "latent differs from the decode replay" + assert torch.equal(state_v, state_d), "pair ring differs from the decode replay" + assert torch.equal(cache_v, cache_d), "cache bytes differ from the decode replay" + + +@pytest.mark.parametrize("draft_len", DRAFT_LENS) +def test_verify_reads_the_ring_only_on_the_first_row(draft_len): + """A block's first row is the only one allowed to consume the ring. Move the + slot it reads and its latent must move with it; every later row pairs inside + the block and must not notice.""" + bs, dim = 4, HEAD_DIMS[0] + # Odd heads, so every block opens by completing a group against the ring. + starts = 2 * torch.arange(bs, device="cuda") + 5 + kv_input, kv_state, positions, req, raw_out_loc, ring_size = _verify_inputs( + bs, draft_len, dim, seed=21000 + draft_len, starts=starts + ) + norm = _norm(dim, 21001 + draft_len) + _, freqs_cis = _freqs(int(positions.max().item()) + 2, 21002 + draft_len) + slots_max = int((raw_out_loc // RATIO).max().item()) + kw = dict(ring_size=ring_size, draft_len=draft_len, freqs_cis=freqs_cis) + + base, _ = _run_verify( + kv_input, + kv_state.clone(), + norm, + positions, + req, + raw_out_loc, + cache=_cache(slots_max), + **kw, + ) + heads = torch.arange(0, bs * draft_len, draft_len, device="cuda") + read = req[heads] * ring_size + (positions[heads].to(torch.int64) - 1) % ring_size + moved = kv_state.clone() + moved[read] += 1.0 + got, _ = _run_verify( + kv_input, + moved, + norm, + positions, + req, + raw_out_loc, + cache=_cache(slots_max), + **kw, + ) + + rest = torch.ones(bs * draft_len, dtype=torch.bool, device="cuda") + rest[heads] = False + assert not torch.equal(got[heads], base[heads]), "a block head ignored the ring" + assert torch.equal(got[rest], base[rest]), "a later row went through the ring" + + +def test_verify_rejects_a_ring_narrower_than_the_block(): + """`ring_size > draft_len` is the whole reason a block's own publishes stay + off the slot its first row reads, so the kernel refuses a narrower ring + rather than racing quietly.""" + bs, draft_len, dim = 2, 4, HEAD_DIMS[0] + kv_input, kv_state, positions, req, raw_out_loc, _ = _verify_inputs( + bs, draft_len, dim, seed=22000 + ) + norm = _norm(dim, 22001) + _, freqs_cis = _freqs(int(positions.max().item()) + 2, 22002) + with pytest.raises(Exception, match="must be wider than the draft length"): + _run_verify( + kv_input, + kv_state, + norm, + positions, + req, + raw_out_loc, + cache=_cache(int((raw_out_loc // RATIO).max().item())), + freqs_cis=freqs_cis, + ring_size=draft_len, + draft_len=draft_len, + ) + + +@pytest.mark.parametrize("start", (31, 32)) +@pytest.mark.parametrize("dtype", (torch.int32, torch.int64)) +def test_rejected_prefix_then_next_verify(start, dtype): + # Compare with a decode history that never saw the rejected suffix. The + # next verify starts at the committed position, including across ring wrap. + bs, draft_len, dim = 3, 6, 512 + inputs, initial, pos, req, loc, ring = _verify_inputs( + bs, + draft_len, + dim, + 23000, + starts=torch.full((bs,), start, device="cuda"), + ) + pos, loc = pos.to(dtype), loc.to(dtype) + norm = _norm(dim, 23001) + _, freqs = _freqs(start + 2 * draft_len + 2, 23002) + rows = torch.arange(bs * draft_len, device="cuda").view(bs, draft_len) + for accepted in range(1, draft_len + 1): + state = initial.clone() + reference = initial.clone() + cache, ref_cache = _cache(128), _cache(128) + kw = dict(draft_len=draft_len, ring_size=ring, freqs_cis=freqs) + _run_verify(inputs, state, norm, pos, req, loc, cache=cache, **kw) + for j in range(accepted): + idx = rows[:, j] + c2_decode_or_verify_norm_rope_store( + inputs[idx], + reference, + norm.weight.data, + pos[idx], + req[idx], + loc[idx], + EPS, + freqs, + ref_cache, + page_size=PAGE_SIZE, + ring_size=ring, + ) + # Do not count speculative cache bytes that have no committed reader. + cache.zero_() + ref_cache.zero_() + next_inputs = inputs.flip(0).contiguous() + got, _ = _run_verify( + next_inputs, + state, + norm, + pos + accepted, + req, + loc, + cache=cache, + **kw, + ) + expected = _decode_replay( + next_inputs, + reference, + norm, + pos + accepted, + req, + loc, + freqs, + ref_cache, + ring_size=ring, + draft_len=draft_len, + ) + assert torch.equal(got, expected), f"{start=} {accepted=}: latent differs" + assert torch.equal(cache, ref_cache), f"{start=} {accepted=}: cache differs" + + +def test_verify_cuda_graph_replay(): + bs, draft_len, dim = 3, 6, 512 + inputs, initial, pos, req, loc, ring = _verify_inputs(bs, draft_len, dim, 24000) + norm = _norm(dim, 24001) + _, freqs = _freqs(32, 24002) + state, cache = initial.clone(), _cache(128) + kw = dict(ring_size=ring, draft_len=draft_len, freqs_cis=freqs, cache=cache) + _run_verify(inputs, state, norm, pos, req, loc, **kw) + torch.cuda.synchronize() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + got, _ = _run_verify(inputs, state, norm, pos, req, loc, **kw) + state.copy_(initial) + cache.zero_() + graph.replay() + ref_cache = _cache(128) + expected = _decode_replay( + inputs, + initial, + norm, + pos, + req, + loc, + freqs, + ref_cache, + ring_size=ring, + draft_len=draft_len, + ) + assert torch.equal(got, expected) + assert torch.equal(state, initial) + assert torch.equal(cache, ref_cache) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__])) diff --git a/test/registered/kernel/attention/dsv4/test_v41_kv_store.py b/test/registered/kernel/attention/dsv4/test_v41_kv_store.py new file mode 100644 index 000000000000..638ee2a99df0 --- /dev/null +++ b/test/registered/kernel/attention/dsv4/test_v41_kv_store.py @@ -0,0 +1,433 @@ +"""Byte-exactness of the V4.1 (fp8 / fp4) FlashMLA KV cache store kernels. + +Every store kernel is compared byte for byte with the pure-torch quantizers of +the two formats (``torch_quant.quantize_k_cache_v41`` / ``_v41_fp4``), which +follow the decode kernel's own reference quantizer. +""" + +import unittest + +import torch + +from sglang.kernels.ops.attention.dsv4.kv_layout import KVLayout +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci(est_time=120, stage="base-b-kernel-unit", runner_config="4-gpu-b200") + + +from typing import Optional + +import torch + +FP8_MAX = 448.0 +FP4_MAX = 6.0 +FP8_BLOCK_SIZE = 32 +FP4_BLOCK_SIZE = 32 +FP4_AMAX_FLOOR = 6 * 2.0**-126 + + +def ceil_pow2(x: torch.Tensor) -> torch.Tensor: + """2 ** ceil(log2(x)) for positive fp32 x, computed on the IEEE bits so the + result is exact at powers of two.""" + bits = x.contiguous().view(torch.int32) + exponent = ((bits >> 23) & 0xFF) - 127 + has_mantissa = (bits & 0x7FFFFF) != 0 + exponent = exponent + has_mantissa.to(torch.int32) + return ((exponent + 127) << 23).view(torch.float32) + + +def block_scale(x: torch.Tensor, block_size: int, fmax: float, amax_floor: float): + """Per-block ue8m0 scale, as fp32 powers of two, shape [..., N // block_size].""" + amax = x.float().unflatten(-1, (-1, block_size)).abs().amax(dim=-1) + amax = amax.clamp_min(amax_floor) + # The kernel multiplies by the fp32 reciprocal rather than dividing. A Python + # scalar keeps this free of host tensors, so it can run under CUDA graph capture. + return ceil_pow2(amax * (1.0 / fmax)) + + +def round_fp4(x: torch.Tensor) -> torch.Tensor: + """Round fp32 values in [-6, 6] onto the e2m1 grid with round-to-nearest-even.""" + magnitude = x.abs() + step = torch.where(magnitude < 2.0, 0.5, torch.where(magnitude < 4.0, 1.0, 2.0)) + return torch.round(magnitude / step) * step * torch.sign(x) + + +def fake_quant_fp4(x: torch.Tensor, block_size: int = FP4_BLOCK_SIZE) -> torch.Tensor: + """Quantize to fp4 (per-block ue8m0 scale) and back, in x's dtype.""" + scale = block_scale(x, block_size, FP4_MAX, FP4_AMAX_FLOOR) + scaled = x.float().unflatten(-1, (-1, block_size)) / scale.unsqueeze(-1) + deq = round_fp4(scaled.clamp(-FP4_MAX, FP4_MAX)) * scale.unsqueeze(-1) + return deq.flatten(-2).to(x.dtype) + + +def fake_quant_compressed_kv(x: torch.Tensor) -> torch.Tensor: + """FP4 round-trip with one E4M3FN scale per 16 compressed-KV elements. + + Round amax / 6 to E4M3 with ties to even, clamping the scale to its + positive finite range [2**-9, 448]. Zero blocks remain zero. + """ + blocks = x.float().unflatten(-1, (-1, 16)) + amax = blocks.abs().amax(dim=-1, keepdim=True) + scale = (amax * (1.0 / FP4_MAX)).clamp(min=2**-9, max=FP8_MAX) + scale = scale.to(torch.float8_e4m3fn).float() + scaled = (blocks / scale).clamp(-FP4_MAX, FP4_MAX) + deq = round_fp4(scaled) * scale + return deq.flatten(-2).to(x.dtype) + + +# --------------------------------------------------------------------------- +# Pure-torch references of the paged V4.1 KV cache formats read by the sparse +# decode kernel (528 B/token fp8 "V41", 288 B/token fp4 "V41_FP4"). They follow +# the kernel's own reference quantizer and are what the store / dequant kernels +# and the tests are checked against, byte for byte. +# --------------------------------------------------------------------------- + +_E2M1_MAGNITUDES = (0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0) + + +def cast_scale_inv_to_ue8m0(scale_inv: torch.Tensor) -> torch.Tensor: + """``2 ** ceil(log2(max(scale_inv, 1e-4)))`` as fp32, computed on the IEEE + bits so that it is exact at (and just above) powers of two.""" + scale_inv = scale_inv.float() + scale = ceil_pow2(torch.clamp_min(scale_inv, 1e-4)) + # ceil_pow2 works on the bits of a finite value; a NaN or inf amax passes + # through (both become the ue8m0 NaN byte, but only the NaN one turns the + # whole tile's payload into NaN). + return torch.where(torch.isfinite(scale_inv), scale, scale_inv) + + +def quantize_to_e2m1_codes(x: torch.Tensor) -> torch.Tensor: + """Round to the nearest e2m1 value with the semantics of + ``cvt.rn.satfinite.e2m1x2.f32`` (ties to even, saturating to +-6) and return + the 4-bit codes as uint8. The sign is kept for values that round to zero + (``-0.0`` and small negatives give the code ``0x8``); NaN maps to code 0.""" + x = x.float() + mags = torch.tensor(_E2M1_MAGNITUDES, dtype=torch.float32, device=x.device) + sign = torch.signbit(x).to(torch.uint8) << 3 + a = torch.nan_to_num(x.abs(), nan=0.0, posinf=6.0).clamp_max(6.0) + mids = (mags[:-1] + mags[1:]) / 2 + code = torch.bucketize(a, mids, right=True) + on_tie = (a.unsqueeze(-1) == mids).any(dim=-1) + tie_code = torch.bucketize(a, mids, right=False) + code = torch.where(on_tie, tie_code + (tie_code & 1), code) + return sign | code.to(torch.uint8) + + +def dequantize_e2m1_codes(codes: torch.Tensor) -> torch.Tensor: + mags = torch.tensor(_E2M1_MAGNITUDES, dtype=torch.float32, device=codes.device) + val = mags[(codes & 7).long()] + return torch.where((codes & 8) != 0, -val, val) + + +def quantize_k_cache_v41( + k: torch.Tensor, page_bytes: Optional[int] = None +) -> torch.Tensor: + """``k`` ``[num_pages, page_size, 512]`` -> uint8 ``[num_pages, page_bytes]`` + pages of the V41 layout: 512 e4m3 per token, then 16 ue8m0 scales per token + (one per 32 values), ``scale = 2 ** ceil(log2(max(amax / 448, 1e-4)))``.""" + num_pages, page_size, d = k.shape + assert d == 512 + x = k.float().view(num_pages, page_size, 16, 32) + scale = cast_scale_inv_to_ue8m0(x.abs().amax(dim=-1) / 448.0) + data = (x / scale.unsqueeze(-1)).to(torch.float8_e4m3fn).view(torch.uint8) + scale_u8 = scale.to(torch.float8_e8m0fnu).view(torch.uint8) + raw = page_size * 528 + if page_bytes is None: + page_bytes = -(-raw // 512) * 512 + assert page_bytes >= raw + out = torch.zeros((num_pages, page_bytes), dtype=torch.uint8, device=k.device) + out[:, : page_size * 512] = data.reshape(num_pages, page_size * 512) + out[:, page_size * 512 : raw] = scale_u8.reshape(num_pages, page_size * 16) + return out + + +def dequantize_k_cache_v41(pages: torch.Tensor, page_size: int) -> torch.Tensor: + """Inverse of :func:`quantize_k_cache_v41`: ``[num_pages, page_size, 512]`` bf16.""" + num_pages = pages.shape[0] + pages = pages.view(torch.uint8) + data = pages[:, : page_size * 512].reshape(num_pages, page_size, 512) + scale = pages[:, page_size * 512 : page_size * 528].reshape( + num_pages, page_size, 16 + ) + values = data.view(torch.float8_e4m3fn).to(torch.bfloat16) + scale_bf16 = scale.view(torch.float8_e8m0fnu).to(torch.bfloat16) + return (values.view(num_pages, page_size, 16, 32) * scale_bf16.unsqueeze(-1)).view( + num_pages, page_size, 512 + ) + + +def quantize_k_cache_v41_fp4( + k: torch.Tensor, page_bytes: Optional[int] = None +) -> torch.Tensor: + """``k`` ``[num_pages, page_size, 512]`` -> uint8 ``[num_pages, page_bytes]`` + pages of the V41_FP4 layout: 256 B of e2m1 codes per token (even index in + the low nibble), then 32 e4m3 scales per token (one per 16 values), + ``scale = e4m3(clamp(amax / 6, 2**-9, 448))``. A NaN element poisons its + tile: NaN scale, zero codes.""" + num_pages, page_size, d = k.shape + assert d == 512 + x = k.float().view(num_pages, page_size, 32, 16) + amax = torch.nan_to_num(x.abs(), nan=float("inf")).amax(dim=-1) + scale = torch.clamp(amax / 6.0, 2.0**-9, 448.0).to(torch.float8_e4m3fn) + scale = torch.where(torch.isinf(amax), torch.full_like(scale, float("nan")), scale) + codes = quantize_to_e2m1_codes(x / scale.float().unsqueeze(-1)) + codes = codes.view(num_pages, page_size, 512) + packed = codes[..., 0::2] | (codes[..., 1::2] << 4) + raw = page_size * 288 + if page_bytes is None: + page_bytes = -(-raw // 256) * 256 + assert page_bytes >= raw + out = torch.zeros((num_pages, page_bytes), dtype=torch.uint8, device=k.device) + out[:, : page_size * 256] = packed.reshape(num_pages, page_size * 256) + out[:, page_size * 256 : raw] = scale.view(torch.uint8).reshape( + num_pages, page_size * 32 + ) + return out + + +def dequantize_k_cache_v41_fp4(pages: torch.Tensor, page_size: int) -> torch.Tensor: + """Inverse of :func:`quantize_k_cache_v41_fp4`: ``[num_pages, page_size, 512]`` + bf16. ``e2m1 * e4m3`` has at most 2 + 4 significant bits, so the product is + exact in bf16, as in the kernel.""" + num_pages = pages.shape[0] + pages = pages.view(torch.uint8) + data = pages[:, : page_size * 256].reshape(num_pages, page_size, 256) + scale = pages[:, page_size * 256 : page_size * 288].reshape( + num_pages, page_size, 32 + ) + codes = torch.empty( + (num_pages, page_size, 512), dtype=torch.uint8, device=pages.device + ) + codes[..., 0::2] = data & 0xF + codes[..., 1::2] = data >> 4 + values = dequantize_e2m1_codes(codes).view(num_pages, page_size, 32, 16) + out = values * scale.view(torch.float8_e4m3fn).float().unsqueeze(-1) + return out.view(num_pages, page_size, 512).to(torch.bfloat16) + + +def rope_tail( + x: torch.Tensor, freqs: torch.Tensor, rope_dim: int, inverse: bool = False +) -> torch.Tensor: + """Rotate the last rope_dim features of x [T, ..., D] with complex freqs [T, rope_dim // 2].""" + head, tail = x[..., :-rope_dim], x[..., -rope_dim:] + tc = torch.view_as_complex(tail.float().unflatten(-1, (-1, 2)).contiguous()) + f = freqs.conj() if inverse else freqs + f = f.view(x.shape[0], *([1] * (x.ndim - 2)), rope_dim // 2) + rotated = torch.view_as_real(tc * f).flatten(-2).to(x.dtype) + return torch.cat([head, rotated], dim=-1) + + +REFERENCE = { + KVLayout.V41: quantize_k_cache_v41, + KVLayout.V41_FP4: quantize_k_cache_v41_fp4, +} + + +def _sm100(): + return ( + torch.cuda.is_available() + and torch.version.cuda is not None + and torch.cuda.get_device_capability()[0] >= 10 + ) + + +def token_rows(pages, layout, page_size, locs): + """The (data row, scale row) bytes of the tokens at ``locs``.""" + locs = locs.long() + page, offset = locs // page_size, locs % page_size + data_cols = torch.arange(layout.data_bytes, device=pages.device) + scale_cols = torch.arange(layout.scale_bytes, device=pages.device) + data = pages[page[:, None], offset[:, None] * layout.data_bytes + data_cols] + scale = pages[ + page[:, None], + layout.scale_offset(page_size) + + offset[:, None] * layout.scale_bytes + + scale_cols, + ] + return data, scale + + +def reference_pages(layout, page_size, num_pages, locs, values, page_bytes): + full = torch.zeros( + num_pages, page_size, 512, device=values.device, dtype=values.dtype + ) + full.view(-1, 512)[locs.long()] = values + return REFERENCE[layout](full, page_bytes=page_bytes) + + +@unittest.skipUnless(_sm100(), "the V4.1 KV layouts are SM100 kernels") +class TestV41KVStore(CustomTestCase): + def assert_tokens_equal(self, cache, ref, layout, page_size, locs): + got_data, got_scale = token_rows(cache, layout, page_size, locs) + exp_data, exp_scale = token_rows(ref, layout, page_size, locs) + self.assertTrue(torch.equal(got_scale, exp_scale), "scale rows differ") + self.assertTrue(torch.equal(got_data, exp_data), "data rows differ") + + def assert_untouched_zero(self, cache, layout, page_size, locs): + num_slots = cache.shape[0] * page_size + written = torch.zeros(num_slots, dtype=torch.bool, device=cache.device) + written[locs.long()] = True + others = torch.arange(num_slots, device=cache.device)[~written] + data, scale = token_rows(cache, layout, page_size, others) + self.assertEqual(int(data.sum()) + int(scale.sum()), 0) + + def test_c1_c2_decode_store(self): + """The ratio-1 / ratio-2 decode compressors write the V4.1 layouts: the cache + holds the quantized rope_tail of the pre-RoPE latent the kernel publishes + (bitwise; the fp8 layout after the model's fp4 fake quantization), and the + latent is the torch RMSNorm to within an fp32-reduction-order bf16 ulp.""" + + from sglang.kernels.ops.attention.dsv4.c1 import c1_decode_norm_rope_store + from sglang.kernels.ops.attention.dsv4.c2 import ( + c2_decode_or_verify_norm_rope_store, + ) + + g = torch.Generator(device="cuda").manual_seed(4) + eps = 1e-6 + angles = torch.randn(4096, 32, generator=g, device="cuda") + freqs = torch.polar(torch.ones_like(angles), angles) + freqs_real = torch.view_as_real(freqs).flatten(-2) + w = (torch.randn(512, generator=g, device="cuda") * 0.3 + 1).to(torch.bfloat16) + + def torch_norm(x): + xf = x.float() + return ( + xf * torch.rsqrt(xf.square().mean(-1, keepdim=True) + eps) * w.float() + ).to(torch.bfloat16) + + def stored_reference(layout, rotated, page_size, num_pages, locs, page_bytes): + values = ( + rotated + if layout is KVLayout.V41_FP4 + else fake_quant_compressed_kv(rotated) + ) + return reference_pages( + layout, page_size, num_pages, locs, values, page_bytes + ) + + n = 200 + for layout in (KVLayout.V41, KVLayout.V41_FP4): + for page_size, num_pages in ((256, 4), (128, 8)): + with self.subTest(kernel="c1", layout=layout.name, page_size=page_size): + x = (torch.randn(n, 512, generator=g, device="cuda") * 2).to( + torch.bfloat16 + ) + pos = torch.randint( + 0, 4096, (n,), generator=g, device="cuda", dtype=torch.int64 + ) + out_loc = ( + torch.randperm( + num_pages * page_size - 1, generator=g, device="cuda" + )[:n].to(torch.int32) + + 1 + ) + out_loc[3] = 0 # a padded graph row publishes nothing + cache = torch.zeros( + num_pages, + layout.page_bytes(page_size), + dtype=torch.uint8, + device="cuda", + ) + latent = c1_decode_norm_rope_store( + x, + w, + pos, + out_loc, + eps, + freqs_real, + cache, + page_size=page_size, + layout=layout, + ) + torch.testing.assert_close( + latent, torch_norm(x), rtol=2**-7, atol=2**-14 + ) + valid = out_loc > 0 + rotated = rope_tail(latent, freqs[pos], 64) + ref = stored_reference( + layout, + rotated[valid], + page_size, + num_pages, + out_loc[valid], + cache.shape[1], + ) + self.assert_tokens_equal( + cache, ref, layout, page_size, out_loc[valid] + ) + self.assert_untouched_zero(cache, layout, page_size, out_loc[valid]) + with self.subTest(kernel="c2", layout=layout.name, page_size=page_size): + ring = 2 + kv_new = torch.randn(n, 512, generator=g, device="cuda") * 2 + score = torch.randn(n, 512, generator=g, device="cuda") + kv_old = torch.randn(n, 512, generator=g, device="cuda") * 2 + kv_input = torch.cat([kv_new, score], dim=-1).contiguous() + req = torch.arange(n, device="cuda", dtype=torch.int64) + # Odd positions complete a pair; one even (pending) row and one padded row. + pos = ( + 2 + * torch.randint( + 0, 2000, (n,), generator=g, device="cuda", dtype=torch.int64 + ) + + 1 + ) + pos[5] = 4 + state = torch.randn(n * ring + 4, 1024, generator=g, device="cuda") + read_rows = req * ring + (pos - 1) % ring + # Equal scores make the pair pool the exact mean. + state[read_rows, :512] = kv_old + state[read_rows, 512:] = score + raw_out_loc = ( + torch.randperm( + num_pages * page_size - 1, generator=g, device="cuda" + )[:n].to(torch.int32) + + 1 + ) * 2 + raw_out_loc[7] = 0 + cache = torch.zeros( + num_pages, + layout.page_bytes(page_size), + dtype=torch.uint8, + device="cuda", + ) + latent = c2_decode_or_verify_norm_rope_store( + kv_input, + state, + w, + pos, + req, + raw_out_loc, + eps, + freqs_real, + cache, + page_size=page_size, + ring_size=ring, + layout=layout, + ) + valid = (raw_out_loc != 0) & (pos % 2 == 1) + pooled = ((kv_old + kv_new) / 2).to(torch.bfloat16) + torch.testing.assert_close( + latent[valid], + torch_norm(pooled)[valid], + rtol=2**-7, + atol=2**-14, + ) + rotated = rope_tail(latent, freqs[(pos - 1).clamp_min(0)], 64) + slots = raw_out_loc >> 1 + ref = stored_reference( + layout, + rotated[valid], + page_size, + num_pages, + slots[valid], + cache.shape[1], + ) + self.assert_tokens_equal( + cache, ref, layout, page_size, slots[valid] + ) + self.assert_untouched_zero(cache, layout, page_size, slots[valid]) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/kernel/attention/test_dsv41_small_metadata.py b/test/registered/kernel/attention/test_dsv41_small_metadata.py new file mode 100644 index 000000000000..723d561686c7 --- /dev/null +++ b/test/registered/kernel/attention/test_dsv41_small_metadata.py @@ -0,0 +1,56 @@ +"""Integer metadata equivalence with changing CUDA graph inputs.""" + +import sys + +import pytest +import torch + +from sglang.kernels.ops.attention.dsv4_attn_metadata_kernels import ( + BuildPageTablePositions, +) +from sglang.kernels.ops.attention.dsv41_small_metadata import ( + page_table_positions_small, +) +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=60, stage="base-b-kernel-unit", runner_config="1-gpu-large") + + +@pytest.mark.parametrize("rows", [1, 5, 6, 8]) +@pytest.mark.parametrize("pages", [1, 17, 4098]) +@pytest.mark.parametrize("dtype", [torch.int32, torch.int64]) +def test_pages_graph(rows, pages, dtype): + # Rows have a larger physical stride than the logical table. + mapping = torch.randint(-256, 1 << 24, (9, pages * 256 + 512), device="cuda") + reqs = torch.arange(rows, device="cuda", dtype=dtype) + lens = torch.arange(rows, device="cuda", dtype=dtype) + args = dict( + req_to_token=mapping, + req_pool_indices_repeated=reqs, + seq_lens_casual=lens, + max_seq_len=pages * 256, + page_size=256, + swa_window=128, + ) + for _ in range(3): + page_table_positions_small(**args) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + out = page_table_positions_small(**args) + for replay in range(5): + mapping.random_(-256, 1 << 24) + reqs.copy_((torch.arange(rows, device="cuda") + replay) % 9) + lens.copy_(torch.arange(rows, device="cuda") * 127 + replay - 1) + graph.replay() + ref = BuildPageTablePositions.triton(**args) + for name in ( + "seq_lens_casual", + "positions_casual", + "page_table", + "swa_topk_lengths", + ): + assert torch.equal(getattr(out, name), getattr(ref, name)), name + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__])) From b14c741c2ebb38d64af3b38fd4b7bdafb480f745 Mon Sep 17 00:00:00 2001 From: hnyls2002 Date: Tue, 15 Sep 2026 15:53:41 -0700 Subject: [PATCH 11/30] dsv4.1: extract KV store and dequantization paths --- .../csrc/deepseek_v4/fused_norm_rope_v2.cuh | 63 +- .../jit/csrc/deepseek_v4/main_norm_rope.cuh | 59 +- .../kernels/jit/csrc/deepseek_v4/store.cuh | 167 ++++- .../sglang/kernels/ops/attention/dsv4/attn.py | 51 +- .../kernels/ops/attention/dsv4/compress.py | 18 +- .../ops/attention/dsv4/dequant_k_cache.py | 170 +++++- .../kernels/ops/attention/dsv4/elementwise.py | 17 +- .../srt/layers/attention/dsv4/torch_quant.py | 194 ++++++ .../attention/dsv4/test_v41_kv_dequant.py | 132 ++++ .../attention/dsv4/test_v41_kv_store.py | 575 ++++++++++++------ 10 files changed, 1192 insertions(+), 254 deletions(-) create mode 100644 python/sglang/srt/layers/attention/dsv4/torch_quant.py create mode 100644 test/registered/kernel/attention/dsv4/test_v41_kv_dequant.py diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh index 0fda6b8a1972..8b4772f2a30d 100644 --- a/python/sglang/kernels/jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh @@ -9,6 +9,7 @@ #include #include +#include #include @@ -386,14 +387,17 @@ constexpr int64_t kFp8TwoPoolRowBytes = 512; // ---------------------------------------------------------------------------- // FlashMLA variant: kHeadDim = 512, 1 token per *block* (256 threads). // Each thread loads kVecSize=2 BF16, so 256 threads cover the full 512 elems. -// Cache layout: 584 bytes/token = 448 fp8 nope + 64 (=32 bf16x2) rope + 8 scale. +// Cache layout (kLayout): V4 = 584 bytes/token = 448 fp8 nope + 64 (=32 bf16x2) +// rope + 8 scale; V41 / V41_FP4 = the V4.1 fp8 (528 B) / fp4 (288 B) formats in +// which every dim is quantized, one scale per 32 / 16 values. // ---------------------------------------------------------------------------- template < typename DType, ForwardMode kMode, int32_t kPageBits, + bool kBf16Store, + deepseek_v4::KVLayout kLayout, bool kUsePDL, - bool kBf16Store = false, bool kFp8TwoPool = false> FLASHMLA_KERNEL void fused_norm_rope_flashmla(const __grid_constant__ FusedNormRopeStoreParams params) { using namespace device; @@ -407,12 +411,15 @@ FLASHMLA_KERNEL void fused_norm_rope_flashmla(const __grid_constant__ FusedNormR constexpr uint32_t kRopeWarp = kNumWarps - 1; // kBf16Store: write the whole head_dim as plain BF16 (no fp8 / no scale) into a // [num_slots, head_dim] bf16 cache (page_size==1) at row out_loc + static_assert(!(kBf16Store && kLayout != deepseek_v4::KVLayout::V4), "the bf16 store is not a paged layout"); + using Paged = deepseek_v4::PagedKV; // kFp8TwoPool: 512 B row holding the 448 fp8 nope + its UE8M0 scales, with rope // split off into a second [num_slots, kRopeDim] bf16 pool at the same row static_assert(!(kBf16Store && kFp8TwoPool)); constexpr int64_t kRowBytes = kBf16Store ? (kHeadDim * 2ll) : (kFp8TwoPool ? kFp8TwoPoolRowBytes : 576ll); constexpr int64_t kPageBytes = (kBf16Store || kFp8TwoPool) ? (kRowBytes << kPageBits) : host::div_ceil(584ll << kPageBits, 576) * 576; + static_assert(!(kFp8TwoPool && kLayout != deepseek_v4::KVLayout::V4), "the fp8 two-pool store is a V4 cache"); static_assert(kHeadDim == kBlockSize * kVecSize); static_assert(kRopeDim == kWarpThreads * kVecSize); static_assert(kHeadDim - kRopeDim == kRopeWarp * kWarpThreads * kVecSize); @@ -480,10 +487,37 @@ FLASHMLA_KERNEL void fused_norm_rope_flashmla(const __grid_constant__ FusedNormR } } + const auto row = Paged::row(params.kvcache, out_loc); + + if constexpr (kLayout != deepseek_v4::KVLayout::V4) { + // V4.1 layouts: the whole row is quantized. Match the unfused path, which + // quantizes the bf16 tensor the norm produces and rotates the tail in bf16: + // round the normed values, rotate, round again, then quantize the row. + using Packed = packed_t; + PDLTriggerSecondary(); + + auto rounded = cast(cast(fp32x2_t{data[0], data[1]})); + if (warp_id == kRopeWarp) { + const auto x_real = rounded.x; + const auto x_imag = rounded.y; + const auto freq_real = freq[0]; + const auto freq_imag = freq[1]; + rounded = cast( + cast(fp32x2_t{x_real * freq_real - x_imag * freq_imag, x_real * freq_imag + x_imag * freq_real})); + } + const float v[2] = {rounded.x, rounded.y}; + deepseek_v4::v41::store_row(row.data, row.scale, tx, v); + return; + } + + // V4 rows come from the paged helper. The bf16 cache is dense [num_slots, head_dim] + // rows and the fp8 two-pool cache is kRowBytes rows, both addressed by out_loc. const int64_t page = out_loc >> kPageBits; const int64_t offset = out_loc & ((1 << kPageBits) - 1); const auto page_ptr = params.kvcache + page * kPageBytes; - const auto value_ptr = page_ptr + offset * kRowBytes; + const auto value_ptr = kBf16Store ? params.kvcache + static_cast(out_loc) * (kHeadDim * 2) + : kFp8TwoPool ? page_ptr + offset * kRowBytes + : row.data; PDLTriggerSecondary(); @@ -534,8 +568,7 @@ FLASHMLA_KERNEL void fused_norm_rope_flashmla(const __grid_constant__ FusedNormR scale_ptr[0] = scale_ue8m0; scale_ptr[1] = scale_ue8m0; } else { - const auto scale_ptr = page_ptr + (576 << kPageBits) + offset * 8; - static_cast(scale_ptr)[warp_id] = scale_ue8m0; + static_cast(row.scale)[warp_id] = scale_ue8m0; } } } @@ -546,15 +579,19 @@ template < int64_t kHeadDim, int64_t kRopeDim, uint32_t kPageSize, - bool kUsePDL, - int32_t kPreshuffleSize = 0, - bool kBf16Store = false> + int32_t kPreshuffleSize, + bool kBf16Store, + deepseek_v4::KVLayout kLayout, + bool kUsePDL> struct FusedNormRopeKernel { static constexpr int32_t kLogPageSize = std::countr_zero(kPageSize); static constexpr bool kIsIndexer = (kHeadDim == 128); static_assert(!(kIsIndexer && kBf16Store), "bf16 store only for flashmla head_dim=512"); + static_assert( + !(kIsIndexer && kLayout != deepseek_v4::KVLayout::V4), "the V4.1 layouts are FlashMLA (head_dim=512) caches"); static constexpr int64_t kIndexerBytes = 132 * kPageSize; - static constexpr int64_t kFlashMLABytes = host::div_ceil(584 * kPageSize, 576) * 576; + static constexpr int64_t kFlashMLABytes = deepseek_v4::kv_page_bytes(kPageSize); + static_assert(kLayout != deepseek_v4::KVLayout::V4 || kFlashMLABytes == host::div_ceil(584 * kPageSize, 576) * 576); static constexpr int64_t kBf16Bytes = kHeadDim * 2 * kPageSize; // plain bf16 cache static constexpr int64_t kPageBytes = kBf16Store ? kBf16Bytes : (kIsIndexer ? kIndexerBytes : kFlashMLABytes); @@ -567,7 +604,7 @@ struct FusedNormRopeKernel { if constexpr (kIsIndexer) { return fused_norm_rope_indexer; } else { - return fused_norm_rope_flashmla; + return fused_norm_rope_flashmla; } } @@ -575,7 +612,8 @@ struct FusedNormRopeKernel { static constexpr auto select_fp8_2buff_kernel() { static_assert(!kIsIndexer, "fp8 two-pool store is only defined for the flashmla latent"); static_assert(!kBf16Store, "fp8 two-pool store and bf16 store are separate layouts"); - return fused_norm_rope_flashmla; + static_assert(kLayout == deepseek_v4::KVLayout::V4, "the fp8 two-pool store is a V4 cache"); + return fused_norm_rope_flashmla; } template @@ -791,4 +829,7 @@ struct FusedNormRopeKernel { } }; +// The JIT wrappers name the layouts as plain enumerators. +using enum deepseek_v4::KVLayout; + } // namespace sglang diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/main_norm_rope.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/main_norm_rope.cuh index 483d6fe1cd95..3afa029a6c42 100644 --- a/python/sglang/kernels/jit/csrc/deepseek_v4/main_norm_rope.cuh +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/main_norm_rope.cuh @@ -9,6 +9,7 @@ #include #include +#include #include @@ -259,13 +260,20 @@ struct FusedKNormRopeFlashMLAParams { float eps; }; -template +template < + typename DType, + int64_t kHeadDim, + int64_t kRopeDim, + typename PosT, + int32_t kPageBits, + deepseek_v4::KVLayout kLayout, + bool kUsePDL> K_KERNEL void fused_k_norm_rope_flashmla(const __grid_constant__ FusedKNormRopeFlashMLAParams params) { using namespace device; constexpr int64_t kVecSize = 2; constexpr uint32_t kRopeWarp = kFusedKNumWarps - 1; - constexpr int64_t kPageBytes = host::div_ceil(584ll << kPageBits, 576) * 576; + using Paged = deepseek_v4::PagedKV; static_assert(kHeadDim == kFusedKBlockSize * kVecSize); static_assert(kRopeDim == kWarpThreads * kVecSize); static_assert(kHeadDim - kRopeDim == kRopeWarp * kWarpThreads * kVecSize); @@ -325,10 +333,30 @@ K_KERNEL void fused_k_norm_rope_flashmla(const __grid_constant__ FusedKNormRopeF // here, not at the load, so the out_loc prefetch overlaps the norm above. if (out_loc < 0) return; - const int32_t page = out_loc >> kPageBits; - const int32_t offset = out_loc & ((1 << kPageBits) - 1); - const auto page_ptr = params.kvcache + page * kPageBytes; - const auto value_ptr = page_ptr + offset * 576; + const auto row = Paged::row(params.kvcache, out_loc); + + if constexpr (kLayout != deepseek_v4::KVLayout::V4) { + // V4.1 layouts: every dim is quantized, with one scale per 32 (fp8) or 16 (fp4) + // values. The reference quantizes the bf16 tensor kv_norm produces and rotates + // the tail in bf16, so round the normed values to the storage dtype, rotate, + // round again, then quantize the whole row. + using Packed = packed_t; + PDLTriggerSecondary(); + + auto rounded = cast(cast(fp32x2_t{data[0], data[1]})); + if (warp_id == kRopeWarp) { + const auto x_real = rounded.x; + const auto x_imag = rounded.y; + const auto freq_real = freq[0]; + const auto freq_imag = freq[1]; + rounded = cast( + cast(fp32x2_t{x_real * freq_real - x_imag * freq_imag, x_real * freq_imag + x_imag * freq_real})); + } + const float v[2] = {rounded.x, rounded.y}; + return deepseek_v4::v41::store_row(row.data, row.scale, tx, v); + } + + const auto value_ptr = row.data; PDLTriggerSecondary(); @@ -351,22 +379,30 @@ K_KERNEL void fused_k_norm_rope_flashmla(const __grid_constant__ FusedKNormRopeF const auto scale_ue8m0 = cast_to_ue8m0(scale_raw); const auto inv_scale = inv_scale_ue8m0(scale_ue8m0); const auto result = pack_fp8(x * inv_scale, y * inv_scale); - const auto scale_ptr = page_ptr + (576 << kPageBits) + offset * 8; + const auto scale_ptr = row.scale; reinterpret_cast(value_ptr)[tx] = result; if (lane_id == 0) static_cast(scale_ptr)[warp_id] = scale_ue8m0; } } -template +template < + typename DType, + int64_t kHeadDim, + int64_t kRopeDim, + uint32_t kPageSize, + deepseek_v4::KVLayout kLayout, + bool kUsePDL> struct FusedKNormRopeFlashMLAKernel { static constexpr int32_t kLogPageSize = std::countr_zero(kPageSize); - static constexpr int64_t kPageBytes = host::div_ceil(584 * kPageSize, 576) * 576; + static constexpr int64_t kPageBytes = deepseek_v4::kv_page_bytes(kPageSize); + static_assert(kLayout != deepseek_v4::KVLayout::V4 || kPageBytes == host::div_ceil(584 * kPageSize, 576) * 576); static_assert(std::has_single_bit(kPageSize), "kPageSize must be a power of 2"); static_assert(1 << kLogPageSize == kPageSize); static_assert(kHeadDim == 512 && kRopeDim == 64, "FlashMLA layout requires (512, 64)"); template - static constexpr auto kernel = fused_k_norm_rope_flashmla; + static constexpr auto kernel = + fused_k_norm_rope_flashmla; static void forward( const tvm::ffi::TensorView kv, @@ -881,4 +917,7 @@ struct FusedQIndexerRopeHadamardFp4QuantKernel { } }; +// The JIT wrappers name the layouts as plain enumerators. +using enum deepseek_v4::KVLayout; + } // namespace sglang diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/store.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/store.cuh index c08256e460e7..9d91f34e36c4 100644 --- a/python/sglang/kernels/jit/csrc/deepseek_v4/store.cuh +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/store.cuh @@ -8,6 +8,7 @@ #include #include +#include #include #include @@ -15,6 +16,7 @@ #include #include #include +#include namespace sglang { @@ -29,12 +31,21 @@ struct FusedStoreCacheParam { uint32_t num_tokens; }; +/// Parameters of the V4.1 (fp8 / fp4) FlashMLA store; `freqs_cis` is the per-token +/// (real, imag) pairs of the 64 RoPE dims, or nullptr when the input is already rotated. +struct FusedStoreCacheV41Param { + const void* __restrict__ input; + void* __restrict__ cache; + const void* __restrict__ indices; + const float* __restrict__ freqs_cis; + uint32_t num_tokens; +}; + template __global__ void fused_store_flashmla_cache(const __grid_constant__ FusedStoreCacheParam param) { using namespace device; - /// NOTE: 584 = 576 + 8 - constexpr int64_t kPageBytes = host::div_ceil(584 << kPageBits, 576) * 576; + using Paged = deepseek_v4::PagedKV; // each warp handles 64 elements, 8 warps, each block handles 1 row const auto& [input, cache, indices, num_tokens] = param; @@ -56,21 +67,82 @@ __global__ void fused_store_flashmla_cache(const __grid_constant__ FusedStoreCac const auto scale_ue8m0 = cast_to_ue8m0(scale_raw); const auto inv_scale = inv_scale_ue8m0(scale_ue8m0); const auto result = pack_fp8(x * inv_scale, y * inv_scale); - const int32_t page = index >> kPageBits; - const int32_t offset = index & ((1 << kPageBits) - 1); - const auto page_ptr = pointer::offset(cache, page * kPageBytes); - const auto value_ptr = pointer::offset(page_ptr, offset * 576); - const auto scale_ptr = pointer::offset(page_ptr, 576 << kPageBits, offset * 8); - static_cast(value_ptr)[tid] = result; - static_cast(scale_ptr)[wid] = scale_ue8m0; + const auto row = Paged::row(static_cast(cache), index); + reinterpret_cast(row.data)[tid] = result; + row.scale[wid] = scale_ue8m0; } else { const auto result = cast(elems); - const int32_t page = index >> kPageBits; - const int32_t offset = index & ((1 << kPageBits) - 1); - const auto page_ptr = pointer::offset(cache, page * kPageBytes); - const auto value_ptr = pointer::offset(page_ptr, offset * 576, 448); - static_cast(value_ptr)[tid - 7 * 32] = result; + const auto row = Paged::row(static_cast(cache), index); + reinterpret_cast(row.data + 448)[tid - 7 * 32] = result; + } + + PDLTriggerSecondary(); +} + +/// V4.1 store: one 256-thread block per token, thread `tx` owns elements (2tx, 2tx + 1) of the +/// 512-wide row. With kRope the last warp first rotates its (real, imag) pairs -- the RoPE +/// tail -- and rounds them back to the input dtype, as `rope_tail` does, so that the caller can +/// hand in the un-rotated, un-quantized latent and the e2m1 / e4m3 rounding happens exactly once. +/// Elements per thread of the V4.1 store, `512 / vec` threads per token: 4 for the fp8 rows +/// (8-byte loads, half the threads), 2 for the fp4 rows, whose per-element IEEE divisions +/// are better spread over more threads (measured on B200, bs 1..512). +constexpr uint32_t v41_store_vec_size(deepseek_v4::KVLayout layout) { + return layout == deepseek_v4::KVLayout::V41 ? 4 : 2; +} + +template < + typename Float, + typename IndicesT, + uint32_t kPageBits, + deepseek_v4::KVLayout kLayout, + bool kRope, + bool kUsePDL> +__global__ void fused_store_flashmla_cache_v41(const __grid_constant__ FusedStoreCacheV41Param param) { + using namespace device; + using Paged = deepseek_v4::PagedKV; + static_assert(kLayout != deepseek_v4::KVLayout::V4, "the V4 layout has its own kernel above"); + + constexpr uint32_t kVecSize = v41_store_vec_size(kLayout); + constexpr uint32_t kNopeLanes = (512 - 64) / kVecSize; // threads from here on hold the RoPE tail + using Packed = packed_t; + using Vec = AlignedVector; + + const auto& [input, cache, indices, freqs_cis, num_tokens] = param; + const uint32_t bid = blockIdx.x; + const uint32_t tid = threadIdx.x; + + PDLWaitPrimary(); + + const auto index = static_cast(indices)[bid]; + Vec elems; + elems.load(static_cast(input) + bid * 512, tid); + float v[kVecSize]; +#pragma unroll + for (uint32_t i = 0; i < kVecSize / 2; ++i) { + const auto [x, y] = cast(elems[i]); + v[2 * i] = x; + v[2 * i + 1] = y; } + if constexpr (kRope) { + if (tid >= kNopeLanes) { + // (real, imag) pairs of the tail, rotated and rounded back to the input dtype as + // `rope_tail` does, so that the caller can also pass pre-rotated rows. + AlignedVector freq; + freq.load(freqs_cis + bid * 64, tid - kNopeLanes); +#pragma unroll + for (uint32_t i = 0; i < kVecSize / 2; ++i) { + const auto x = v[2 * i]; + const auto y = v[2 * i + 1]; + const auto rotated = cast( + cast(fp32x2_t{x * freq[2 * i] - y * freq[2 * i + 1], x * freq[2 * i + 1] + y * freq[2 * i]})); + v[2 * i] = rotated.x; + v[2 * i + 1] = rotated.y; + } + } + } + + const auto row = Paged::row(static_cast(cache), index); + deepseek_v4::v41::store_row(row.data, row.scale, tid, v); PDLTriggerSecondary(); } @@ -120,16 +192,41 @@ __global__ void fused_store_indexer_cache(const __grid_constant__ FusedStoreCach PDLTriggerSecondary(); } -template +template struct FusedStoreCacheFlashMLAKernel { static constexpr int32_t kLogSize = std::countr_zero(kPageSize); - static constexpr int64_t kPageBytes = host::div_ceil(584 * kPageSize, 576) * 576; - static constexpr auto kernel = fused_store_flashmla_cache; + static constexpr bool kIsV4 = kLayout == deepseek_v4::KVLayout::V4; + static constexpr int64_t kPageBytes = deepseek_v4::kv_page_bytes(kPageSize); + static_assert(!kIsV4 || kPageBytes == host::div_ceil(584 * kPageSize, 576) * 576); static_assert(std::has_single_bit(kPageSize), "kPageSize must be a power of 2"); static_assert(1 << kLogSize == kPageSize); + template + static constexpr auto v41_kernel = fused_store_flashmla_cache_v41; + + /// Store rows that are already normed and rotated. static void run(tvm::ffi::TensorView input, tvm::ffi::TensorView cache, tvm::ffi::TensorView indices) { + launch(input, cache, indices, std::nullopt); + } + + /// V4.1 layouts only: rotate the RoPE tail in-kernel with the per-token `freqs_cis` + /// (`[num_tokens, 64]` fp32, real / imag interleaved) before quantizing. + static void run_rope( + tvm::ffi::TensorView input, + tvm::ffi::TensorView cache, + tvm::ffi::TensorView indices, + tvm::ffi::TensorView freqs_cis) { + static_assert(!kIsV4, "the V4 layout keeps its RoPE dims in bf16 and has no in-kernel RoPE"); + launch(input, cache, indices, freqs_cis); + } + + private: + static void launch( + tvm::ffi::TensorView input, + tvm::ffi::TensorView cache, + tvm::ffi::TensorView indices, + std::optional freqs_cis) { using namespace host; auto N = SymbolicSize{"num_tokens"}; @@ -148,16 +245,35 @@ struct FusedStoreCacheFlashMLAKernel { .with_dtype() .with_device(device_) .verify(indices); + if (freqs_cis.has_value()) { + // Real / imag interleaved, so the trailing dim is 64, not 32. + TensorMatcher({N, 64}).with_dtype().with_device(device_).verify(*freqs_cis); + } const auto num_tokens = static_cast(N.unwrap()); - const auto params = FusedStoreCacheParam{ - .input = input.data_ptr(), - .cache = cache.data_ptr(), - .indices = indices.data_ptr(), - .num_tokens = num_tokens, - }; + if (num_tokens == 0) return; const auto kBlockSize = 256; const auto num_blocks = num_tokens; - LaunchKernel(num_blocks, kBlockSize, device_.unwrap()).enable_pdl(kUsePDL)(kernel, params); + if constexpr (kIsV4) { + RuntimeCheck(!freqs_cis.has_value(), "the V4 layout has no in-kernel RoPE"); + const auto params = FusedStoreCacheParam{ + .input = input.data_ptr(), + .cache = cache.data_ptr(), + .indices = indices.data_ptr(), + .num_tokens = num_tokens, + }; + constexpr auto kernel = fused_store_flashmla_cache; + LaunchKernel(num_blocks, kBlockSize, device_.unwrap()).enable_pdl(kUsePDL)(kernel, params); + } else { + const auto params = FusedStoreCacheV41Param{ + .input = input.data_ptr(), + .cache = cache.data_ptr(), + .indices = indices.data_ptr(), + .freqs_cis = freqs_cis.has_value() ? static_cast(freqs_cis->data_ptr()) : nullptr, + .num_tokens = num_tokens, + }; + const auto kernel = freqs_cis.has_value() ? v41_kernel : v41_kernel; + LaunchKernel(num_blocks, 512 / v41_store_vec_size(kLayout), device_.unwrap()).enable_pdl(kUsePDL)(kernel, params); + } } }; @@ -202,4 +318,7 @@ struct FusedStoreCacheIndexerKernel { } }; +// The JIT wrappers name the layouts as plain enumerators. +using enum deepseek_v4::KVLayout; + } // namespace sglang diff --git a/python/sglang/kernels/ops/attention/dsv4/attn.py b/python/sglang/kernels/ops/attention/dsv4/attn.py index 8996bb5226f6..6ac304f295cb 100644 --- a/python/sglang/kernels/ops/attention/dsv4/attn.py +++ b/python/sglang/kernels/ops/attention/dsv4/attn.py @@ -1,4 +1,4 @@ -from typing import Literal, Tuple +from typing import Literal, Optional, Tuple, Union import torch import triton @@ -12,6 +12,7 @@ make_cpp_args, ) +from .kv_layout import KVLayout from .utils import make_name @@ -30,20 +31,33 @@ def _jit_fused_store_module( input_dtype: torch.dtype, index_dtype: torch.dtype, page_size: int, + layout: KVLayout, ): - args = make_cpp_args(input_dtype, index_dtype, page_size, is_arch_support_pdl()) - cname = "FlashMLA" if name == "flashmla" else "Indexer" + if name == "flashmla": + args = make_cpp_args( + input_dtype, index_dtype, page_size, layout.cpp_name, is_arch_support_pdl() + ) + # The V4 layout keeps its RoPE dims in bf16 and has no in-kernel RoPE. + cname = "FlashMLA" + wrappers = ["run"] if layout is KVLayout.V4 else ["run", "run_rope"] + else: + assert layout is KVLayout.V4, "only the FlashMLA cache has V4.1 layouts" + args = make_cpp_args(input_dtype, index_dtype, page_size, is_arch_support_pdl()) + cname, wrappers = "Indexer", ["run"] kernel_class = f"FusedStoreCache{cname}Kernel<{args}>" return load_jit( make_name("store_" + name), *args, cuda_files=["deepseek_v4/store.cuh"], - cuda_wrappers=[("run", f"{kernel_class}::run")], + cuda_wrappers=[(w, f"{kernel_class}::{w}") for w in wrappers], ) def get_paged_mqa_logits_metadata(seq_lens: torch.Tensor, page_size: int, num_sm: int): - assert page_size == 64 + # The schedule only depends on the sequence lengths (256-token splits), not + # on the page size; DeepGEMM takes 32 / 64 / 128 on SM100 and we page at 64 + # or 128 slots. + assert page_size in (64, 128), page_size seq_lens = seq_lens.view(-1).to(torch.int32) bs = int(seq_lens.shape[0]) metadata = seq_lens.new_empty(num_sm + 1, 2) @@ -67,8 +81,26 @@ def fused_store_cache( *, page_size: int, type: Literal["flashmla", "indexer"], + layout: Union[KVLayout, str] = KVLayout.V4, + freqs_cis: Optional[torch.Tensor] = None, ) -> None: + """Quantize ``input`` ``[num_tokens, 512]`` (bf16, normed and rotated) into the + paged cache at ``indices``. + + :param layout: the cache's :class:`KVLayout`. ``V4`` is the 584-byte layout + (fp8 nope, bf16 rope); ``V41`` (528 B) and ``V41_FP4`` (288 B) are the + V4.1 formats, fp8 with per-32 ue8m0 scales and e2m1 with per-16 e4m3 + scales over all 512 dims. + :param freqs_cis: V4.1 layouts only. ``[num_tokens, 32]`` complex or + ``[num_tokens, 64]`` fp32 (real / imag interleaved): rotate the 64-dim + RoPE tail in-kernel first, so that the caller passes the un-rotated, + un-quantized latent and the fp4 / fp8 rounding happens exactly once. + """ + layout = KVLayout.parse(layout) if is_hip_runtime(): + assert layout is KVLayout.V4 and freqs_cis is None, ( + "the V4.1 KV layouts are CUDA (sm100) only" + ) from sglang.kernels.ops.kvcache.triton_store_cache import ( triton_fused_store_cache, ) @@ -80,8 +112,15 @@ def fused_store_cache( input_dtype=input.dtype, index_dtype=indices.dtype, page_size=page_size, + layout=layout, ) - module.run(input, cache, indices) + if freqs_cis is None: + module.run(input, cache, indices) + else: + assert layout is not KVLayout.V4, "the V4 layout has no in-kernel RoPE" + if freqs_cis.is_complex(): + freqs_cis = torch.view_as_real(freqs_cis).flatten(-2) + module.run_rope(input, cache, indices, freqs_cis.contiguous()) @triton.jit diff --git a/python/sglang/kernels/ops/attention/dsv4/compress.py b/python/sglang/kernels/ops/attention/dsv4/compress.py index becf52f07d87..fb21173a937c 100644 --- a/python/sglang/kernels/ops/attention/dsv4/compress.py +++ b/python/sglang/kernels/ops/attention/dsv4/compress.py @@ -16,6 +16,7 @@ ) from sglang.srt.utils import is_hip, is_xpu +from .kv_layout import KVLayout from .utils import make_name _is_xpu = is_xpu() @@ -48,7 +49,8 @@ def _jit_compress_norm_rope_module( head_dim: int, rope_dim: int, page_size: int, - bf16_store: bool = False, + bf16_store: bool, + layout: KVLayout, fp8_2buff: bool = False, ) -> Module: args = make_cpp_args( @@ -56,9 +58,10 @@ def _jit_compress_norm_rope_module( head_dim, rope_dim, page_size, - is_arch_support_pdl(), INDEXER_K_CACHE_PRESHUFFLE_TILE if aiter_can_use_preshuffle_paged_mqa() else 0, bf16_store, + layout.cpp_name, + is_arch_support_pdl(), ) cuda_wrappers = [("forward", f"FusedNormRopeKernel<{args}>::forward")] if head_dim == 128: @@ -455,9 +458,18 @@ def compress_norm_rope_store( kvcache_scale: Optional[torch.Tensor] = None, rope_cache: Optional[tuple[torch.Tensor, torch.Tensor]] = None, fp4_k_write_metadata=None, + # Page layout of a FlashMLA (head_dim 512) main-KV cache: the 584-byte V4 + # layout, or the V4.1 fp8 / fp4 formats (CUDA only). + layout: Union[KVLayout, str] = KVLayout.V4, fp8_2buff: bool = False, kvcache_rope: Optional[torch.Tensor] = None, ) -> None: + layout = KVLayout.parse(layout) + if layout is not KVLayout.V4: + assert kv.shape[-1] == 512 and not use_fp4 and not bf16_store, ( + "the V4.1 layouts are paged FlashMLA main-KV caches" + ) + assert not is_hip() and not _is_xpu, "the V4.1 KV layouts are CUDA (sm100) only" if use_fp4: assert kv.shape[-1] == 128 if is_hip() and use_fp4: @@ -482,6 +494,7 @@ def compress_norm_rope_store( if fp8_2buff: assert not (use_fp4 or bf16_store), "fp8 two-pool store is its own layout" + assert layout is KVLayout.V4, "fp8 two-pool store is a V4 (584 B page) cache" assert kv.shape[-1] != 128, "fp8 two-pool store is the latent, not the indexer" assert kvcache_rope is not None, "fp8 two-pool store needs the rope pool" assert not _is_xpu, "fp8 two-pool store is only wired for the CUDA/HIP kernel" @@ -507,6 +520,7 @@ def compress_norm_rope_store( freq_cis.shape[-1], page_size, bf16_store, + layout, fp8_2buff, ) if use_fp4: diff --git a/python/sglang/kernels/ops/attention/dsv4/dequant_k_cache.py b/python/sglang/kernels/ops/attention/dsv4/dequant_k_cache.py index c689de8f7e10..3c0cc8e00838 100644 --- a/python/sglang/kernels/ops/attention/dsv4/dequant_k_cache.py +++ b/python/sglang/kernels/ops/attention/dsv4/dequant_k_cache.py @@ -1,9 +1,10 @@ -from typing import Optional +from typing import Optional, Union import torch import triton import triton.language as tl +from sglang.kernels.ops.attention.dsv4.kv_layout import KVLayout from sglang.kernels.ops.quantization.fp8_kernel import is_fp8_fnuz fp8_dtype = torch.float8_e4m3fnuz if is_fp8_fnuz() else torch.float8_e4m3fn @@ -26,6 +27,7 @@ def dequantize_k_cache_paged( page_table_1_flattened: torch.Tensor, page_size: int, out: Optional[torch.Tensor] = None, + layout: Union[KVLayout, str] = KVLayout.V4, ) -> torch.Tensor: """Dequantize the DeepSeek v4 paged KV cache for a list of token IDs. @@ -36,10 +38,17 @@ def dequantize_k_cache_paged( out: optional (num_tokens, 1, DIM_NOPE + DIM_ROPE) bf16 destination. May be a slice of a larger workspace; the kernel uses out.stride(0) so contiguous-along-dim-0 slices work. + layout: the cache's :class:`KVLayout`; the V4.1 layouts (528-byte fp8, + 288-byte fp4) go through :func:`dequantize_k_cache_paged_v41`. Returns: (num_tokens, 1, DIM_NOPE + DIM_ROPE) bfloat16. """ + layout = KVLayout.parse(layout) + if layout is not KVLayout.V4: + return dequantize_k_cache_paged_v41( + quant_k_cache, page_table_1_flattened, page_size, out=out, layout=layout + ) assert quant_k_cache.is_contiguous() assert page_table_1_flattened.dtype in (torch.int32, torch.int64) @@ -85,6 +94,67 @@ def dequantize_k_cache_paged( return out +def dequantize_k_cache_paged_v41( + quant_k_cache: torch.Tensor, + page_table_1_flattened: torch.Tensor, + page_size: int, + out: Optional[torch.Tensor] = None, + layout: KVLayout = KVLayout.V41, +) -> torch.Tensor: + """Dequantize a V4.1 paged cache (fp8 ``V41`` or fp4 ``V41_FP4``) for a list + of token IDs into ``(num_tokens, 1, 512)`` bf16. + + Bit-exact with the pure-torch dequantizer of the formats: the fp8 value + times its power-of-two ue8m0 scale, or the e2m1 value times its e4m3 scale + (at most 2 + 4 significant bits, so exact), rounded to bf16 once. + """ + layout = KVLayout.parse(layout) + assert layout in (KVLayout.V41, KVLayout.V41_FP4), layout + assert quant_k_cache.is_contiguous() + assert page_table_1_flattened.dtype in (torch.int32, torch.int64) + + quant_k_cache_u8 = quant_k_cache.view(torch.uint8) + num_tokens = page_table_1_flattened.shape[0] + bytes_per_page = quant_k_cache_u8.shape[-1] + assert bytes_per_page >= page_size * layout.bytes_per_token, ( + f"{bytes_per_page=} cannot hold {page_size} tokens of {layout}" + ) + buf_fp8 = quant_k_cache_u8.view(fp8_dtype).reshape(-1) + buf_uint8 = quant_k_cache_u8.reshape(-1) + + if out is None: + out = torch.empty( + (num_tokens, 1, DIM_NOPE + DIM_ROPE), + dtype=torch.bfloat16, + device=quant_k_cache.device, + ) + else: + assert out.shape == (num_tokens, 1, DIM_NOPE + DIM_ROPE) + assert out.dtype == torch.bfloat16 + if num_tokens == 0: + return out + + kernel = ( + _dequantize_k_cache_paged_v41_fp8_kernel + if layout is KVLayout.V41 + else _dequantize_k_cache_paged_v41_fp4_kernel + ) + kernel[(num_tokens,)]( + out, + buf_fp8, + buf_uint8, + page_table_1_flattened, + out.stride(0), + BYTES_PER_PAGE=bytes_per_page, + PAGE_SIZE=page_size, + DATA_BYTES=layout.data_bytes, + SCALE_BYTES=layout.scale_bytes, + TILE_SIZE=layout.tile_size, + S_OFFSET_BYTES=layout.scale_offset(page_size), + ) + return out + + def gather_dequant_requant_fp8_paged( quant_k_cache: torch.Tensor, page_table_1_flattened: torch.Tensor, @@ -270,6 +340,104 @@ def _dequantize_k_cache_paged_kernel( tl.store(output_ptr + out_row_base + DIM_NOPE + rope_offs, rope_data) +@triton.jit +def _ue8m0_to_fp32(scale_u8): + """The ue8m0 byte as fp32: ``2 ** (byte - 127)``, built from the exponent + bits so it is exact; byte 0 is the denormal ``2 ** -127`` and byte 255 NaN, + as ``torch.float8_e8m0fnu`` converts them.""" + normal = (scale_u8.to(tl.int32) << 23).to(tl.float32, bitcast=True) + denormal = tl.full(scale_u8.shape, 0x00400000, tl.int32).to( + tl.float32, bitcast=True + ) + scale = tl.where(scale_u8 == 0, denormal, normal) + return tl.where(scale_u8 == 255, float("nan"), scale) + + +@triton.jit +def _e2m1_code_to_fp32(code): + """The 4-bit e2m1 code (bit 3 sign, bits 0-2 index into + ``[0, 0.5, 1, 1.5, 2, 3, 4, 6]``) as fp32.""" + m = code & 7 + e = m >> 1 + f = (m & 1).to(tl.float32) + mag = tl.where(e == 0, 0.5 * f, tl.exp2((e - 1).to(tl.float32)) * (1.0 + 0.5 * f)) + # Set the sign bit directly: a negated zero must stay -0.0 (code 0x8). + sign = (code & 8).to(tl.int32) << 28 + return (mag.to(tl.int32, bitcast=True) | sign).to(tl.float32, bitcast=True) + + +@triton.jit +def _dequantize_k_cache_paged_v41_fp8_kernel( + output_ptr, + buf_fp8_ptr, + buf_uint8_ptr, + page_table_ptr, + output_stride_0, + BYTES_PER_PAGE: tl.constexpr, + PAGE_SIZE: tl.constexpr, + DATA_BYTES: tl.constexpr, + SCALE_BYTES: tl.constexpr, + TILE_SIZE: tl.constexpr, + S_OFFSET_BYTES: tl.constexpr, +): + # V41: 512 e4m3 values per token, then 16 ue8m0 scales (one per 32 values). + tl.static_assert(DATA_BYTES == 512 and SCALE_BYTES == 16 and TILE_SIZE == 32) + token_id = tl.program_id(0).to(tl.int64) + loc = tl.load(page_table_ptr + token_id).to(tl.int64) + page_idx = loc // PAGE_SIZE + in_page = loc % PAGE_SIZE + page_byte_base = page_idx * BYTES_PER_PAGE + token_data_base = page_byte_base + in_page * DATA_BYTES + token_scale_base = page_byte_base + S_OFFSET_BYTES + in_page * SCALE_BYTES + + offs = tl.arange(0, DATA_BYTES) + vals = tl.load(buf_fp8_ptr + token_data_base + offs).to(tl.float32) + scale_u8 = tl.load(buf_uint8_ptr + token_scale_base + offs // TILE_SIZE) + out = vals * _ue8m0_to_fp32(scale_u8) + tl.store( + output_ptr + token_id * output_stride_0 + offs, + out.to(output_ptr.dtype.element_ty), + ) + + +@triton.jit +def _dequantize_k_cache_paged_v41_fp4_kernel( + output_ptr, + buf_fp8_ptr, + buf_uint8_ptr, + page_table_ptr, + output_stride_0, + BYTES_PER_PAGE: tl.constexpr, + PAGE_SIZE: tl.constexpr, + DATA_BYTES: tl.constexpr, + SCALE_BYTES: tl.constexpr, + TILE_SIZE: tl.constexpr, + S_OFFSET_BYTES: tl.constexpr, +): + # V41_FP4: 512 e2m1 codes packed two per byte (even index in the low nibble), + # then 32 e4m3 scales (one per 16 values). + tl.static_assert(DATA_BYTES == 256 and SCALE_BYTES == 32 and TILE_SIZE == 16) + token_id = tl.program_id(0).to(tl.int64) + loc = tl.load(page_table_ptr + token_id).to(tl.int64) + page_idx = loc // PAGE_SIZE + in_page = loc % PAGE_SIZE + page_byte_base = page_idx * BYTES_PER_PAGE + token_data_base = page_byte_base + in_page * DATA_BYTES + token_scale_base = page_byte_base + S_OFFSET_BYTES + in_page * SCALE_BYTES + + boffs = tl.arange(0, DATA_BYTES) + packed = tl.load(buf_uint8_ptr + token_data_base + boffs) + # Byte j holds elements 2j (low nibble) and 2j + 1, both in tile (2j) // 16. + scale = tl.load(buf_fp8_ptr + token_scale_base + (2 * boffs) // TILE_SIZE).to( + tl.float32 + ) + lo = _e2m1_code_to_fp32(packed & 0xF) * scale + hi = _e2m1_code_to_fp32(packed >> 4) * scale + out_base = output_ptr + token_id * output_stride_0 + tl.store(out_base + 2 * boffs, lo.to(output_ptr.dtype.element_ty)) + tl.store(out_base + 2 * boffs + 1, hi.to(output_ptr.dtype.element_ty)) + + @triton.jit def _gather_dequant_requant_fp8_paged_kernel( output_ptr, diff --git a/python/sglang/kernels/ops/attention/dsv4/elementwise.py b/python/sglang/kernels/ops/attention/dsv4/elementwise.py index 053eb01c6cb2..426bae742d33 100644 --- a/python/sglang/kernels/ops/attention/dsv4/elementwise.py +++ b/python/sglang/kernels/ops/attention/dsv4/elementwise.py @@ -1,4 +1,4 @@ -from typing import Optional, Tuple +from typing import Optional, Tuple, Union import torch @@ -10,6 +10,7 @@ ) from sglang.srt.utils import is_hip, is_xpu +from .kv_layout import KVLayout from .utils import make_name _is_hip = is_hip() @@ -55,9 +56,12 @@ def _jit_main_k_norm_rope_flashmla_module( head_dim: int, rope_dim: int, page_size: int, + layout: KVLayout, ): """Main MLA path K kernel: rmsnorm + RoPE + write to FlashMLA paged cache.""" - args = make_cpp_args(dtype, head_dim, rope_dim, page_size, is_arch_support_pdl()) + args = make_cpp_args( + dtype, head_dim, rope_dim, page_size, layout.cpp_name, is_arch_support_pdl() + ) return load_jit( make_name("main_k_norm_rope_flashmla"), *args, @@ -273,16 +277,23 @@ def fused_k_norm_rope_flashmla( out_loc: torch.Tensor, kvcache: torch.Tensor, page_size: int, + layout: Union[KVLayout, str] = KVLayout.V4, ) -> None: + """RMSNorm + RoPE ``kv`` and write it into the paged FlashMLA cache at ``out_loc``. + + ``layout`` selects the page format: the 584-byte V4 layout, or the V4.1 fp8 + (528 B) / fp4 (288 B) formats, in which every dim is quantized.""" + layout = KVLayout.parse(layout) freqs_real = torch.view_as_real(freqs_cis).flatten(-2) head_dim = kv.shape[-1] rope_dim = freqs_real.shape[-1] if _is_xpu: + assert layout is KVLayout.V4, "the V4.1 KV layouts are CUDA (sm100) only" fused_k_norm_rope_flashmla_xpu( kv, kv_weight, freqs_real, positions, out_loc, kvcache, eps, page_size ) else: module = _jit_main_k_norm_rope_flashmla_module( - kv.dtype, head_dim, rope_dim, page_size + kv.dtype, head_dim, rope_dim, page_size, layout ) module.forward(kv, kv_weight, freqs_real, positions, out_loc, kvcache, eps) diff --git a/python/sglang/srt/layers/attention/dsv4/torch_quant.py b/python/sglang/srt/layers/attention/dsv4/torch_quant.py new file mode 100644 index 000000000000..b8cb0e4dda3e --- /dev/null +++ b/python/sglang/srt/layers/attention/dsv4/torch_quant.py @@ -0,0 +1,194 @@ +"""Pure-torch FP4 fake quantization for DeepSeek-V4.1. + +Indexer values use per-32 UE8M0 scales; compressed KV uses per-16 E4M3 scales. +Both paths round to the E2M1 grid with ties to even. +""" + +from typing import Optional + +import torch + +FP8_MAX = 448.0 +FP4_MAX = 6.0 +FP8_BLOCK_SIZE = 32 +FP4_BLOCK_SIZE = 32 +FP4_AMAX_FLOOR = 6 * 2.0**-126 + + +def ceil_pow2(x: torch.Tensor) -> torch.Tensor: + """2 ** ceil(log2(x)) for positive fp32 x, computed on the IEEE bits so the + result is exact at powers of two.""" + bits = x.contiguous().view(torch.int32) + exponent = ((bits >> 23) & 0xFF) - 127 + has_mantissa = (bits & 0x7FFFFF) != 0 + exponent = exponent + has_mantissa.to(torch.int32) + return ((exponent + 127) << 23).view(torch.float32) + + +def block_scale(x: torch.Tensor, block_size: int, fmax: float, amax_floor: float): + """Per-block ue8m0 scale, as fp32 powers of two, shape [..., N // block_size].""" + amax = x.float().unflatten(-1, (-1, block_size)).abs().amax(dim=-1) + amax = amax.clamp_min(amax_floor) + # The kernel multiplies by the fp32 reciprocal rather than dividing. A Python + # scalar keeps this free of host tensors, so it can run under CUDA graph capture. + return ceil_pow2(amax * (1.0 / fmax)) + + +def round_fp4(x: torch.Tensor) -> torch.Tensor: + """Round fp32 values in [-6, 6] onto the e2m1 grid with round-to-nearest-even.""" + magnitude = x.abs() + step = torch.where(magnitude < 2.0, 0.5, torch.where(magnitude < 4.0, 1.0, 2.0)) + return torch.round(magnitude / step) * step * torch.sign(x) + + +def fake_quant_fp4(x: torch.Tensor, block_size: int = FP4_BLOCK_SIZE) -> torch.Tensor: + """Quantize to fp4 (per-block ue8m0 scale) and back, in x's dtype.""" + scale = block_scale(x, block_size, FP4_MAX, FP4_AMAX_FLOOR) + scaled = x.float().unflatten(-1, (-1, block_size)) / scale.unsqueeze(-1) + deq = round_fp4(scaled.clamp(-FP4_MAX, FP4_MAX)) * scale.unsqueeze(-1) + return deq.flatten(-2).to(x.dtype) + + +def fake_quant_compressed_kv(x: torch.Tensor) -> torch.Tensor: + """FP4 round-trip with one E4M3FN scale per 16 compressed-KV elements. + + Round amax / 6 to E4M3 with ties to even, clamping the scale to its + positive finite range [2**-9, 448]. Zero blocks remain zero. + """ + blocks = x.float().unflatten(-1, (-1, 16)) + amax = blocks.abs().amax(dim=-1, keepdim=True) + scale = (amax * (1.0 / FP4_MAX)).clamp(min=2**-9, max=FP8_MAX) + scale = scale.to(torch.float8_e4m3fn).float() + scaled = (blocks / scale).clamp(-FP4_MAX, FP4_MAX) + deq = round_fp4(scaled) * scale + return deq.flatten(-2).to(x.dtype) + + +# --------------------------------------------------------------------------- +# Pure-torch references of the paged V4.1 KV cache formats read by the sparse +# decode kernel (528 B/token fp8 "V41", 288 B/token fp4 "V41_FP4"). They follow +# the kernel's own reference quantizer and are what the store / dequant kernels +# and the tests are checked against, byte for byte. +# --------------------------------------------------------------------------- + +_E2M1_MAGNITUDES = (0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0) + + +def cast_scale_inv_to_ue8m0(scale_inv: torch.Tensor) -> torch.Tensor: + """``2 ** ceil(log2(max(scale_inv, 1e-4)))`` as fp32, computed on the IEEE + bits so that it is exact at (and just above) powers of two.""" + scale_inv = scale_inv.float() + scale = ceil_pow2(torch.clamp_min(scale_inv, 1e-4)) + # ceil_pow2 works on the bits of a finite value; a NaN or inf amax passes + # through (both become the ue8m0 NaN byte, but only the NaN one turns the + # whole tile's payload into NaN). + return torch.where(torch.isfinite(scale_inv), scale, scale_inv) + + +def quantize_to_e2m1_codes(x: torch.Tensor) -> torch.Tensor: + """Round to the nearest e2m1 value with the semantics of + ``cvt.rn.satfinite.e2m1x2.f32`` (ties to even, saturating to +-6) and return + the 4-bit codes as uint8. The sign is kept for values that round to zero + (``-0.0`` and small negatives give the code ``0x8``); NaN maps to code 0.""" + x = x.float() + mags = torch.tensor(_E2M1_MAGNITUDES, dtype=torch.float32, device=x.device) + sign = torch.signbit(x).to(torch.uint8) << 3 + a = torch.nan_to_num(x.abs(), nan=0.0, posinf=6.0).clamp_max(6.0) + mids = (mags[:-1] + mags[1:]) / 2 + code = torch.bucketize(a, mids, right=True) + on_tie = (a.unsqueeze(-1) == mids).any(dim=-1) + tie_code = torch.bucketize(a, mids, right=False) + code = torch.where(on_tie, tie_code + (tie_code & 1), code) + return sign | code.to(torch.uint8) + + +def dequantize_e2m1_codes(codes: torch.Tensor) -> torch.Tensor: + mags = torch.tensor(_E2M1_MAGNITUDES, dtype=torch.float32, device=codes.device) + val = mags[(codes & 7).long()] + return torch.where((codes & 8) != 0, -val, val) + + +def quantize_k_cache_v41( + k: torch.Tensor, page_bytes: Optional[int] = None +) -> torch.Tensor: + """``k`` ``[num_pages, page_size, 512]`` -> uint8 ``[num_pages, page_bytes]`` + pages of the V41 layout: 512 e4m3 per token, then 16 ue8m0 scales per token + (one per 32 values), ``scale = 2 ** ceil(log2(max(amax / 448, 1e-4)))``.""" + num_pages, page_size, d = k.shape + assert d == 512 + x = k.float().view(num_pages, page_size, 16, 32) + scale = cast_scale_inv_to_ue8m0(x.abs().amax(dim=-1) / 448.0) + data = (x / scale.unsqueeze(-1)).to(torch.float8_e4m3fn).view(torch.uint8) + scale_u8 = scale.to(torch.float8_e8m0fnu).view(torch.uint8) + raw = page_size * 528 + if page_bytes is None: + page_bytes = -(-raw // 512) * 512 + assert page_bytes >= raw + out = torch.zeros((num_pages, page_bytes), dtype=torch.uint8, device=k.device) + out[:, : page_size * 512] = data.reshape(num_pages, page_size * 512) + out[:, page_size * 512 : raw] = scale_u8.reshape(num_pages, page_size * 16) + return out + + +def dequantize_k_cache_v41(pages: torch.Tensor, page_size: int) -> torch.Tensor: + """Inverse of :func:`quantize_k_cache_v41`: ``[num_pages, page_size, 512]`` bf16.""" + num_pages = pages.shape[0] + pages = pages.view(torch.uint8) + data = pages[:, : page_size * 512].reshape(num_pages, page_size, 512) + scale = pages[:, page_size * 512 : page_size * 528].reshape( + num_pages, page_size, 16 + ) + values = data.view(torch.float8_e4m3fn).to(torch.bfloat16) + scale_bf16 = scale.view(torch.float8_e8m0fnu).to(torch.bfloat16) + return (values.view(num_pages, page_size, 16, 32) * scale_bf16.unsqueeze(-1)).view( + num_pages, page_size, 512 + ) + + +def quantize_k_cache_v41_fp4( + k: torch.Tensor, page_bytes: Optional[int] = None +) -> torch.Tensor: + """``k`` ``[num_pages, page_size, 512]`` -> uint8 ``[num_pages, page_bytes]`` + pages of the V41_FP4 layout: 256 B of e2m1 codes per token (even index in + the low nibble), then 32 e4m3 scales per token (one per 16 values), + ``scale = e4m3(clamp(amax / 6, 2**-9, 448))``. A NaN element poisons its + tile: NaN scale, zero codes.""" + num_pages, page_size, d = k.shape + assert d == 512 + x = k.float().view(num_pages, page_size, 32, 16) + amax = torch.nan_to_num(x.abs(), nan=float("inf")).amax(dim=-1) + scale = torch.clamp(amax / 6.0, 2.0**-9, 448.0).to(torch.float8_e4m3fn) + scale = torch.where(torch.isinf(amax), torch.full_like(scale, float("nan")), scale) + codes = quantize_to_e2m1_codes(x / scale.float().unsqueeze(-1)) + codes = codes.view(num_pages, page_size, 512) + packed = codes[..., 0::2] | (codes[..., 1::2] << 4) + raw = page_size * 288 + if page_bytes is None: + page_bytes = -(-raw // 256) * 256 + assert page_bytes >= raw + out = torch.zeros((num_pages, page_bytes), dtype=torch.uint8, device=k.device) + out[:, : page_size * 256] = packed.reshape(num_pages, page_size * 256) + out[:, page_size * 256 : raw] = scale.view(torch.uint8).reshape( + num_pages, page_size * 32 + ) + return out + + +def dequantize_k_cache_v41_fp4(pages: torch.Tensor, page_size: int) -> torch.Tensor: + """Inverse of :func:`quantize_k_cache_v41_fp4`: ``[num_pages, page_size, 512]`` + bf16. ``e2m1 * e4m3`` has at most 2 + 4 significant bits, so the product is + exact in bf16, as in the kernel.""" + num_pages = pages.shape[0] + pages = pages.view(torch.uint8) + data = pages[:, : page_size * 256].reshape(num_pages, page_size, 256) + scale = pages[:, page_size * 256 : page_size * 288].reshape( + num_pages, page_size, 32 + ) + codes = torch.empty( + (num_pages, page_size, 512), dtype=torch.uint8, device=pages.device + ) + codes[..., 0::2] = data & 0xF + codes[..., 1::2] = data >> 4 + values = dequantize_e2m1_codes(codes).view(num_pages, page_size, 32, 16) + out = values * scale.view(torch.float8_e4m3fn).float().unsqueeze(-1) + return out.view(num_pages, page_size, 512).to(torch.bfloat16) diff --git a/test/registered/kernel/attention/dsv4/test_v41_kv_dequant.py b/test/registered/kernel/attention/dsv4/test_v41_kv_dequant.py new file mode 100644 index 000000000000..17049308b06d --- /dev/null +++ b/test/registered/kernel/attention/dsv4/test_v41_kv_dequant.py @@ -0,0 +1,132 @@ +"""The V4.1 paged dequant (bf16 prefill workspace) is bit-exact with the +pure-torch dequantizer of the fp8 (V41) and fp4 (V41_FP4) formats.""" + +import unittest + +import torch + +from sglang.kernels.ops.attention.dsv4.dequant_k_cache import dequantize_k_cache_paged +from sglang.kernels.ops.attention.dsv4.kv_layout import KVLayout +from sglang.srt.layers.attention.dsv4 import torch_quant as tq +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci(est_time=60, stage="base-b-kernel-unit", runner_config="1-gpu-large") +register_cuda_ci(est_time=60, stage="base-b-kernel-unit", runner_config="4-gpu-b200") + +CASES = { + KVLayout.V41: (tq.quantize_k_cache_v41, tq.dequantize_k_cache_v41), + KVLayout.V41_FP4: (tq.quantize_k_cache_v41_fp4, tq.dequantize_k_cache_v41_fp4), +} + + +def bits(t: torch.Tensor) -> torch.Tensor: + """bf16 as int16, so that -0.0 and NaN payloads compare exactly.""" + return t.contiguous().view(torch.int16) + + +@unittest.skipUnless(torch.cuda.is_available(), "requires CUDA") +class TestV41KVDequant(CustomTestCase): + def _gather_ref(self, dequant, pages, page_size, ids): + return dequant(pages, page_size).view(-1, 512)[ids.long()].unsqueeze(1) + + def test_quantized_pages(self): + g = torch.Generator(device="cuda").manual_seed(0) + for layout, (quant, dequant) in CASES.items(): + for page_size, num_pages in ((64, 9), (256, 3), (2, 50)): + with self.subTest(layout=layout.name, page_size=page_size): + k = torch.randn( + num_pages, + page_size, + 512, + generator=g, + device="cuda", + dtype=torch.bfloat16, + ) + k = ( + k + * torch.exp2( + torch.randint( + -12, + 6, + (num_pages, page_size, 1), + generator=g, + device="cuda", + ).float() + ) + ).to(torch.bfloat16) + k[0, 0, :32] = 0 + k[0, 0, 32:48] = -0.0 + pages = quant(k, page_bytes=layout.page_bytes(page_size)) + ids = torch.randint( + 0, + num_pages * page_size, + (777,), + generator=g, + device="cuda", + dtype=torch.int32, + ) + got = dequantize_k_cache_paged(pages, ids, page_size, layout=layout) + self.assertEqual(got.shape, (777, 1, 512)) + self.assertTrue( + torch.equal( + bits(got), + bits(self._gather_ref(dequant, pages, page_size, ids)), + ) + ) + # The fp4 cache dequantizes to the model's fake-quantized value + # (compared by value: the fake quant maps an exact -0.0 to +0.0). + if layout is KVLayout.V41_FP4: + expect = tq.fake_quant_compressed_kv( + k.view(-1, 512)[ids.long()] + ).unsqueeze(1) + self.assertTrue(torch.equal(got, expect)) + + def test_random_bytes_and_workspace_slice(self): + """Arbitrary payload bytes (scales in the quantizer's range) and an + ``out`` that is a strided slice of a larger workspace.""" + g = torch.Generator(device="cuda").manual_seed(1) + for layout, (_, dequant) in CASES.items(): + page_size, num_pages = 64, 7 + with self.subTest(layout=layout.name): + pages = torch.randint( + 0, + 256, + (num_pages, layout.page_bytes(page_size)), + generator=g, + dtype=torch.uint8, + device="cuda", + ) + if layout is KVLayout.V41: + lo = layout.scale_offset(page_size) + hi = lo + page_size * layout.scale_bytes + pages[:, lo:hi] = torch.randint( + 100, + 140, + (num_pages, hi - lo), + generator=g, + dtype=torch.uint8, + device="cuda", + ) + ids = torch.randint( + 0, + num_pages * page_size, + (300,), + generator=g, + device="cuda", + dtype=torch.int64, + ) + ref = self._gather_ref(dequant, pages, page_size, ids) + workspace = torch.zeros( + 305, 1, 512, dtype=torch.bfloat16, device="cuda" + ) + out = dequantize_k_cache_paged( + pages, ids, page_size, out=workspace[5:], layout=layout + ) + # NaN payloads (fp8 0x7F / e4m3 NaN scales) compare through their bits. + self.assertTrue(torch.equal(bits(workspace[5:]), bits(ref))) + self.assertEqual(int(workspace[:5].abs().sum()), 0) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/kernel/attention/dsv4/test_v41_kv_store.py b/test/registered/kernel/attention/dsv4/test_v41_kv_store.py index 638ee2a99df0..12a3b3bc3014 100644 --- a/test/registered/kernel/attention/dsv4/test_v41_kv_store.py +++ b/test/registered/kernel/attention/dsv4/test_v41_kv_store.py @@ -10,200 +10,22 @@ import torch from sglang.kernels.ops.attention.dsv4.kv_layout import KVLayout +from sglang.srt.layers.attention.dsv4 import torch_quant as tq from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.test_utils import CustomTestCase register_cuda_ci(est_time=120, stage="base-b-kernel-unit", runner_config="4-gpu-b200") - -from typing import Optional - -import torch - -FP8_MAX = 448.0 -FP4_MAX = 6.0 -FP8_BLOCK_SIZE = 32 -FP4_BLOCK_SIZE = 32 -FP4_AMAX_FLOOR = 6 * 2.0**-126 - - -def ceil_pow2(x: torch.Tensor) -> torch.Tensor: - """2 ** ceil(log2(x)) for positive fp32 x, computed on the IEEE bits so the - result is exact at powers of two.""" - bits = x.contiguous().view(torch.int32) - exponent = ((bits >> 23) & 0xFF) - 127 - has_mantissa = (bits & 0x7FFFFF) != 0 - exponent = exponent + has_mantissa.to(torch.int32) - return ((exponent + 127) << 23).view(torch.float32) - - -def block_scale(x: torch.Tensor, block_size: int, fmax: float, amax_floor: float): - """Per-block ue8m0 scale, as fp32 powers of two, shape [..., N // block_size].""" - amax = x.float().unflatten(-1, (-1, block_size)).abs().amax(dim=-1) - amax = amax.clamp_min(amax_floor) - # The kernel multiplies by the fp32 reciprocal rather than dividing. A Python - # scalar keeps this free of host tensors, so it can run under CUDA graph capture. - return ceil_pow2(amax * (1.0 / fmax)) - - -def round_fp4(x: torch.Tensor) -> torch.Tensor: - """Round fp32 values in [-6, 6] onto the e2m1 grid with round-to-nearest-even.""" - magnitude = x.abs() - step = torch.where(magnitude < 2.0, 0.5, torch.where(magnitude < 4.0, 1.0, 2.0)) - return torch.round(magnitude / step) * step * torch.sign(x) - - -def fake_quant_fp4(x: torch.Tensor, block_size: int = FP4_BLOCK_SIZE) -> torch.Tensor: - """Quantize to fp4 (per-block ue8m0 scale) and back, in x's dtype.""" - scale = block_scale(x, block_size, FP4_MAX, FP4_AMAX_FLOOR) - scaled = x.float().unflatten(-1, (-1, block_size)) / scale.unsqueeze(-1) - deq = round_fp4(scaled.clamp(-FP4_MAX, FP4_MAX)) * scale.unsqueeze(-1) - return deq.flatten(-2).to(x.dtype) - - -def fake_quant_compressed_kv(x: torch.Tensor) -> torch.Tensor: - """FP4 round-trip with one E4M3FN scale per 16 compressed-KV elements. - - Round amax / 6 to E4M3 with ties to even, clamping the scale to its - positive finite range [2**-9, 448]. Zero blocks remain zero. - """ - blocks = x.float().unflatten(-1, (-1, 16)) - amax = blocks.abs().amax(dim=-1, keepdim=True) - scale = (amax * (1.0 / FP4_MAX)).clamp(min=2**-9, max=FP8_MAX) - scale = scale.to(torch.float8_e4m3fn).float() - scaled = (blocks / scale).clamp(-FP4_MAX, FP4_MAX) - deq = round_fp4(scaled) * scale - return deq.flatten(-2).to(x.dtype) - - -# --------------------------------------------------------------------------- -# Pure-torch references of the paged V4.1 KV cache formats read by the sparse -# decode kernel (528 B/token fp8 "V41", 288 B/token fp4 "V41_FP4"). They follow -# the kernel's own reference quantizer and are what the store / dequant kernels -# and the tests are checked against, byte for byte. -# --------------------------------------------------------------------------- - -_E2M1_MAGNITUDES = (0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0) - - -def cast_scale_inv_to_ue8m0(scale_inv: torch.Tensor) -> torch.Tensor: - """``2 ** ceil(log2(max(scale_inv, 1e-4)))`` as fp32, computed on the IEEE - bits so that it is exact at (and just above) powers of two.""" - scale_inv = scale_inv.float() - scale = ceil_pow2(torch.clamp_min(scale_inv, 1e-4)) - # ceil_pow2 works on the bits of a finite value; a NaN or inf amax passes - # through (both become the ue8m0 NaN byte, but only the NaN one turns the - # whole tile's payload into NaN). - return torch.where(torch.isfinite(scale_inv), scale, scale_inv) - - -def quantize_to_e2m1_codes(x: torch.Tensor) -> torch.Tensor: - """Round to the nearest e2m1 value with the semantics of - ``cvt.rn.satfinite.e2m1x2.f32`` (ties to even, saturating to +-6) and return - the 4-bit codes as uint8. The sign is kept for values that round to zero - (``-0.0`` and small negatives give the code ``0x8``); NaN maps to code 0.""" - x = x.float() - mags = torch.tensor(_E2M1_MAGNITUDES, dtype=torch.float32, device=x.device) - sign = torch.signbit(x).to(torch.uint8) << 3 - a = torch.nan_to_num(x.abs(), nan=0.0, posinf=6.0).clamp_max(6.0) - mids = (mags[:-1] + mags[1:]) / 2 - code = torch.bucketize(a, mids, right=True) - on_tie = (a.unsqueeze(-1) == mids).any(dim=-1) - tie_code = torch.bucketize(a, mids, right=False) - code = torch.where(on_tie, tie_code + (tie_code & 1), code) - return sign | code.to(torch.uint8) - - -def dequantize_e2m1_codes(codes: torch.Tensor) -> torch.Tensor: - mags = torch.tensor(_E2M1_MAGNITUDES, dtype=torch.float32, device=codes.device) - val = mags[(codes & 7).long()] - return torch.where((codes & 8) != 0, -val, val) - - -def quantize_k_cache_v41( - k: torch.Tensor, page_bytes: Optional[int] = None -) -> torch.Tensor: - """``k`` ``[num_pages, page_size, 512]`` -> uint8 ``[num_pages, page_bytes]`` - pages of the V41 layout: 512 e4m3 per token, then 16 ue8m0 scales per token - (one per 32 values), ``scale = 2 ** ceil(log2(max(amax / 448, 1e-4)))``.""" - num_pages, page_size, d = k.shape - assert d == 512 - x = k.float().view(num_pages, page_size, 16, 32) - scale = cast_scale_inv_to_ue8m0(x.abs().amax(dim=-1) / 448.0) - data = (x / scale.unsqueeze(-1)).to(torch.float8_e4m3fn).view(torch.uint8) - scale_u8 = scale.to(torch.float8_e8m0fnu).view(torch.uint8) - raw = page_size * 528 - if page_bytes is None: - page_bytes = -(-raw // 512) * 512 - assert page_bytes >= raw - out = torch.zeros((num_pages, page_bytes), dtype=torch.uint8, device=k.device) - out[:, : page_size * 512] = data.reshape(num_pages, page_size * 512) - out[:, page_size * 512 : raw] = scale_u8.reshape(num_pages, page_size * 16) - return out - - -def dequantize_k_cache_v41(pages: torch.Tensor, page_size: int) -> torch.Tensor: - """Inverse of :func:`quantize_k_cache_v41`: ``[num_pages, page_size, 512]`` bf16.""" - num_pages = pages.shape[0] - pages = pages.view(torch.uint8) - data = pages[:, : page_size * 512].reshape(num_pages, page_size, 512) - scale = pages[:, page_size * 512 : page_size * 528].reshape( - num_pages, page_size, 16 - ) - values = data.view(torch.float8_e4m3fn).to(torch.bfloat16) - scale_bf16 = scale.view(torch.float8_e8m0fnu).to(torch.bfloat16) - return (values.view(num_pages, page_size, 16, 32) * scale_bf16.unsqueeze(-1)).view( - num_pages, page_size, 512 - ) - - -def quantize_k_cache_v41_fp4( - k: torch.Tensor, page_bytes: Optional[int] = None -) -> torch.Tensor: - """``k`` ``[num_pages, page_size, 512]`` -> uint8 ``[num_pages, page_bytes]`` - pages of the V41_FP4 layout: 256 B of e2m1 codes per token (even index in - the low nibble), then 32 e4m3 scales per token (one per 16 values), - ``scale = e4m3(clamp(amax / 6, 2**-9, 448))``. A NaN element poisons its - tile: NaN scale, zero codes.""" - num_pages, page_size, d = k.shape - assert d == 512 - x = k.float().view(num_pages, page_size, 32, 16) - amax = torch.nan_to_num(x.abs(), nan=float("inf")).amax(dim=-1) - scale = torch.clamp(amax / 6.0, 2.0**-9, 448.0).to(torch.float8_e4m3fn) - scale = torch.where(torch.isinf(amax), torch.full_like(scale, float("nan")), scale) - codes = quantize_to_e2m1_codes(x / scale.float().unsqueeze(-1)) - codes = codes.view(num_pages, page_size, 512) - packed = codes[..., 0::2] | (codes[..., 1::2] << 4) - raw = page_size * 288 - if page_bytes is None: - page_bytes = -(-raw // 256) * 256 - assert page_bytes >= raw - out = torch.zeros((num_pages, page_bytes), dtype=torch.uint8, device=k.device) - out[:, : page_size * 256] = packed.reshape(num_pages, page_size * 256) - out[:, page_size * 256 : raw] = scale.view(torch.uint8).reshape( - num_pages, page_size * 32 - ) - return out - - -def dequantize_k_cache_v41_fp4(pages: torch.Tensor, page_size: int) -> torch.Tensor: - """Inverse of :func:`quantize_k_cache_v41_fp4`: ``[num_pages, page_size, 512]`` - bf16. ``e2m1 * e4m3`` has at most 2 + 4 significant bits, so the product is - exact in bf16, as in the kernel.""" - num_pages = pages.shape[0] - pages = pages.view(torch.uint8) - data = pages[:, : page_size * 256].reshape(num_pages, page_size, 256) - scale = pages[:, page_size * 256 : page_size * 288].reshape( - num_pages, page_size, 32 - ) - codes = torch.empty( - (num_pages, page_size, 512), dtype=torch.uint8, device=pages.device - ) - codes[..., 0::2] = data & 0xF - codes[..., 1::2] = data >> 4 - values = dequantize_e2m1_codes(codes).view(num_pages, page_size, 32, 16) - out = values * scale.view(torch.float8_e4m3fn).float().unsqueeze(-1) - return out.view(num_pages, page_size, 512).to(torch.bfloat16) +REFERENCE = { + KVLayout.V41: tq.quantize_k_cache_v41, + KVLayout.V41_FP4: tq.quantize_k_cache_v41_fp4, +} +DEQUANT = { + KVLayout.V41: tq.dequantize_k_cache_v41, + KVLayout.V41_FP4: tq.dequantize_k_cache_v41_fp4, +} +# One quantization step, relative: e4m3 has 3 mantissa bits, e2m1 one. +ONE_CODE_RTOL = {KVLayout.V41: 0.13, KVLayout.V41_FP4: 0.51} def rope_tail( @@ -218,12 +40,6 @@ def rope_tail( return torch.cat([head, rotated], dim=-1) -REFERENCE = { - KVLayout.V41: quantize_k_cache_v41, - KVLayout.V41_FP4: quantize_k_cache_v41_fp4, -} - - def _sm100(): return ( torch.cuda.is_available() @@ -248,6 +64,20 @@ def token_rows(pages, layout, page_size, locs): return data, scale +def random_rows(n, generator, device="cuda"): + """bf16 rows over a wide dynamic range, with zero, negative-zero and tiny tiles.""" + x = torch.randn(n, 512, generator=generator, device=device, dtype=torch.bfloat16) + scale = torch.exp2( + torch.randint(-10, 6, (n, 1), generator=generator, device=device).float() + ) + x = (x * scale).to(torch.bfloat16) + if n >= 3: + x[0, :32] = 0 + x[1, 32:48] = -0.0 + x[2, 100] = -1e-10 + return x + + def reference_pages(layout, page_size, num_pages, locs, values, page_bytes): full = torch.zeros( num_pages, page_size, 512, device=values.device, dtype=values.dtype @@ -264,6 +94,21 @@ def assert_tokens_equal(self, cache, ref, layout, page_size, locs): self.assertTrue(torch.equal(got_scale, exp_scale), "scale rows differ") self.assertTrue(torch.equal(got_data, exp_data), "data rows differ") + def assert_rows_close(self, cache, ref, layout, page_size, locs): + """For rows that went through the kernel's fp32 RMSNorm: torch sums the + squares in another order, and the fp32 ulp this can cost is occasionally + kept by a bf16 rounding boundary and then by the quantizer. Allow a + one-code difference in a handful of elements.""" + got_data, got_scale = token_rows(cache, layout, page_size, locs) + exp_data, exp_scale = token_rows(ref, layout, page_size, locs) + self.assertGreater((got_scale == exp_scale).float().mean().item(), 0.999) + self.assertGreater((got_data == exp_data).float().mean().item(), 0.999) + deq_got = DEQUANT[layout](cache, page_size).view(-1, 512)[locs.long()].float() + deq_exp = DEQUANT[layout](ref, page_size).view(-1, 512)[locs.long()].float() + torch.testing.assert_close( + deq_got, deq_exp, rtol=ONE_CODE_RTOL[layout], atol=0.06 + ) + def assert_untouched_zero(self, cache, layout, page_size, locs): num_slots = cache.shape[0] * page_size written = torch.zeros(num_slots, dtype=torch.bool, device=cache.device) @@ -272,12 +117,224 @@ def assert_untouched_zero(self, cache, layout, page_size, locs): data, scale = token_rows(cache, layout, page_size, others) self.assertEqual(int(data.sum()) + int(scale.sum()), 0) + def test_fused_store_cache(self): + """Ragged page fills: a random subset of slots is written, the rest stays zero.""" + from sglang.kernels.ops.attention.dsv4.attn import fused_store_cache + + g = torch.Generator(device="cuda").manual_seed(0) + for layout in (KVLayout.V41, KVLayout.V41_FP4): + for page_size, num_pages, n in ((64, 9, 333), (256, 3, 500), (2, 40, 37)): + for idx_dtype in (torch.int32, torch.int64): + with self.subTest( + layout=layout.name, page_size=page_size, idx=idx_dtype + ): + x = random_rows(n, g) + locs = torch.randperm( + num_pages * page_size, generator=g, device="cuda" + )[:n].to(idx_dtype) + cache = torch.zeros( + num_pages, + layout.page_bytes(page_size), + dtype=torch.uint8, + device="cuda", + ) + fused_store_cache( + x, + cache, + locs, + page_size=page_size, + type="flashmla", + layout=layout, + ) + ref = reference_pages( + layout, page_size, num_pages, locs, x, cache.shape[1] + ) + self.assert_tokens_equal(cache, ref, layout, page_size, locs) + self.assert_untouched_zero(cache, layout, page_size, locs) + + def test_fused_store_cache_with_rope(self): + """The in-kernel RoPE tail equals rope_tail (bf16-rounded) before quantizing, + so the fp4 cache holds exactly fake_quant_compressed_kv(rope_tail(x)).""" + from sglang.kernels.ops.attention.dsv4.attn import fused_store_cache + + g = torch.Generator(device="cuda").manual_seed(1) + for layout in (KVLayout.V41, KVLayout.V41_FP4): + for page_size, num_pages, n in ((64, 5, 200), (256, 2, 129)): + with self.subTest(layout=layout.name, page_size=page_size): + x = random_rows(n, g) + angles = torch.randn(n, 32, generator=g, device="cuda") + freqs = torch.polar(torch.ones_like(angles), angles) + locs = torch.randperm( + num_pages * page_size, generator=g, device="cuda" + )[:n] + cache = torch.zeros( + num_pages, + layout.page_bytes(page_size), + dtype=torch.uint8, + device="cuda", + ) + fused_store_cache( + x, + cache, + locs, + page_size=page_size, + type="flashmla", + layout=layout, + freqs_cis=freqs, + ) + rotated = rope_tail(x, freqs, 64) + ref = reference_pages( + layout, page_size, num_pages, locs, rotated, cache.shape[1] + ) + self.assert_tokens_equal(cache, ref, layout, page_size, locs) + if layout is KVLayout.V41_FP4: + deq = tq.dequantize_k_cache_v41_fp4(cache, page_size).view( + -1, 512 + )[locs] + self.assertTrue( + torch.equal(deq, tq.fake_quant_compressed_kv(rotated)) + ) + + def test_boundary_tiles(self): + """Tie, saturation and subnormal-scale tiles follow the reference. (NaN / inf + rows are not fed: the kernels do not reproduce the reference's handling of them.)""" + from sglang.kernels.ops.attention.dsv4.attn import fused_store_cache + + maxima = [ + 0, + 2**-12, + 6 * 2**-9, + 6 * 1.0625, + 6 * 1.1875, + 6 * 448, + 1e6, + 8.25, + 448.0, + 449.0, + ] + rows = [] + for m in maxima: + row = torch.full((512,), m, dtype=torch.bfloat16) + row[1::2] *= -1 + row[64:80] = torch.tensor( + [8.25, 4.125, 3.4375, -4.8125, 0.0, -0.0, 2.5, -3.5] * 2, + dtype=torch.bfloat16, + ) + rows.append(row) + x = torch.stack(rows).cuda() + n = x.shape[0] + page_size = 16 + locs = torch.arange(n, device="cuda", dtype=torch.int32) + for layout in (KVLayout.V41, KVLayout.V41_FP4): + with self.subTest(layout=layout.name): + cache = torch.zeros( + 1, layout.page_bytes(page_size), dtype=torch.uint8, device="cuda" + ) + fused_store_cache( + x, cache, locs, page_size=page_size, type="flashmla", layout=layout + ) + ref = reference_pages(layout, page_size, 1, locs, x, cache.shape[1]) + self.assert_tokens_equal(cache, ref, layout, page_size, locs) + self.assert_untouched_zero(cache, layout, page_size, locs) + + def test_fused_k_norm_rope_store(self): + """The fused RMSNorm + RoPE + store (the SWA write) equals norm -> rope_tail -> + fused_store_cache: exact-norm rows bitwise against the torch quantizer, and + general rows bitwise against the unfused kernel chain.""" + from sglang.kernels.ops.attention.dsv4.attn import fused_store_cache + from sglang.kernels.ops.attention.dsv4.elementwise import ( + fused_k_norm_rope_flashmla, + ) + + g = torch.Generator(device="cuda").manual_seed(2) + page_size, num_pages, n = 256, 3, 300 + angles = torch.randn(1024, 32, generator=g, device="cuda") + freqs_table = torch.polar(torch.ones_like(angles), angles) + pos = torch.randint( + 0, 1024, (n,), generator=g, device="cuda", dtype=torch.int64 + ) + locs = torch.randperm(num_pages * page_size, generator=g, device="cuda")[:n].to( + torch.int32 + ) + locs[7] = -1 # a row without a write target is skipped + valid = locs >= 0 + for layout in (KVLayout.V41, KVLayout.V41_FP4): + with self.subTest(layout=layout.name, rows="exact-norm"): + # x = +-2^k per row with eps = 0 normalizes to +-1 exactly, so the + # normed row is sign * weight and the reference is exact. + signs = torch.where( + torch.rand(n, 512, generator=g, device="cuda") < 0.5, -1.0, 1.0 + ) + k = torch.randint(-6, 6, (n, 1), generator=g, device="cuda").float() + x = (signs * torch.exp2(k)).to(torch.bfloat16) + w = ( + torch.randn(512, generator=g, device="cuda") + * torch.exp2( + torch.randint(-6, 4, (512,), generator=g, device="cuda").float() + ) + ).to(torch.bfloat16) + cache = torch.zeros( + num_pages, + layout.page_bytes(page_size), + dtype=torch.uint8, + device="cuda", + ) + fused_k_norm_rope_flashmla( + x, w, 0.0, freqs_table, pos, locs, cache, page_size, layout=layout + ) + rotated = rope_tail( + (signs * w.float()).to(torch.bfloat16), freqs_table[pos], 64 + ) + ref = reference_pages( + layout, + page_size, + num_pages, + locs[valid], + rotated[valid], + cache.shape[1], + ) + self.assert_tokens_equal(cache, ref, layout, page_size, locs[valid]) + self.assert_untouched_zero(cache, layout, page_size, locs[valid]) + with self.subTest(layout=layout.name, rows="general"): + x = (torch.randn(n, 512, generator=g, device="cuda") * 3).to( + torch.bfloat16 + ) + w = (torch.randn(512, generator=g, device="cuda") * 0.3 + 1).to( + torch.bfloat16 + ) + eps = 1e-6 + cache = torch.zeros( + num_pages, + layout.page_bytes(page_size), + dtype=torch.uint8, + device="cuda", + ) + fused_k_norm_rope_flashmla( + x, w, eps, freqs_table, pos, locs, cache, page_size, layout=layout + ) + xf = x.float() + normed = ( + xf + * torch.rsqrt(xf.square().mean(-1, keepdim=True) + eps) + * w.float() + ).to(torch.bfloat16) + rotated = rope_tail(normed, freqs_table[pos], 64) + unfused = torch.zeros_like(cache) + fused_store_cache( + rotated[valid], + unfused, + locs[valid], + page_size=page_size, + type="flashmla", + layout=layout, + ) + self.assert_rows_close(cache, unfused, layout, page_size, locs[valid]) + def test_c1_c2_decode_store(self): """The ratio-1 / ratio-2 decode compressors write the V4.1 layouts: the cache holds the quantized rope_tail of the pre-RoPE latent the kernel publishes (bitwise; the fp8 layout after the model's fp4 fake quantization), and the latent is the torch RMSNorm to within an fp32-reduction-order bf16 ulp.""" - from sglang.kernels.ops.attention.dsv4.c1 import c1_decode_norm_rope_store from sglang.kernels.ops.attention.dsv4.c2 import ( c2_decode_or_verify_norm_rope_store, @@ -300,7 +357,7 @@ def stored_reference(layout, rotated, page_size, num_pages, locs, page_bytes): values = ( rotated if layout is KVLayout.V41_FP4 - else fake_quant_compressed_kv(rotated) + else tq.fake_quant_compressed_kv(rotated) ) return reference_pages( layout, page_size, num_pages, locs, values, page_bytes @@ -428,6 +485,130 @@ def stored_reference(layout, rotated, page_size, num_pages, locs, page_bytes): ) self.assert_untouched_zero(cache, layout, page_size, slots[valid]) + def test_compress_norm_rope_store(self): + """The ratio-4 / ratio-128 writer (norm + RoPE + store from a decode plan) in + the V4.1 layouts: bitwise on exact-norm rows, one-code close on general rows.""" + from sglang.kernels.ops.attention.dsv4.compress import ( + CompressorDecodePlan, + compress_norm_rope_store, + ) + + g = torch.Generator(device="cuda").manual_seed(5) + ratio = 4 + angles = torch.randn(4096, 32, generator=g, device="cuda") + freqs = torch.polar(torch.ones_like(angles), angles) + n = 150 + for layout in (KVLayout.V41, KVLayout.V41_FP4): + for page_size, num_pages in ((64, 8), (2, 300)): + seq_lens = ( + torch.randint( + 1, 1000, (n,), generator=g, device="cuda", dtype=torch.int64 + ) + * ratio + ) + seq_lens[3] += 1 # not a group boundary: writes nothing + plan = CompressorDecodePlan.generate_legacy( + ratio, torch.arange(n, device="cuda", dtype=torch.int64), seq_lens + ) + valid = seq_lens % ratio == 0 + pos = seq_lens - ratio + out_loc = torch.randperm( + num_pages * page_size, generator=g, device="cuda" + )[:n].to(torch.int64) + with self.subTest( + layout=layout.name, page_size=page_size, rows="exact-norm" + ): + signs = torch.where( + torch.rand(n, 512, generator=g, device="cuda") < 0.5, -1.0, 1.0 + ) + k = torch.randint(-6, 6, (n, 1), generator=g, device="cuda").float() + kv = (signs * torch.exp2(k)).to(torch.bfloat16) + w = ( + torch.randn(512, generator=g, device="cuda") + * torch.exp2( + torch.randint( + -6, 4, (512,), generator=g, device="cuda" + ).float() + ) + ).to(torch.bfloat16) + cache = torch.zeros( + num_pages, + layout.page_bytes(page_size), + dtype=torch.uint8, + device="cuda", + ) + compress_norm_rope_store( + kv, + plan, + norm_weight=w, + norm_eps=0.0, + freq_cis=freqs, + out_loc=out_loc, + kvcache=cache, + page_size=page_size, + layout=layout, + ) + rotated = rope_tail( + (signs * w.float()).to(torch.bfloat16), freqs[pos], 64 + ) + ref = reference_pages( + layout, + page_size, + num_pages, + out_loc[valid], + rotated[valid], + cache.shape[1], + ) + self.assert_tokens_equal( + cache, ref, layout, page_size, out_loc[valid] + ) + self.assert_untouched_zero(cache, layout, page_size, out_loc[valid]) + with self.subTest( + layout=layout.name, page_size=page_size, rows="general" + ): + eps = 1e-6 + kv = (torch.randn(n, 512, generator=g, device="cuda") * 2).to( + torch.bfloat16 + ) + w = (torch.randn(512, generator=g, device="cuda") * 0.3 + 1).to( + torch.bfloat16 + ) + cache = torch.zeros( + num_pages, + layout.page_bytes(page_size), + dtype=torch.uint8, + device="cuda", + ) + compress_norm_rope_store( + kv, + plan, + norm_weight=w, + norm_eps=eps, + freq_cis=freqs, + out_loc=out_loc, + kvcache=cache, + page_size=page_size, + layout=layout, + ) + xf = kv.float() + normed = ( + xf + * torch.rsqrt(xf.square().mean(-1, keepdim=True) + eps) + * w.float() + ).to(torch.bfloat16) + rotated = rope_tail(normed, freqs[pos], 64) + ref = reference_pages( + layout, + page_size, + num_pages, + out_loc[valid], + rotated[valid], + cache.shape[1], + ) + self.assert_rows_close( + cache, ref, layout, page_size, out_loc[valid] + ) + if __name__ == "__main__": unittest.main() From a0ae65a1e7273332053f95b33bc4df06565a4ace Mon Sep 17 00:00:00 2001 From: hnyls2002 Date: Tue, 15 Sep 2026 15:53:41 -0700 Subject: [PATCH 12/30] dsv4.1: preserve 64K prefill planner indices and reject sentinel collisions --- .../kernels/jit/csrc/deepseek_v4/c_plan.cuh | 10 ++-- ...est_deepseek_v4_compress_plan_draft_pad.py | 55 ++++++++++++++++++- 2 files changed, 60 insertions(+), 5 deletions(-) diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/c_plan.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/c_plan.cuh index b157eebea28b..45ff0e203286 100644 --- a/python/sglang/kernels/jit/csrc/deepseek_v4/c_plan.cuh +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/c_plan.cuh @@ -505,10 +505,12 @@ inline PrefillPlan plan_compress_prefill( const auto f2s_ptr = static_cast(full_to_state.data_ptr()); const auto batch_size = static_cast(B.unwrap()); - constexpr auto kMaxTokens = static_cast(std::numeric_limits::max()); + // ragged_id is a zero-based uint16 index, so a 64K-token batch is valid. + constexpr auto kMaxTokens = static_cast(std::numeric_limits::max()) + 1; RuntimeCheck(compress_ratio == 4 || compress_ratio == 128); RuntimeCheck(!use_req_ring || compress_ratio == 4); - RuntimeCheck(batch_size <= num_q_tokens && num_q_tokens <= kMaxTokens); + // Keep batch_id below 65535: pack_w(65535, 65535, ...) is the invalid sentinel. + RuntimeCheck(batch_size < kMaxTokens && batch_size <= num_q_tokens && num_q_tokens <= kMaxTokens); // `swa_page_size` >= `ring_size` >= `compress_ratio` RuntimeCheck(swa_page_size % ring_size == 0 && ring_size % compress_ratio == 0); // Write pad: trailing tokens kept resident so a verify batch's committed tail survives @@ -750,9 +752,9 @@ inline PrefillPlan plan_compress_prefill_legacy( const auto window_size = compress_ratio * (is_overlap ? 2 : 1); const auto batch_size = static_cast(B.unwrap()); - constexpr auto kMaxTokens = static_cast(std::numeric_limits::max()); + constexpr auto kMaxTokens = static_cast(std::numeric_limits::max()) + 1; RuntimeCheck(compress_ratio == 4 || compress_ratio == 128); - RuntimeCheck(batch_size <= num_q_tokens && num_q_tokens <= kMaxTokens); + RuntimeCheck(batch_size < kMaxTokens && batch_size <= num_q_tokens && num_q_tokens <= kMaxTokens); uint32_t counter = 0; uint32_t counter_c = 0; diff --git a/test/registered/kernels/ops/attention/test_deepseek_v4_compress_plan_draft_pad.py b/test/registered/kernels/ops/attention/test_deepseek_v4_compress_plan_draft_pad.py index 8765cb7c8df4..d00a8323de1d 100644 --- a/test/registered/kernels/ops/attention/test_deepseek_v4_compress_plan_draft_pad.py +++ b/test/registered/kernels/ops/attention/test_deepseek_v4_compress_plan_draft_pad.py @@ -21,7 +21,11 @@ import torch from sglang.test.ci.ci_register import register_cuda_ci -from sglang.test.kernels.deepseek_v4.common import make_paged_context, to_seq_extend +from sglang.test.kernels.deepseek_v4.common import ( + make_legacy_context, + make_paged_context, + to_seq_extend, +) from sglang.test.test_utils import CustomTestCase register_cuda_ci(est_time=20, stage="base-b-kernel-unit", runner_config="1-gpu-large") @@ -58,6 +62,55 @@ def _written_positions(plan_w: torch.Tensor, prefix_len: int) -> set[int]: class TestCompressWritePlanDraftPad(CustomTestCase): + def test_64k_prefill_preserves_last_token(self): + """65536 tokens fit uint16 indices; the last token must not wrap or vanish.""" + for cr in (4, 128): + paged = make_paged_context( + bs=16, compress_ratio=cr, num_swa_pages_per_req=16 + ) + legacy = make_legacy_context(bs=16, compress_ratio=cr) + seq_lens, extend_lens, num_q = to_seq_extend([(4096, 4096)] * 16) + for ctx, on_gpu in ((paged, False), (paged, True), (legacy, False)): + with self.subTest(cr=cr, paged=ctx is paged, on_gpu=on_gpu): + device = "cuda" if on_gpu else "cpu" + plan = ctx.make_prefill_plan( + seq_lens.to(device), extend_lens.to(device), num_q + ) + c = plan.plan_c.cpu().view(torch.int32).view(-1, 4) + valid_c = c[:, 0] != -1 + ids = c[valid_c, 1].bitwise_and(0xFFFF).sort().values + torch.testing.assert_close( + ids, torch.arange(cr - 1, num_q, cr, dtype=torch.int32) + ) + w = plan.plan_w.cpu().view(torch.int32).view(-1, 2) + last = w[w[:, 0] == 65535] + if cr == 4: + self.assertEqual(len(last), 1) + self.assertEqual(int(last[0, 1]), ctx.state_loc(15, 4095)) + else: + # Non-overlapping C128 consumed the complete final block; + # no raw tail remains to persist into the state ring. + self.assertEqual(len(w[w[:, 0] != -1]), 0) + + def test_prefill_rejects_uint16_index_overflow(self): + for ctx in ( + make_paged_context(bs=16, compress_ratio=4, num_swa_pages_per_req=17), + make_legacy_context(bs=16, compress_ratio=4), + ): + seq_lens, extend_lens, num_q = to_seq_extend( + [(4096, 4096)] * 15 + [(4097, 4097)] + ) + with self.assertRaisesRegex(RuntimeError, "plan_compress_prefill"): + ctx.make_prefill_plan(seq_lens, extend_lens, num_q) + + def test_prefill_rejects_packed_invalid_sentinel(self): + # A 65536-request, one-token-per-request batch makes the last packed + # (batch_id, ragged_id) equal (65535, 65535), the invalid write sentinel. + ctx = make_legacy_context(bs=65536, compress_ratio=4) + seq_lens, extend_lens, num_q = to_seq_extend([(1, 1)] * 65536) + with self.assertRaisesRegex(RuntimeError, "plan_compress_prefill"): + ctx.make_prefill_plan(seq_lens, extend_lens, num_q) + def _make_plan_positions( self, *, From e36a479bc90ab12876c44a6e279fe95b37714c10 Mon Sep 17 00:00:00 2001 From: BBuf <1182563586@qq.com> Date: Wed, 16 Sep 2026 12:21:40 +0800 Subject: [PATCH 13/30] fix c2 padding and trim metadata tests --- .../kernels/jit/csrc/deepseek_v4/c2.cuh | 10 +- .../ops/attention/dsv4/dequant_k_cache.py | 5 + .../kernel/attention/dsv4/test_c2_verify.py | 397 +----------------- .../attention/dsv4/test_v41_kv_dequant.py | 132 ------ .../attention/dsv4/test_v41_kv_store.py | 16 +- .../attention/test_dsv41_small_metadata.py | 56 --- ...est_deepseek_v4_compress_plan_draft_pad.py | 200 +-------- 7 files changed, 46 insertions(+), 770 deletions(-) delete mode 100644 test/registered/kernel/attention/dsv4/test_v41_kv_dequant.py delete mode 100644 test/registered/kernel/attention/test_dsv41_small_metadata.py diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/c2.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/c2.cuh index f16c308885cf..c50e1283e143 100644 --- a/python/sglang/kernels/jit/csrc/deepseek_v4/c2.cuh +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/c2.cuh @@ -102,13 +102,16 @@ __global__ __launch_bounds__(kHeadDim / kC2VecSize) void flash_c2_decode_kernel( const auto tx = threadIdx.x; // Verify gives each request a CTA column; decode a flat grid of one row each. const auto row = kVerify ? blockIdx.y * gridDim.x + blockIdx.x : blockIdx.x; + PDLWaitPrimary(); // Slots fit in int32 whatever width the scheduler hands them in. const auto raw_out_loc = static_cast(static_cast(params.raw_out_loc)[row]); + // CUDA graph padding is a completely inert row: do not read schedule, input, + // state, or RoPE data, and do not publish output or update the cache. + if (raw_out_loc == 0) return PDLTriggerSecondary(); const auto pos = static_cast(params.positions)[row]; // A completing row reads the slot left by `pos - 1`; // a pending row writes its own slot, so reads and writes stay disjoint. const auto rid = params.req[row]; - PDLWaitPrimary(); fp32_vec_t kv_new, score_new; kv_new.load(params.kv_input + row * kStride, tx); @@ -127,8 +130,6 @@ __global__ __launch_bounds__(kHeadDim / kC2VecSize) void flash_c2_decode_kernel( score_old.load(partner, tx + kCTASize); if ((pos & 1) == 0) { - // padded case - if (raw_out_loc == 0) return PDLTriggerSecondary(); kv_new.store(params.kv_state + write_row * kStride, tx); score_new.store(params.kv_state + write_row * kStride, tx + kCTASize); return PDLTriggerSecondary(); @@ -220,7 +221,6 @@ __global__ __launch_bounds__(kHeadDim / kC2VecSize) void flash_c2_decode_kernel( if constexpr (kLayout == KVLayout::V41_FP4) { // The fp4 cache takes the rotated bf16 value as is: its row quantizer is the // fake quantization, minus the dequantization. - if (raw_out_loc == 0) return; const int32_t out_loc = raw_out_loc >> 1; const auto kv_row = Paged::row(params.kvcache, out_loc); return deepseek_v4::v41::store_row(kv_row.data, kv_row.scale, tx, staged); @@ -244,8 +244,6 @@ __global__ __launch_bounds__(kHeadDim / kC2VecSize) void flash_c2_decode_kernel( } } - // padded case - if (raw_out_loc == 0) return; // `raw_out_loc / ratio`; ratio 2 makes it a shift. const int32_t out_loc = raw_out_loc >> 1; const auto kv_row = Paged::row(params.kvcache, out_loc); diff --git a/python/sglang/kernels/ops/attention/dsv4/dequant_k_cache.py b/python/sglang/kernels/ops/attention/dsv4/dequant_k_cache.py index 3c0cc8e00838..13b6e953a2eb 100644 --- a/python/sglang/kernels/ops/attention/dsv4/dequant_k_cache.py +++ b/python/sglang/kernels/ops/attention/dsv4/dequant_k_cache.py @@ -4,6 +4,7 @@ import triton import triton.language as tl +from sglang.kernels.jit.utils import get_jit_cuda_arch, is_hip_runtime from sglang.kernels.ops.attention.dsv4.kv_layout import KVLayout from sglang.kernels.ops.quantization.fp8_kernel import is_fp8_fnuz @@ -110,6 +111,10 @@ def dequantize_k_cache_paged_v41( """ layout = KVLayout.parse(layout) assert layout in (KVLayout.V41, KVLayout.V41_FP4), layout + if is_hip_runtime() or get_jit_cuda_arch().major < 10: + raise RuntimeError( + "DeepSeek V4.1 KV cache dequantization requires CUDA SM100 or newer" + ) assert quant_k_cache.is_contiguous() assert page_table_1_flattened.dtype in (torch.int32, torch.int64) diff --git a/test/registered/kernel/attention/dsv4/test_c2_verify.py b/test/registered/kernel/attention/dsv4/test_c2_verify.py index 7eb55dc45ef9..be637cf26670 100644 --- a/test/registered/kernel/attention/dsv4/test_c2_verify.py +++ b/test/registered/kernel/attention/dsv4/test_c2_verify.py @@ -1,391 +1,38 @@ -"""Verify compressor: exact decode replay, padding, wrap and rejected prefixes.""" - -import sys +"""C2 graph-padding rows must leave every destination untouched.""" import pytest import torch -from torch import nn -from sglang.kernels.ops.attention.dsv4.c2 import c2_decode_or_verify_norm_rope_store -from sglang.kernels.ops.attention.dsv4.rmsnorm_fp32 import rmsnorm_fp32 -from sglang.srt.model_loader.utils import set_default_torch_dtype +from sglang.kernels.ops.attention.dsv4.c2 import c2_decode_norm from sglang.test.ci.ci_register import register_cuda_ci -register_cuda_ci(est_time=60, stage="base-b-kernel-unit", runner_config="4-gpu-b200") -pytestmark = pytest.mark.skipif( - not torch.cuda.is_available() - or torch.version.cuda is None - or torch.cuda.get_device_capability()[0] < 10, - reason="the compressor packs FP4 with SM100 instructions", -) - -EPS = 1e-6 -HEAD_DIMS = (512,) -ROPE_DIM = 64 -RATIO = 2 -PAGE_SIZE = 128 -PAGE_BYTES = -(-584 * PAGE_SIZE // 576) * 576 -DRAFT_LENS = (2, 5, 6, 9) -VERIFY_BATCHES = (1, 3, 8) - - -class RMSNorm(nn.Module): - """fp32 statistics and fp32 weight multiply, cast back at the very end.""" - - def __init__(self, dim: int, eps: float): - super().__init__() - self.eps = eps - self.weight = nn.Parameter(torch.ones(dim)) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - if ( - x.is_cuda - and torch.version.cuda is not None - and x.dtype in (torch.bfloat16, torch.float32) - and self.weight.dtype in (torch.bfloat16, torch.float32) - and x.shape[-1] in (128, 512) - and x.is_contiguous() - and self.weight.is_contiguous() - ): - return rmsnorm_fp32(x, self.weight, self.eps) - dtype = x.dtype - x = x.float() - x = x * torch.rsqrt(x.square().mean(-1, keepdim=True) + self.eps) - return (self.weight * x).to(dtype) - - -def _norm(dim: int, seed: int) -> RMSNorm: - """`DeepseekV41Compressor.norm` as the model holds it: bf16 weight, because - the parameter is created inside `set_default_torch_dtype(model dtype)`.""" - with set_default_torch_dtype(torch.bfloat16): - norm = RMSNorm(dim, EPS).cuda() - assert norm.weight.dtype == torch.bfloat16 - # Not `ones`: a constant weight cannot catch a wrong per-element index. - g = torch.Generator(device="cuda").manual_seed(seed) - with torch.no_grad(): - norm.weight.copy_(torch.randn(dim, generator=g, device="cuda")) - return norm +register_cuda_ci(est_time=5, stage="base-b-kernel-unit", runner_config="1-gpu-large") -def _freqs(max_pos, seed): - """`layer.freqs_cis` and the fp32 real/imag-interleaved view the kernel - indexes itself, at `positions - 1`.""" - g = torch.Generator(device="cuda").manual_seed(seed) - ang = torch.randn(max_pos, ROPE_DIM // 2, generator=g, device="cuda") - freqs = torch.polar(torch.ones_like(ang), ang) - return freqs, torch.view_as_real(freqs).flatten(-2).contiguous().float() +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +def test_padded_even_and_odd_rows_are_inert() -> None: + num_tokens, head_dim, ring_size = 2, 512, 2 + kv_input = torch.randn(num_tokens, 2 * head_dim, device="cuda") + kv_state = torch.randn(num_tokens * ring_size, 2 * head_dim, device="cuda") + state_before = kv_state.clone() + norm_weight = torch.randn(head_dim, dtype=torch.bfloat16, device="cuda") + positions = torch.tensor([2, 3], dtype=torch.int64, device="cuda") + req = torch.arange(num_tokens, dtype=torch.int64, device="cuda") + raw_out_loc = torch.zeros(num_tokens, dtype=torch.int32, device="cuda") + out = torch.full((num_tokens, head_dim), 123, dtype=torch.bfloat16, device="cuda") + out_before = out.clone() - -def _cache(max_slot): - """A compressed-pool buffer wide enough for `max_slot`, zeroed so an - untouched slot is recognizable.""" - return torch.zeros( - max_slot // PAGE_SIZE + 2, PAGE_BYTES, dtype=torch.uint8, device="cuda" - ) - - -def _spec_ring_size(draft_len): - """`get_compress_state_ring_size(2, is_speculative=True, draft_len)`.""" - return 1 << (draft_len + 1).bit_length() - - -def _verify_inputs(bs, draft_len, dim, seed, *, starts=None, pad_reqs=0): - """A target-verify batch: `draft_len` consecutive positions per request, - request-major. Block heads alternate parity by default, so some blocks open - by consuming the ring and some by parking into it.""" - g = torch.Generator(device="cuda").manual_seed(seed) - n = bs * draft_len - ring_size = _spec_ring_size(draft_len) - kv_input = torch.randn(n, 2 * dim, generator=g, device="cuda", dtype=torch.float32) - kv_state = torch.randn( - (bs + 1) * ring_size, 2 * dim, generator=g, device="cuda", dtype=torch.float32 - ) - if starts is None: - starts = torch.arange(bs, device="cuda") + 4 - offsets = torch.arange(draft_len, device="cuda") - positions = (starts[:, None] + offsets[None, :]).flatten().to(torch.int32) - req = torch.arange(bs, device="cuda", dtype=torch.int64).repeat_interleave( - draft_len - ) - raw_out_loc = torch.arange(n, device="cuda", dtype=torch.int32) * 2 + 3 - if pad_reqs: - # Graph padding pads whole request slots, aliasing a live - # `req_pool_idx`, and leaves their position buffer at zero. - pad = req >= bs - pad_reqs - raw_out_loc[pad] = 0 - positions[pad] = 0 - req[pad] = 0 - return kv_input, kv_state, positions, req, raw_out_loc, ring_size - - -def _decode_replay( - kv_input, kv_state, norm, positions, req, raw_out_loc, freqs_cis, cache, **kw -): - """The same block, one position per launch -- what one fused verify launch - has to reproduce. Step `j` is a decode step over row `j` of every request, - carrying the pair ring between steps exactly as the served decode path does. - """ - draft_len = kw["draft_len"] - n, dim = positions.shape[0], kv_input.shape[1] // 2 - rows = torch.arange(n, device="cuda").view(-1, draft_len) - out = torch.zeros(n, dim, device="cuda", dtype=torch.bfloat16) - for j in range(draft_len): - idx = rows[:, j] - out[idx] = c2_decode_or_verify_norm_rope_store( - kv_input[idx].contiguous(), - kv_state, - norm.weight.data, - positions[idx].contiguous(), - req[idx].contiguous(), - raw_out_loc[idx].contiguous(), - EPS, - freqs_cis, - cache, - page_size=PAGE_SIZE, - ring_size=kw["ring_size"], - out=torch.zeros(idx.numel(), dim, device="cuda", dtype=torch.bfloat16), - ) - return out - - -def _run_verify(kv_input, kv_state, norm, positions, req, raw_out_loc, **kw): - """One `c2_verify_norm_rope_store` call; returns `(latent, cache)`. `out` is - zeroed so rows the kernel skips compare equal to the replay's.""" - n, dim = positions.shape[0], kv_input.shape[1] // 2 - freqs_cis, cache = kw["freqs_cis"], kw["cache"] - got = c2_decode_or_verify_norm_rope_store( + result = c2_decode_norm( kv_input, kv_state, - norm.weight.data, - positions, - req, - raw_out_loc, - EPS, - freqs_cis, - cache, - page_size=PAGE_SIZE, - ring_size=kw["ring_size"], - draft_len=kw["draft_len"], - out=torch.zeros(n, dim, device="cuda", dtype=torch.bfloat16), - ) - return got, cache - - -@pytest.mark.parametrize("pad_reqs", (0, 1)) -@pytest.mark.parametrize("draft_len", DRAFT_LENS) -@pytest.mark.parametrize("bs", VERIFY_BATCHES) -def test_verify_matches_decode_replay(bs, draft_len, pad_reqs): - """The load-bearing property: one verify launch over a block equals - `draft_len` decode launches over the same rows -- same latents, same pair - ring, same cache bytes, bitwise. Verify takes the in-block partner from - `kv_input` where decode takes it from the ring, and that substitution has to - be invisible.""" - if pad_reqs >= bs: - pytest.skip("an all-padded batch has no live block to compare") - dim = HEAD_DIMS[0] - kv_input, kv_state, positions, req, raw_out_loc, ring_size = _verify_inputs( - bs, draft_len, dim, seed=20000 + bs * 97 + draft_len, pad_reqs=pad_reqs - ) - norm = _norm(dim, 20001 + bs + draft_len) - _, freqs_cis = _freqs(int(positions.max().item()) + 2, 20002 + draft_len) - slots_max = int((raw_out_loc // RATIO).max().item()) - cache_v, cache_d = _cache(slots_max), _cache(slots_max) - state_v, state_d = kv_state.clone(), kv_state.clone() - kw = dict(ring_size=ring_size, draft_len=draft_len) - - got, _ = _run_verify( - kv_input, - state_v, - norm, + norm_weight, positions, req, raw_out_loc, - freqs_cis=freqs_cis, - cache=cache_v, - **kw, + 1e-6, + ring_size=ring_size, + out=out, ) - expected = _decode_replay( - kv_input, state_d, norm, positions, req, raw_out_loc, freqs_cis, cache_d, **kw - ) - - assert cache_d.any(), "the replay stored nothing, so the comparison is empty" - assert torch.equal(got, expected), "latent differs from the decode replay" - assert torch.equal(state_v, state_d), "pair ring differs from the decode replay" - assert torch.equal(cache_v, cache_d), "cache bytes differ from the decode replay" - - -@pytest.mark.parametrize("draft_len", DRAFT_LENS) -def test_verify_reads_the_ring_only_on_the_first_row(draft_len): - """A block's first row is the only one allowed to consume the ring. Move the - slot it reads and its latent must move with it; every later row pairs inside - the block and must not notice.""" - bs, dim = 4, HEAD_DIMS[0] - # Odd heads, so every block opens by completing a group against the ring. - starts = 2 * torch.arange(bs, device="cuda") + 5 - kv_input, kv_state, positions, req, raw_out_loc, ring_size = _verify_inputs( - bs, draft_len, dim, seed=21000 + draft_len, starts=starts - ) - norm = _norm(dim, 21001 + draft_len) - _, freqs_cis = _freqs(int(positions.max().item()) + 2, 21002 + draft_len) - slots_max = int((raw_out_loc // RATIO).max().item()) - kw = dict(ring_size=ring_size, draft_len=draft_len, freqs_cis=freqs_cis) - - base, _ = _run_verify( - kv_input, - kv_state.clone(), - norm, - positions, - req, - raw_out_loc, - cache=_cache(slots_max), - **kw, - ) - heads = torch.arange(0, bs * draft_len, draft_len, device="cuda") - read = req[heads] * ring_size + (positions[heads].to(torch.int64) - 1) % ring_size - moved = kv_state.clone() - moved[read] += 1.0 - got, _ = _run_verify( - kv_input, - moved, - norm, - positions, - req, - raw_out_loc, - cache=_cache(slots_max), - **kw, - ) - - rest = torch.ones(bs * draft_len, dtype=torch.bool, device="cuda") - rest[heads] = False - assert not torch.equal(got[heads], base[heads]), "a block head ignored the ring" - assert torch.equal(got[rest], base[rest]), "a later row went through the ring" - - -def test_verify_rejects_a_ring_narrower_than_the_block(): - """`ring_size > draft_len` is the whole reason a block's own publishes stay - off the slot its first row reads, so the kernel refuses a narrower ring - rather than racing quietly.""" - bs, draft_len, dim = 2, 4, HEAD_DIMS[0] - kv_input, kv_state, positions, req, raw_out_loc, _ = _verify_inputs( - bs, draft_len, dim, seed=22000 - ) - norm = _norm(dim, 22001) - _, freqs_cis = _freqs(int(positions.max().item()) + 2, 22002) - with pytest.raises(Exception, match="must be wider than the draft length"): - _run_verify( - kv_input, - kv_state, - norm, - positions, - req, - raw_out_loc, - cache=_cache(int((raw_out_loc // RATIO).max().item())), - freqs_cis=freqs_cis, - ring_size=draft_len, - draft_len=draft_len, - ) - - -@pytest.mark.parametrize("start", (31, 32)) -@pytest.mark.parametrize("dtype", (torch.int32, torch.int64)) -def test_rejected_prefix_then_next_verify(start, dtype): - # Compare with a decode history that never saw the rejected suffix. The - # next verify starts at the committed position, including across ring wrap. - bs, draft_len, dim = 3, 6, 512 - inputs, initial, pos, req, loc, ring = _verify_inputs( - bs, - draft_len, - dim, - 23000, - starts=torch.full((bs,), start, device="cuda"), - ) - pos, loc = pos.to(dtype), loc.to(dtype) - norm = _norm(dim, 23001) - _, freqs = _freqs(start + 2 * draft_len + 2, 23002) - rows = torch.arange(bs * draft_len, device="cuda").view(bs, draft_len) - for accepted in range(1, draft_len + 1): - state = initial.clone() - reference = initial.clone() - cache, ref_cache = _cache(128), _cache(128) - kw = dict(draft_len=draft_len, ring_size=ring, freqs_cis=freqs) - _run_verify(inputs, state, norm, pos, req, loc, cache=cache, **kw) - for j in range(accepted): - idx = rows[:, j] - c2_decode_or_verify_norm_rope_store( - inputs[idx], - reference, - norm.weight.data, - pos[idx], - req[idx], - loc[idx], - EPS, - freqs, - ref_cache, - page_size=PAGE_SIZE, - ring_size=ring, - ) - # Do not count speculative cache bytes that have no committed reader. - cache.zero_() - ref_cache.zero_() - next_inputs = inputs.flip(0).contiguous() - got, _ = _run_verify( - next_inputs, - state, - norm, - pos + accepted, - req, - loc, - cache=cache, - **kw, - ) - expected = _decode_replay( - next_inputs, - reference, - norm, - pos + accepted, - req, - loc, - freqs, - ref_cache, - ring_size=ring, - draft_len=draft_len, - ) - assert torch.equal(got, expected), f"{start=} {accepted=}: latent differs" - assert torch.equal(cache, ref_cache), f"{start=} {accepted=}: cache differs" - - -def test_verify_cuda_graph_replay(): - bs, draft_len, dim = 3, 6, 512 - inputs, initial, pos, req, loc, ring = _verify_inputs(bs, draft_len, dim, 24000) - norm = _norm(dim, 24001) - _, freqs = _freqs(32, 24002) - state, cache = initial.clone(), _cache(128) - kw = dict(ring_size=ring, draft_len=draft_len, freqs_cis=freqs, cache=cache) - _run_verify(inputs, state, norm, pos, req, loc, **kw) - torch.cuda.synchronize() - graph = torch.cuda.CUDAGraph() - with torch.cuda.graph(graph): - got, _ = _run_verify(inputs, state, norm, pos, req, loc, **kw) - state.copy_(initial) - cache.zero_() - graph.replay() - ref_cache = _cache(128) - expected = _decode_replay( - inputs, - initial, - norm, - pos, - req, - loc, - freqs, - ref_cache, - ring_size=ring, - draft_len=draft_len, - ) - assert torch.equal(got, expected) - assert torch.equal(state, initial) - assert torch.equal(cache, ref_cache) - -if __name__ == "__main__": - sys.exit(pytest.main([__file__])) + assert torch.equal(result, out_before) + assert torch.equal(kv_state, state_before) diff --git a/test/registered/kernel/attention/dsv4/test_v41_kv_dequant.py b/test/registered/kernel/attention/dsv4/test_v41_kv_dequant.py deleted file mode 100644 index 17049308b06d..000000000000 --- a/test/registered/kernel/attention/dsv4/test_v41_kv_dequant.py +++ /dev/null @@ -1,132 +0,0 @@ -"""The V4.1 paged dequant (bf16 prefill workspace) is bit-exact with the -pure-torch dequantizer of the fp8 (V41) and fp4 (V41_FP4) formats.""" - -import unittest - -import torch - -from sglang.kernels.ops.attention.dsv4.dequant_k_cache import dequantize_k_cache_paged -from sglang.kernels.ops.attention.dsv4.kv_layout import KVLayout -from sglang.srt.layers.attention.dsv4 import torch_quant as tq -from sglang.test.ci.ci_register import register_cuda_ci -from sglang.test.test_utils import CustomTestCase - -register_cuda_ci(est_time=60, stage="base-b-kernel-unit", runner_config="1-gpu-large") -register_cuda_ci(est_time=60, stage="base-b-kernel-unit", runner_config="4-gpu-b200") - -CASES = { - KVLayout.V41: (tq.quantize_k_cache_v41, tq.dequantize_k_cache_v41), - KVLayout.V41_FP4: (tq.quantize_k_cache_v41_fp4, tq.dequantize_k_cache_v41_fp4), -} - - -def bits(t: torch.Tensor) -> torch.Tensor: - """bf16 as int16, so that -0.0 and NaN payloads compare exactly.""" - return t.contiguous().view(torch.int16) - - -@unittest.skipUnless(torch.cuda.is_available(), "requires CUDA") -class TestV41KVDequant(CustomTestCase): - def _gather_ref(self, dequant, pages, page_size, ids): - return dequant(pages, page_size).view(-1, 512)[ids.long()].unsqueeze(1) - - def test_quantized_pages(self): - g = torch.Generator(device="cuda").manual_seed(0) - for layout, (quant, dequant) in CASES.items(): - for page_size, num_pages in ((64, 9), (256, 3), (2, 50)): - with self.subTest(layout=layout.name, page_size=page_size): - k = torch.randn( - num_pages, - page_size, - 512, - generator=g, - device="cuda", - dtype=torch.bfloat16, - ) - k = ( - k - * torch.exp2( - torch.randint( - -12, - 6, - (num_pages, page_size, 1), - generator=g, - device="cuda", - ).float() - ) - ).to(torch.bfloat16) - k[0, 0, :32] = 0 - k[0, 0, 32:48] = -0.0 - pages = quant(k, page_bytes=layout.page_bytes(page_size)) - ids = torch.randint( - 0, - num_pages * page_size, - (777,), - generator=g, - device="cuda", - dtype=torch.int32, - ) - got = dequantize_k_cache_paged(pages, ids, page_size, layout=layout) - self.assertEqual(got.shape, (777, 1, 512)) - self.assertTrue( - torch.equal( - bits(got), - bits(self._gather_ref(dequant, pages, page_size, ids)), - ) - ) - # The fp4 cache dequantizes to the model's fake-quantized value - # (compared by value: the fake quant maps an exact -0.0 to +0.0). - if layout is KVLayout.V41_FP4: - expect = tq.fake_quant_compressed_kv( - k.view(-1, 512)[ids.long()] - ).unsqueeze(1) - self.assertTrue(torch.equal(got, expect)) - - def test_random_bytes_and_workspace_slice(self): - """Arbitrary payload bytes (scales in the quantizer's range) and an - ``out`` that is a strided slice of a larger workspace.""" - g = torch.Generator(device="cuda").manual_seed(1) - for layout, (_, dequant) in CASES.items(): - page_size, num_pages = 64, 7 - with self.subTest(layout=layout.name): - pages = torch.randint( - 0, - 256, - (num_pages, layout.page_bytes(page_size)), - generator=g, - dtype=torch.uint8, - device="cuda", - ) - if layout is KVLayout.V41: - lo = layout.scale_offset(page_size) - hi = lo + page_size * layout.scale_bytes - pages[:, lo:hi] = torch.randint( - 100, - 140, - (num_pages, hi - lo), - generator=g, - dtype=torch.uint8, - device="cuda", - ) - ids = torch.randint( - 0, - num_pages * page_size, - (300,), - generator=g, - device="cuda", - dtype=torch.int64, - ) - ref = self._gather_ref(dequant, pages, page_size, ids) - workspace = torch.zeros( - 305, 1, 512, dtype=torch.bfloat16, device="cuda" - ) - out = dequantize_k_cache_paged( - pages, ids, page_size, out=workspace[5:], layout=layout - ) - # NaN payloads (fp8 0x7F / e4m3 NaN scales) compare through their bits. - self.assertTrue(torch.equal(bits(workspace[5:]), bits(ref))) - self.assertEqual(int(workspace[:5].abs().sum()), 0) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/kernel/attention/dsv4/test_v41_kv_store.py b/test/registered/kernel/attention/dsv4/test_v41_kv_store.py index 12a3b3bc3014..b6f428bf7cbf 100644 --- a/test/registered/kernel/attention/dsv4/test_v41_kv_store.py +++ b/test/registered/kernel/attention/dsv4/test_v41_kv_store.py @@ -421,7 +421,8 @@ def stored_reference(layout, rotated, page_size, num_pages, locs, page_bytes): kv_old = torch.randn(n, 512, generator=g, device="cuda") * 2 kv_input = torch.cat([kv_new, score], dim=-1).contiguous() req = torch.arange(n, device="cuda", dtype=torch.int64) - # Odd positions complete a pair; one even (pending) row and one padded row. + # Odd positions complete a pair; row 5 is even and rows 5/7 + # are padded graph rows that must remain completely inert. pos = ( 2 * torch.randint( @@ -441,13 +442,19 @@ def stored_reference(layout, rotated, page_size, num_pages, locs, page_bytes): )[:n].to(torch.int32) + 1 ) * 2 - raw_out_loc[7] = 0 + raw_out_loc[[5, 7]] = 0 cache = torch.zeros( num_pages, layout.page_bytes(page_size), dtype=torch.uint8, device="cuda", ) + latent_out = torch.full( + (n, 512), 123, dtype=torch.bfloat16, device="cuda" + ) + padded_out_before = latent_out[[5, 7]].clone() + padded_state_row = req[5] * ring + pos[5] % ring + padded_state_before = state[padded_state_row].clone() latent = c2_decode_or_verify_norm_rope_store( kv_input, state, @@ -461,6 +468,11 @@ def stored_reference(layout, rotated, page_size, num_pages, locs, page_bytes): page_size=page_size, ring_size=ring, layout=layout, + out=latent_out, + ) + self.assertTrue(torch.equal(latent[[5, 7]], padded_out_before)) + self.assertTrue( + torch.equal(state[padded_state_row], padded_state_before) ) valid = (raw_out_loc != 0) & (pos % 2 == 1) pooled = ((kv_old + kv_new) / 2).to(torch.bfloat16) diff --git a/test/registered/kernel/attention/test_dsv41_small_metadata.py b/test/registered/kernel/attention/test_dsv41_small_metadata.py deleted file mode 100644 index 723d561686c7..000000000000 --- a/test/registered/kernel/attention/test_dsv41_small_metadata.py +++ /dev/null @@ -1,56 +0,0 @@ -"""Integer metadata equivalence with changing CUDA graph inputs.""" - -import sys - -import pytest -import torch - -from sglang.kernels.ops.attention.dsv4_attn_metadata_kernels import ( - BuildPageTablePositions, -) -from sglang.kernels.ops.attention.dsv41_small_metadata import ( - page_table_positions_small, -) -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=60, stage="base-b-kernel-unit", runner_config="1-gpu-large") - - -@pytest.mark.parametrize("rows", [1, 5, 6, 8]) -@pytest.mark.parametrize("pages", [1, 17, 4098]) -@pytest.mark.parametrize("dtype", [torch.int32, torch.int64]) -def test_pages_graph(rows, pages, dtype): - # Rows have a larger physical stride than the logical table. - mapping = torch.randint(-256, 1 << 24, (9, pages * 256 + 512), device="cuda") - reqs = torch.arange(rows, device="cuda", dtype=dtype) - lens = torch.arange(rows, device="cuda", dtype=dtype) - args = dict( - req_to_token=mapping, - req_pool_indices_repeated=reqs, - seq_lens_casual=lens, - max_seq_len=pages * 256, - page_size=256, - swa_window=128, - ) - for _ in range(3): - page_table_positions_small(**args) - graph = torch.cuda.CUDAGraph() - with torch.cuda.graph(graph): - out = page_table_positions_small(**args) - for replay in range(5): - mapping.random_(-256, 1 << 24) - reqs.copy_((torch.arange(rows, device="cuda") + replay) % 9) - lens.copy_(torch.arange(rows, device="cuda") * 127 + replay - 1) - graph.replay() - ref = BuildPageTablePositions.triton(**args) - for name in ( - "seq_lens_casual", - "positions_casual", - "page_table", - "swa_topk_lengths", - ): - assert torch.equal(getattr(out, name), getattr(ref, name)), name - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__])) diff --git a/test/registered/kernels/ops/attention/test_deepseek_v4_compress_plan_draft_pad.py b/test/registered/kernels/ops/attention/test_deepseek_v4_compress_plan_draft_pad.py index d00a8323de1d..ce0649534a10 100644 --- a/test/registered/kernels/ops/attention/test_deepseek_v4_compress_plan_draft_pad.py +++ b/test/registered/kernels/ops/attention/test_deepseek_v4_compress_plan_draft_pad.py @@ -1,18 +1,4 @@ -"""Kernel-level tests for the DSV4 compress write-plan (`plan_prefill`). - -`plan_w` decides which tokens' raw KV get persisted into the compress-state ring -for a *future* compression window to read. A speculative verify batch plans from -the optimistic `seq_len = prefix + num_draft_tokens` but rolls back to -`prefix + accept_len`, so every committed token must stay resident whatever the -accept length -- i.e. the plan must write all of `[prefix, seq_len)`. - -`c_plan.cuh` used to cap that pad at 4 (`kMaxMTPDraftTokens`), silently -under-writing the ring for larger draft counts -- no IMA, no NaN, just wrong -compressed state. The pad now comes from the ring itself -(`ring_size - window_size + 2`), covering every draft count the ring can serve. -Tests pin the invariant on both planner paths (CPU host loop and GPU -`plan_compress_prefill_kernel0`) and both compress ratios. -""" +"""Boundary tests for the packed indices in the DSV4 prefill write plan.""" from __future__ import annotations @@ -30,36 +16,6 @@ register_cuda_ci(est_time=20, stage="base-b-kernel-unit", runner_config="1-gpu-large") -C4_RING_SIZE = 16 # get_compress_state_ring_size(4, is_speculative=True) -C128_RING_SIZE = 256 # get_compress_state_ring_size(128, is_speculative=True) -C4_RING_SIZE_NO_SPEC = 8 # get_compress_state_ring_size(4, is_speculative=False) -C128_RING_SIZE_NO_SPEC = 128 # get_compress_state_ring_size(128, is_speculative=False) - - -def _window_size(compress_ratio: int) -> int: - """Tokens read by one compression: c4 overlaps two chunks, c128 does not.""" - return compress_ratio * (2 if compress_ratio == 4 else 1) - - -def _max_draft_tokens(compress_ratio: int, ring_size: int) -> int: - """Largest draft count this ring serves; mirrors `mtp_pad` in c_plan.cuh.""" - window = _window_size(compress_ratio) - return ring_size - window + 2 if ring_size > window else 0 - - -def _written_positions(plan_w: torch.Tensor, prefix_len: int) -> set[int]: - """Decode `plan_w` into the set of positions written, for a bs=1 plan. - - `plan_w` is `[n, 8]` uint8 = (uint32 ragged_id, int32 write_loc). Stage 1 - overwrites `write_loc` with the final state slot, but `ragged_id` survives and - equals the token's index within the ragged layout, so for a single request - `position = prefix_len + ragged_id`. - """ - words = plan_w.cpu().view(torch.uint32).view(-1, 2) - ragged_ids = words[:, 0] - valid = ragged_ids != 0xFFFFFFFF - return {prefix_len + int(r) for r in ragged_ids[valid]} - class TestCompressWritePlanDraftPad(CustomTestCase): def test_64k_prefill_preserves_last_token(self): @@ -111,160 +67,6 @@ def test_prefill_rejects_packed_invalid_sentinel(self): with self.assertRaisesRegex(RuntimeError, "plan_compress_prefill"): ctx.make_prefill_plan(seq_lens, extend_lens, num_q) - def _make_plan_positions( - self, - *, - compress_ratio: int, - ring_size: int, - prefix_len: int, - num_draft_tokens: int, - on_gpu: bool = False, - ) -> set[int]: - """Build a bs=1 verify plan and return the positions it writes. - - `on_gpu=True` moves the planner inputs to device, which routes - `plan_prefill` to `plan_compress_prefill_kernel0` instead of the host loop. - """ - ctx = make_paged_context( - bs=1, compress_ratio=compress_ratio, ring_size=ring_size - ) - seq_lens, extend_lens, num_q = to_seq_extend( - [(prefix_len + num_draft_tokens, num_draft_tokens)] - ) - if on_gpu: - seq_lens = seq_lens.to(ctx.req_to_token.device) - extend_lens = extend_lens.to(ctx.req_to_token.device) - plan = ctx.make_prefill_plan(seq_lens, extend_lens, num_q) - return _written_positions(plan.plan_w, prefix_len) - - def _assert_ring_residency(self, compress_ratio: int, ring_size: int): - """Every committed token must be written, for each (D, sl mod cr) combo. - - This is the sufficient condition, which is why there is no multi-step replay - test: if a step writes all of `[prefix, prefix + D)`, then whatever the accept - length, the tokens the next compression window needs are either from this step - (written here) or older (written by an earlier step, same invariant by - induction). - """ - max_d = _max_draft_tokens(compress_ratio, ring_size) - # Vary `seq_len % compress_ratio`: that residue decides whether the unpadded rule - # alone would have sufficed. Four is enough -- with the pad in place it dominates - # `last_c_pos` for every residue, so the rest repeat one branch. Bases are - # page-aligned so the swa-page-boundary clause does not mask the pad. - bases = [512 + off for off in range(min(compress_ratio, 4))] - draft_counts = sorted( - {1, 2, 3, 4, 5, max_d - 1, max_d} & set(range(1, max_d + 1)) - ) - for num_draft_tokens in draft_counts: - for prefix_len in bases: - with self.subTest( - cr=compress_ratio, - D=num_draft_tokens, - prefix=prefix_len, - ): - written = self._make_plan_positions( - compress_ratio=compress_ratio, - ring_size=ring_size, - prefix_len=prefix_len, - num_draft_tokens=num_draft_tokens, - ) - seq_len = prefix_len + num_draft_tokens - missing = set(range(prefix_len, seq_len)) - written - self.assertEqual( - missing, - set(), - f"plan_w skipped committed positions {sorted(missing)}; " - f"a later compression would read stale ring slots", - ) - - def test_c4_ring_residency(self): - self._assert_ring_residency(4, C4_RING_SIZE) - - def test_c128_ring_residency(self): - self._assert_ring_residency(128, C128_RING_SIZE) - - def test_cpu_and_gpu_planner_agree(self): - """Both planner paths must emit the same write set. - - The residency invariants above are checked on the host-loop plan; this - pins the GPU `plan_compress_prefill_kernel0` plan to it, so the pad fix - has to hold on both paths. - """ - for compress_ratio, ring_size in ((4, C4_RING_SIZE), (128, C128_RING_SIZE)): - max_d = _max_draft_tokens(compress_ratio, ring_size) - for num_draft_tokens in (1, 4, max_d): - for prefix_len in (512, 513, 515): - with self.subTest( - cr=compress_ratio, D=num_draft_tokens, prefix=prefix_len - ): - kwargs = dict( - compress_ratio=compress_ratio, - ring_size=ring_size, - prefix_len=prefix_len, - num_draft_tokens=num_draft_tokens, - ) - self.assertEqual( - self._make_plan_positions(**kwargs, on_gpu=False), - self._make_plan_positions(**kwargs, on_gpu=True), - ) - - def test_plain_prefill_write_set(self): - """A non-speculative ring is exactly one window wide, so the pad is 0 and the - base write rule stands unchanged for both ratios.""" - for compress_ratio, ring_size in ( - (4, C4_RING_SIZE_NO_SPEC), - (128, C128_RING_SIZE_NO_SPEC), - ): - self.assertEqual(_max_draft_tokens(compress_ratio, ring_size), 0) - is_overlap = compress_ratio == 4 - for seq_len in (512, 600, 777): - with self.subTest(cr=compress_ratio, sl=seq_len): - ctx = make_paged_context( - bs=1, compress_ratio=compress_ratio, ring_size=ring_size - ) - seq_lens, extend_lens, num_q = to_seq_extend([(seq_len, seq_len)]) - plan = ctx.make_prefill_plan(seq_lens, extend_lens, num_q) - written = _written_positions(plan.plan_w, 0) - - last_c_pos = seq_len // compress_ratio * compress_ratio - first_w_pos = last_c_pos - (compress_ratio if is_overlap else 0) - sps = ctx.swa_page_size - expected = { - p - for p in range(seq_len) - if p >= first_w_pos - or (is_overlap and p % sps >= sps - compress_ratio) - } - self.assertEqual(written, expected) - - def test_over_capacity_under_writes(self): - """Beyond the ring's capacity the plan silently under-writes. - - The planner cannot tell an over-configured verify batch from an ordinary long - prefill, so it cannot fail loudly -- hence the startup check in - `DSV4PoolConfigurator._assert_ring_serves_draft_tokens`. - """ - for compress_ratio, ring_size in ((4, C4_RING_SIZE), (128, C128_RING_SIZE)): - max_d = _max_draft_tokens(compress_ratio, ring_size) - # Far enough over that the `last_c_pos` term cannot cover the gap for any - # residue of `seq_len % compress_ratio`. - too_many = max_d + compress_ratio + 1 - prefix_len = 512 - with self.subTest(cr=compress_ratio, D=too_many): - written = self._make_plan_positions( - compress_ratio=compress_ratio, - ring_size=ring_size, - prefix_len=prefix_len, - num_draft_tokens=too_many, - ) - missing = set(range(prefix_len, prefix_len + too_many)) - written - self.assertNotEqual( - missing, - set(), - "expected the plan to under-write past the ring capacity; if this " - "now covers everything, the startup bound can be relaxed", - ) - if __name__ == "__main__": unittest.main() From 1b94c719327f0ea15f5a1257dcc465be9ed1f137 Mon Sep 17 00:00:00 2001 From: hnyls2002 Date: Tue, 15 Sep 2026 22:22:44 -0700 Subject: [PATCH 14/30] dsv4.1: restore C2 padding test entry point --- test/registered/kernel/attention/dsv4/test_c2_verify.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/test/registered/kernel/attention/dsv4/test_c2_verify.py b/test/registered/kernel/attention/dsv4/test_c2_verify.py index be637cf26670..947febfebfb2 100644 --- a/test/registered/kernel/attention/dsv4/test_c2_verify.py +++ b/test/registered/kernel/attention/dsv4/test_c2_verify.py @@ -1,5 +1,7 @@ """C2 graph-padding rows must leave every destination untouched.""" +import sys + import pytest import torch @@ -36,3 +38,7 @@ def test_padded_even_and_odd_rows_are_inert() -> None: assert torch.equal(result, out_before) assert torch.equal(kv_state, state_before) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v"])) From ed2981664636bd6afed05b187a64cd34a83e05a1 Mon Sep 17 00:00:00 2001 From: BBuf <1182563586@qq.com> Date: Wed, 16 Sep 2026 13:25:46 +0800 Subject: [PATCH 15/30] remove standalone dsv4 metadata tests --- .../kernel/attention/dsv4/test_c2_verify.py | 44 -- .../attention/dsv4/test_v41_kv_store.py | 626 ------------------ 2 files changed, 670 deletions(-) delete mode 100644 test/registered/kernel/attention/dsv4/test_c2_verify.py delete mode 100644 test/registered/kernel/attention/dsv4/test_v41_kv_store.py diff --git a/test/registered/kernel/attention/dsv4/test_c2_verify.py b/test/registered/kernel/attention/dsv4/test_c2_verify.py deleted file mode 100644 index 947febfebfb2..000000000000 --- a/test/registered/kernel/attention/dsv4/test_c2_verify.py +++ /dev/null @@ -1,44 +0,0 @@ -"""C2 graph-padding rows must leave every destination untouched.""" - -import sys - -import pytest -import torch - -from sglang.kernels.ops.attention.dsv4.c2 import c2_decode_norm -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=5, stage="base-b-kernel-unit", runner_config="1-gpu-large") - - -@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") -def test_padded_even_and_odd_rows_are_inert() -> None: - num_tokens, head_dim, ring_size = 2, 512, 2 - kv_input = torch.randn(num_tokens, 2 * head_dim, device="cuda") - kv_state = torch.randn(num_tokens * ring_size, 2 * head_dim, device="cuda") - state_before = kv_state.clone() - norm_weight = torch.randn(head_dim, dtype=torch.bfloat16, device="cuda") - positions = torch.tensor([2, 3], dtype=torch.int64, device="cuda") - req = torch.arange(num_tokens, dtype=torch.int64, device="cuda") - raw_out_loc = torch.zeros(num_tokens, dtype=torch.int32, device="cuda") - out = torch.full((num_tokens, head_dim), 123, dtype=torch.bfloat16, device="cuda") - out_before = out.clone() - - result = c2_decode_norm( - kv_input, - kv_state, - norm_weight, - positions, - req, - raw_out_loc, - 1e-6, - ring_size=ring_size, - out=out, - ) - - assert torch.equal(result, out_before) - assert torch.equal(kv_state, state_before) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/kernel/attention/dsv4/test_v41_kv_store.py b/test/registered/kernel/attention/dsv4/test_v41_kv_store.py deleted file mode 100644 index b6f428bf7cbf..000000000000 --- a/test/registered/kernel/attention/dsv4/test_v41_kv_store.py +++ /dev/null @@ -1,626 +0,0 @@ -"""Byte-exactness of the V4.1 (fp8 / fp4) FlashMLA KV cache store kernels. - -Every store kernel is compared byte for byte with the pure-torch quantizers of -the two formats (``torch_quant.quantize_k_cache_v41`` / ``_v41_fp4``), which -follow the decode kernel's own reference quantizer. -""" - -import unittest - -import torch - -from sglang.kernels.ops.attention.dsv4.kv_layout import KVLayout -from sglang.srt.layers.attention.dsv4 import torch_quant as tq -from sglang.test.ci.ci_register import register_cuda_ci -from sglang.test.test_utils import CustomTestCase - -register_cuda_ci(est_time=120, stage="base-b-kernel-unit", runner_config="4-gpu-b200") - -REFERENCE = { - KVLayout.V41: tq.quantize_k_cache_v41, - KVLayout.V41_FP4: tq.quantize_k_cache_v41_fp4, -} -DEQUANT = { - KVLayout.V41: tq.dequantize_k_cache_v41, - KVLayout.V41_FP4: tq.dequantize_k_cache_v41_fp4, -} -# One quantization step, relative: e4m3 has 3 mantissa bits, e2m1 one. -ONE_CODE_RTOL = {KVLayout.V41: 0.13, KVLayout.V41_FP4: 0.51} - - -def rope_tail( - x: torch.Tensor, freqs: torch.Tensor, rope_dim: int, inverse: bool = False -) -> torch.Tensor: - """Rotate the last rope_dim features of x [T, ..., D] with complex freqs [T, rope_dim // 2].""" - head, tail = x[..., :-rope_dim], x[..., -rope_dim:] - tc = torch.view_as_complex(tail.float().unflatten(-1, (-1, 2)).contiguous()) - f = freqs.conj() if inverse else freqs - f = f.view(x.shape[0], *([1] * (x.ndim - 2)), rope_dim // 2) - rotated = torch.view_as_real(tc * f).flatten(-2).to(x.dtype) - return torch.cat([head, rotated], dim=-1) - - -def _sm100(): - return ( - torch.cuda.is_available() - and torch.version.cuda is not None - and torch.cuda.get_device_capability()[0] >= 10 - ) - - -def token_rows(pages, layout, page_size, locs): - """The (data row, scale row) bytes of the tokens at ``locs``.""" - locs = locs.long() - page, offset = locs // page_size, locs % page_size - data_cols = torch.arange(layout.data_bytes, device=pages.device) - scale_cols = torch.arange(layout.scale_bytes, device=pages.device) - data = pages[page[:, None], offset[:, None] * layout.data_bytes + data_cols] - scale = pages[ - page[:, None], - layout.scale_offset(page_size) - + offset[:, None] * layout.scale_bytes - + scale_cols, - ] - return data, scale - - -def random_rows(n, generator, device="cuda"): - """bf16 rows over a wide dynamic range, with zero, negative-zero and tiny tiles.""" - x = torch.randn(n, 512, generator=generator, device=device, dtype=torch.bfloat16) - scale = torch.exp2( - torch.randint(-10, 6, (n, 1), generator=generator, device=device).float() - ) - x = (x * scale).to(torch.bfloat16) - if n >= 3: - x[0, :32] = 0 - x[1, 32:48] = -0.0 - x[2, 100] = -1e-10 - return x - - -def reference_pages(layout, page_size, num_pages, locs, values, page_bytes): - full = torch.zeros( - num_pages, page_size, 512, device=values.device, dtype=values.dtype - ) - full.view(-1, 512)[locs.long()] = values - return REFERENCE[layout](full, page_bytes=page_bytes) - - -@unittest.skipUnless(_sm100(), "the V4.1 KV layouts are SM100 kernels") -class TestV41KVStore(CustomTestCase): - def assert_tokens_equal(self, cache, ref, layout, page_size, locs): - got_data, got_scale = token_rows(cache, layout, page_size, locs) - exp_data, exp_scale = token_rows(ref, layout, page_size, locs) - self.assertTrue(torch.equal(got_scale, exp_scale), "scale rows differ") - self.assertTrue(torch.equal(got_data, exp_data), "data rows differ") - - def assert_rows_close(self, cache, ref, layout, page_size, locs): - """For rows that went through the kernel's fp32 RMSNorm: torch sums the - squares in another order, and the fp32 ulp this can cost is occasionally - kept by a bf16 rounding boundary and then by the quantizer. Allow a - one-code difference in a handful of elements.""" - got_data, got_scale = token_rows(cache, layout, page_size, locs) - exp_data, exp_scale = token_rows(ref, layout, page_size, locs) - self.assertGreater((got_scale == exp_scale).float().mean().item(), 0.999) - self.assertGreater((got_data == exp_data).float().mean().item(), 0.999) - deq_got = DEQUANT[layout](cache, page_size).view(-1, 512)[locs.long()].float() - deq_exp = DEQUANT[layout](ref, page_size).view(-1, 512)[locs.long()].float() - torch.testing.assert_close( - deq_got, deq_exp, rtol=ONE_CODE_RTOL[layout], atol=0.06 - ) - - def assert_untouched_zero(self, cache, layout, page_size, locs): - num_slots = cache.shape[0] * page_size - written = torch.zeros(num_slots, dtype=torch.bool, device=cache.device) - written[locs.long()] = True - others = torch.arange(num_slots, device=cache.device)[~written] - data, scale = token_rows(cache, layout, page_size, others) - self.assertEqual(int(data.sum()) + int(scale.sum()), 0) - - def test_fused_store_cache(self): - """Ragged page fills: a random subset of slots is written, the rest stays zero.""" - from sglang.kernels.ops.attention.dsv4.attn import fused_store_cache - - g = torch.Generator(device="cuda").manual_seed(0) - for layout in (KVLayout.V41, KVLayout.V41_FP4): - for page_size, num_pages, n in ((64, 9, 333), (256, 3, 500), (2, 40, 37)): - for idx_dtype in (torch.int32, torch.int64): - with self.subTest( - layout=layout.name, page_size=page_size, idx=idx_dtype - ): - x = random_rows(n, g) - locs = torch.randperm( - num_pages * page_size, generator=g, device="cuda" - )[:n].to(idx_dtype) - cache = torch.zeros( - num_pages, - layout.page_bytes(page_size), - dtype=torch.uint8, - device="cuda", - ) - fused_store_cache( - x, - cache, - locs, - page_size=page_size, - type="flashmla", - layout=layout, - ) - ref = reference_pages( - layout, page_size, num_pages, locs, x, cache.shape[1] - ) - self.assert_tokens_equal(cache, ref, layout, page_size, locs) - self.assert_untouched_zero(cache, layout, page_size, locs) - - def test_fused_store_cache_with_rope(self): - """The in-kernel RoPE tail equals rope_tail (bf16-rounded) before quantizing, - so the fp4 cache holds exactly fake_quant_compressed_kv(rope_tail(x)).""" - from sglang.kernels.ops.attention.dsv4.attn import fused_store_cache - - g = torch.Generator(device="cuda").manual_seed(1) - for layout in (KVLayout.V41, KVLayout.V41_FP4): - for page_size, num_pages, n in ((64, 5, 200), (256, 2, 129)): - with self.subTest(layout=layout.name, page_size=page_size): - x = random_rows(n, g) - angles = torch.randn(n, 32, generator=g, device="cuda") - freqs = torch.polar(torch.ones_like(angles), angles) - locs = torch.randperm( - num_pages * page_size, generator=g, device="cuda" - )[:n] - cache = torch.zeros( - num_pages, - layout.page_bytes(page_size), - dtype=torch.uint8, - device="cuda", - ) - fused_store_cache( - x, - cache, - locs, - page_size=page_size, - type="flashmla", - layout=layout, - freqs_cis=freqs, - ) - rotated = rope_tail(x, freqs, 64) - ref = reference_pages( - layout, page_size, num_pages, locs, rotated, cache.shape[1] - ) - self.assert_tokens_equal(cache, ref, layout, page_size, locs) - if layout is KVLayout.V41_FP4: - deq = tq.dequantize_k_cache_v41_fp4(cache, page_size).view( - -1, 512 - )[locs] - self.assertTrue( - torch.equal(deq, tq.fake_quant_compressed_kv(rotated)) - ) - - def test_boundary_tiles(self): - """Tie, saturation and subnormal-scale tiles follow the reference. (NaN / inf - rows are not fed: the kernels do not reproduce the reference's handling of them.)""" - from sglang.kernels.ops.attention.dsv4.attn import fused_store_cache - - maxima = [ - 0, - 2**-12, - 6 * 2**-9, - 6 * 1.0625, - 6 * 1.1875, - 6 * 448, - 1e6, - 8.25, - 448.0, - 449.0, - ] - rows = [] - for m in maxima: - row = torch.full((512,), m, dtype=torch.bfloat16) - row[1::2] *= -1 - row[64:80] = torch.tensor( - [8.25, 4.125, 3.4375, -4.8125, 0.0, -0.0, 2.5, -3.5] * 2, - dtype=torch.bfloat16, - ) - rows.append(row) - x = torch.stack(rows).cuda() - n = x.shape[0] - page_size = 16 - locs = torch.arange(n, device="cuda", dtype=torch.int32) - for layout in (KVLayout.V41, KVLayout.V41_FP4): - with self.subTest(layout=layout.name): - cache = torch.zeros( - 1, layout.page_bytes(page_size), dtype=torch.uint8, device="cuda" - ) - fused_store_cache( - x, cache, locs, page_size=page_size, type="flashmla", layout=layout - ) - ref = reference_pages(layout, page_size, 1, locs, x, cache.shape[1]) - self.assert_tokens_equal(cache, ref, layout, page_size, locs) - self.assert_untouched_zero(cache, layout, page_size, locs) - - def test_fused_k_norm_rope_store(self): - """The fused RMSNorm + RoPE + store (the SWA write) equals norm -> rope_tail -> - fused_store_cache: exact-norm rows bitwise against the torch quantizer, and - general rows bitwise against the unfused kernel chain.""" - from sglang.kernels.ops.attention.dsv4.attn import fused_store_cache - from sglang.kernels.ops.attention.dsv4.elementwise import ( - fused_k_norm_rope_flashmla, - ) - - g = torch.Generator(device="cuda").manual_seed(2) - page_size, num_pages, n = 256, 3, 300 - angles = torch.randn(1024, 32, generator=g, device="cuda") - freqs_table = torch.polar(torch.ones_like(angles), angles) - pos = torch.randint( - 0, 1024, (n,), generator=g, device="cuda", dtype=torch.int64 - ) - locs = torch.randperm(num_pages * page_size, generator=g, device="cuda")[:n].to( - torch.int32 - ) - locs[7] = -1 # a row without a write target is skipped - valid = locs >= 0 - for layout in (KVLayout.V41, KVLayout.V41_FP4): - with self.subTest(layout=layout.name, rows="exact-norm"): - # x = +-2^k per row with eps = 0 normalizes to +-1 exactly, so the - # normed row is sign * weight and the reference is exact. - signs = torch.where( - torch.rand(n, 512, generator=g, device="cuda") < 0.5, -1.0, 1.0 - ) - k = torch.randint(-6, 6, (n, 1), generator=g, device="cuda").float() - x = (signs * torch.exp2(k)).to(torch.bfloat16) - w = ( - torch.randn(512, generator=g, device="cuda") - * torch.exp2( - torch.randint(-6, 4, (512,), generator=g, device="cuda").float() - ) - ).to(torch.bfloat16) - cache = torch.zeros( - num_pages, - layout.page_bytes(page_size), - dtype=torch.uint8, - device="cuda", - ) - fused_k_norm_rope_flashmla( - x, w, 0.0, freqs_table, pos, locs, cache, page_size, layout=layout - ) - rotated = rope_tail( - (signs * w.float()).to(torch.bfloat16), freqs_table[pos], 64 - ) - ref = reference_pages( - layout, - page_size, - num_pages, - locs[valid], - rotated[valid], - cache.shape[1], - ) - self.assert_tokens_equal(cache, ref, layout, page_size, locs[valid]) - self.assert_untouched_zero(cache, layout, page_size, locs[valid]) - with self.subTest(layout=layout.name, rows="general"): - x = (torch.randn(n, 512, generator=g, device="cuda") * 3).to( - torch.bfloat16 - ) - w = (torch.randn(512, generator=g, device="cuda") * 0.3 + 1).to( - torch.bfloat16 - ) - eps = 1e-6 - cache = torch.zeros( - num_pages, - layout.page_bytes(page_size), - dtype=torch.uint8, - device="cuda", - ) - fused_k_norm_rope_flashmla( - x, w, eps, freqs_table, pos, locs, cache, page_size, layout=layout - ) - xf = x.float() - normed = ( - xf - * torch.rsqrt(xf.square().mean(-1, keepdim=True) + eps) - * w.float() - ).to(torch.bfloat16) - rotated = rope_tail(normed, freqs_table[pos], 64) - unfused = torch.zeros_like(cache) - fused_store_cache( - rotated[valid], - unfused, - locs[valid], - page_size=page_size, - type="flashmla", - layout=layout, - ) - self.assert_rows_close(cache, unfused, layout, page_size, locs[valid]) - - def test_c1_c2_decode_store(self): - """The ratio-1 / ratio-2 decode compressors write the V4.1 layouts: the cache - holds the quantized rope_tail of the pre-RoPE latent the kernel publishes - (bitwise; the fp8 layout after the model's fp4 fake quantization), and the - latent is the torch RMSNorm to within an fp32-reduction-order bf16 ulp.""" - from sglang.kernels.ops.attention.dsv4.c1 import c1_decode_norm_rope_store - from sglang.kernels.ops.attention.dsv4.c2 import ( - c2_decode_or_verify_norm_rope_store, - ) - - g = torch.Generator(device="cuda").manual_seed(4) - eps = 1e-6 - angles = torch.randn(4096, 32, generator=g, device="cuda") - freqs = torch.polar(torch.ones_like(angles), angles) - freqs_real = torch.view_as_real(freqs).flatten(-2) - w = (torch.randn(512, generator=g, device="cuda") * 0.3 + 1).to(torch.bfloat16) - - def torch_norm(x): - xf = x.float() - return ( - xf * torch.rsqrt(xf.square().mean(-1, keepdim=True) + eps) * w.float() - ).to(torch.bfloat16) - - def stored_reference(layout, rotated, page_size, num_pages, locs, page_bytes): - values = ( - rotated - if layout is KVLayout.V41_FP4 - else tq.fake_quant_compressed_kv(rotated) - ) - return reference_pages( - layout, page_size, num_pages, locs, values, page_bytes - ) - - n = 200 - for layout in (KVLayout.V41, KVLayout.V41_FP4): - for page_size, num_pages in ((256, 4), (128, 8)): - with self.subTest(kernel="c1", layout=layout.name, page_size=page_size): - x = (torch.randn(n, 512, generator=g, device="cuda") * 2).to( - torch.bfloat16 - ) - pos = torch.randint( - 0, 4096, (n,), generator=g, device="cuda", dtype=torch.int64 - ) - out_loc = ( - torch.randperm( - num_pages * page_size - 1, generator=g, device="cuda" - )[:n].to(torch.int32) - + 1 - ) - out_loc[3] = 0 # a padded graph row publishes nothing - cache = torch.zeros( - num_pages, - layout.page_bytes(page_size), - dtype=torch.uint8, - device="cuda", - ) - latent = c1_decode_norm_rope_store( - x, - w, - pos, - out_loc, - eps, - freqs_real, - cache, - page_size=page_size, - layout=layout, - ) - torch.testing.assert_close( - latent, torch_norm(x), rtol=2**-7, atol=2**-14 - ) - valid = out_loc > 0 - rotated = rope_tail(latent, freqs[pos], 64) - ref = stored_reference( - layout, - rotated[valid], - page_size, - num_pages, - out_loc[valid], - cache.shape[1], - ) - self.assert_tokens_equal( - cache, ref, layout, page_size, out_loc[valid] - ) - self.assert_untouched_zero(cache, layout, page_size, out_loc[valid]) - with self.subTest(kernel="c2", layout=layout.name, page_size=page_size): - ring = 2 - kv_new = torch.randn(n, 512, generator=g, device="cuda") * 2 - score = torch.randn(n, 512, generator=g, device="cuda") - kv_old = torch.randn(n, 512, generator=g, device="cuda") * 2 - kv_input = torch.cat([kv_new, score], dim=-1).contiguous() - req = torch.arange(n, device="cuda", dtype=torch.int64) - # Odd positions complete a pair; row 5 is even and rows 5/7 - # are padded graph rows that must remain completely inert. - pos = ( - 2 - * torch.randint( - 0, 2000, (n,), generator=g, device="cuda", dtype=torch.int64 - ) - + 1 - ) - pos[5] = 4 - state = torch.randn(n * ring + 4, 1024, generator=g, device="cuda") - read_rows = req * ring + (pos - 1) % ring - # Equal scores make the pair pool the exact mean. - state[read_rows, :512] = kv_old - state[read_rows, 512:] = score - raw_out_loc = ( - torch.randperm( - num_pages * page_size - 1, generator=g, device="cuda" - )[:n].to(torch.int32) - + 1 - ) * 2 - raw_out_loc[[5, 7]] = 0 - cache = torch.zeros( - num_pages, - layout.page_bytes(page_size), - dtype=torch.uint8, - device="cuda", - ) - latent_out = torch.full( - (n, 512), 123, dtype=torch.bfloat16, device="cuda" - ) - padded_out_before = latent_out[[5, 7]].clone() - padded_state_row = req[5] * ring + pos[5] % ring - padded_state_before = state[padded_state_row].clone() - latent = c2_decode_or_verify_norm_rope_store( - kv_input, - state, - w, - pos, - req, - raw_out_loc, - eps, - freqs_real, - cache, - page_size=page_size, - ring_size=ring, - layout=layout, - out=latent_out, - ) - self.assertTrue(torch.equal(latent[[5, 7]], padded_out_before)) - self.assertTrue( - torch.equal(state[padded_state_row], padded_state_before) - ) - valid = (raw_out_loc != 0) & (pos % 2 == 1) - pooled = ((kv_old + kv_new) / 2).to(torch.bfloat16) - torch.testing.assert_close( - latent[valid], - torch_norm(pooled)[valid], - rtol=2**-7, - atol=2**-14, - ) - rotated = rope_tail(latent, freqs[(pos - 1).clamp_min(0)], 64) - slots = raw_out_loc >> 1 - ref = stored_reference( - layout, - rotated[valid], - page_size, - num_pages, - slots[valid], - cache.shape[1], - ) - self.assert_tokens_equal( - cache, ref, layout, page_size, slots[valid] - ) - self.assert_untouched_zero(cache, layout, page_size, slots[valid]) - - def test_compress_norm_rope_store(self): - """The ratio-4 / ratio-128 writer (norm + RoPE + store from a decode plan) in - the V4.1 layouts: bitwise on exact-norm rows, one-code close on general rows.""" - from sglang.kernels.ops.attention.dsv4.compress import ( - CompressorDecodePlan, - compress_norm_rope_store, - ) - - g = torch.Generator(device="cuda").manual_seed(5) - ratio = 4 - angles = torch.randn(4096, 32, generator=g, device="cuda") - freqs = torch.polar(torch.ones_like(angles), angles) - n = 150 - for layout in (KVLayout.V41, KVLayout.V41_FP4): - for page_size, num_pages in ((64, 8), (2, 300)): - seq_lens = ( - torch.randint( - 1, 1000, (n,), generator=g, device="cuda", dtype=torch.int64 - ) - * ratio - ) - seq_lens[3] += 1 # not a group boundary: writes nothing - plan = CompressorDecodePlan.generate_legacy( - ratio, torch.arange(n, device="cuda", dtype=torch.int64), seq_lens - ) - valid = seq_lens % ratio == 0 - pos = seq_lens - ratio - out_loc = torch.randperm( - num_pages * page_size, generator=g, device="cuda" - )[:n].to(torch.int64) - with self.subTest( - layout=layout.name, page_size=page_size, rows="exact-norm" - ): - signs = torch.where( - torch.rand(n, 512, generator=g, device="cuda") < 0.5, -1.0, 1.0 - ) - k = torch.randint(-6, 6, (n, 1), generator=g, device="cuda").float() - kv = (signs * torch.exp2(k)).to(torch.bfloat16) - w = ( - torch.randn(512, generator=g, device="cuda") - * torch.exp2( - torch.randint( - -6, 4, (512,), generator=g, device="cuda" - ).float() - ) - ).to(torch.bfloat16) - cache = torch.zeros( - num_pages, - layout.page_bytes(page_size), - dtype=torch.uint8, - device="cuda", - ) - compress_norm_rope_store( - kv, - plan, - norm_weight=w, - norm_eps=0.0, - freq_cis=freqs, - out_loc=out_loc, - kvcache=cache, - page_size=page_size, - layout=layout, - ) - rotated = rope_tail( - (signs * w.float()).to(torch.bfloat16), freqs[pos], 64 - ) - ref = reference_pages( - layout, - page_size, - num_pages, - out_loc[valid], - rotated[valid], - cache.shape[1], - ) - self.assert_tokens_equal( - cache, ref, layout, page_size, out_loc[valid] - ) - self.assert_untouched_zero(cache, layout, page_size, out_loc[valid]) - with self.subTest( - layout=layout.name, page_size=page_size, rows="general" - ): - eps = 1e-6 - kv = (torch.randn(n, 512, generator=g, device="cuda") * 2).to( - torch.bfloat16 - ) - w = (torch.randn(512, generator=g, device="cuda") * 0.3 + 1).to( - torch.bfloat16 - ) - cache = torch.zeros( - num_pages, - layout.page_bytes(page_size), - dtype=torch.uint8, - device="cuda", - ) - compress_norm_rope_store( - kv, - plan, - norm_weight=w, - norm_eps=eps, - freq_cis=freqs, - out_loc=out_loc, - kvcache=cache, - page_size=page_size, - layout=layout, - ) - xf = kv.float() - normed = ( - xf - * torch.rsqrt(xf.square().mean(-1, keepdim=True) + eps) - * w.float() - ).to(torch.bfloat16) - rotated = rope_tail(normed, freqs[pos], 64) - ref = reference_pages( - layout, - page_size, - num_pages, - out_loc[valid], - rotated[valid], - cache.shape[1], - ) - self.assert_rows_close( - cache, ref, layout, page_size, out_loc[valid] - ) - - -if __name__ == "__main__": - unittest.main() From 1333f2a5ff96f9aae7a0dfa770c839f28f024447 Mon Sep 17 00:00:00 2001 From: hnyls2002 Date: Tue, 15 Sep 2026 15:53:42 -0700 Subject: [PATCH 16/30] dsv4.1: extract candidate indexer library --- .../attention/dsv4/candidate_deep_gemm.py | 273 ++++++++++++++++++ .../attention/dsv4/candidate_indexer.py | 72 +++++ .../layers/attention/dsv4/candidate_torch.py | 208 +++++++++++++ .../srt/layers/attention/dsv4/indexer.py | 85 ++++++ .../layers/deep_gemm_wrapper/configurer.py | 17 ++ .../dsv4/test_dsv41_sparse_indexer.py | 234 +++++++++++++++ 6 files changed, 889 insertions(+) create mode 100644 python/sglang/srt/layers/attention/dsv4/candidate_deep_gemm.py create mode 100644 python/sglang/srt/layers/attention/dsv4/candidate_indexer.py create mode 100644 python/sglang/srt/layers/attention/dsv4/candidate_torch.py create mode 100644 test/registered/kernel/attention/dsv4/test_dsv41_sparse_indexer.py diff --git a/python/sglang/srt/layers/attention/dsv4/candidate_deep_gemm.py b/python/sglang/srt/layers/attention/dsv4/candidate_deep_gemm.py new file mode 100644 index 000000000000..c2d11b81f5cc --- /dev/null +++ b/python/sglang/srt/layers/attention/dsv4/candidate_deep_gemm.py @@ -0,0 +1,273 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Optional + +import torch + +from sglang.kernels.ops.attention.dsv4.candidate_blocks import candidate_row_lens +from sglang.kernels.ops.attention.dsv4.topk import ( + amax8_varlen, + plan_topk_v2, + sort_candidate_blocks, + topk_transform_bf16_small, + topk_transform_paged_v2, +) +from sglang.srt.layers.attention.dsv4.candidate_indexer import ( + CandidateMetadata, + IndexerInputs, +) +from sglang.srt.layers.attention.dsv4.indexer import ( + fp4_paged_mqa_logits, +) + +CANDIDATE_BLOCK_SIZE = 8 # positions per block; DeepGEMM accepts 8 or 16 + + +@dataclass +class SparseBlockTable(CandidateMetadata): + # [rows, topk_blocks] int32: ascending logical block ids, valid for the first + # min(topk_blocks, ceil(seq_len / 8)) entries of a row; DeepGEMM reads only those + blocks: torch.Tensor + # DeepGEMM's schedule metadata (uint8) for them + schedule: torch.Tensor + # [rows, topk_blocks] int32: the same blocks as pool slots / 8, so a consumer's + # top-k maps column j of the sparse row to slot phys_blocks[b, j // 8] * 8 + j % 8 + # with the plain page-table transform at page size 8 + phys_blocks: torch.Tensor + # [rows] int32: length of each row of the sparse logits (see `valid_lens`) + valid_lens: torch.Tensor + + +def valid_lens(seq_lens: torch.Tensor, topk_blocks: int) -> torch.Tensor: + """Length of each row of the sparse logits: the published blocks laid out + block by block, the newest (highest) block possibly partial. Torch reference + of ``candidate_row_lens``; the decode path takes the kernel's value.""" + block = CANDIDATE_BLOCK_SIZE + num = ((seq_lens + block - 1) // block).clamp_max(topk_blocks) + return block * (num - 1) + (seq_lens - 1) % block + 1 + + +def amax_topk_blocks( + logits: torch.Tensor, + seq_lens: torch.Tensor, + nblocks: torch.Tensor, + topk_blocks: int, + max_seq_len: Optional[int] = None, +) -> torch.Tensor: + """Per row the ``topk_blocks`` blocks of 8 positions with the largest block + maximum among its first ``seq_lens[b]`` positions, the newest block always + included: block ids in no particular order, ``-1`` past the row's count + (``sort_candidate_blocks`` turns them into the published table). ``nblocks`` + is ``ceil(seq_lens / 8)`` as int32, from ``candidate_row_lens``.""" + rows = logits.shape[0] + block = CANDIDATE_BLOCK_SIZE + if max_seq_len is None: + max_seq_len = logits.shape[1] + # NOTE: plan cannot be the previous kernel of topk_transform_paged_v2 + plan = plan_topk_v2(nblocks) + # block maxima, the newest block +inf; the top-k reads each row up to nblocks + # only, so nothing past a row's keys is initialised (v2 needs stride % 4 == 0) + keys = logits.new_empty(rows, -(-max_seq_len // (4 * block)) * 4) + amax8_varlen(logits, seq_lens, out=keys) + blocks = torch.empty(rows, topk_blocks, dtype=torch.int32, device=logits.device) + topk_transform_paged_v2(keys, nblocks, None, blocks, 1, plan) + return blocks + + +_ROW_IDS: dict = {} + + +def _row_ids(rows: int, device: torch.device) -> torch.Tensor: + """``arange(rows)`` int32 from a cached buffer (grown in steps of 8192), so the + every-row-its-own-request case costs no launch.""" + buf = _ROW_IDS.get(device) + if buf is None or buf.numel() < rows: + size = max(8192, -(-rows // 8192) * 8192) + buf = _ROW_IDS[device] = torch.arange(size, dtype=torch.int32, device=device) + return buf[:rows] + + +def build_sparse_indexer_schedule( + blocks: torch.Tensor, + seq_lens: torch.Tensor, + page_table: torch.Tensor, + page_size: int, + q_dtype: torch.dtype, + request_ids: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """DeepGEMM's schedule for the published blocks: ``seq_lens`` ``[rows]`` + int32, ``page_table`` ``[rows, pages]`` int32 at the index pool's page size. + ``request_ids`` ``[rows]``, one id per row with a request's rows consecutive + (verify: its draft tokens), lets DeepGEMM pair two rows of a request on one + KV pass; each row keeps its own block list and output layout. Paired rows + must share their page-table row. None: every row is its own request.""" + import deep_gemm + + rows = blocks.shape[0] + if request_ids is None: + request_ids = _row_ids(rows, blocks.device) # cached, no launch + else: + # the scheduler keeps request indices as int64; one small cast per publish + request_ids = request_ids[:rows].to(torch.int32).contiguous() + return deep_gemm.get_paged_sparse_mqa_logits_metadata( + seq_lens.contiguous(), + page_table, + request_ids, + page_size, + blocks, + q_dtype, + CANDIDATE_BLOCK_SIZE, + ) + + +def sparse_logits( + q_fp4: torch.Tensor, + q_sf: torch.Tensor, + k_cache: torch.Tensor, + weights: torch.Tensor, + table: SparseBlockTable, +) -> torch.Tensor: + """bf16 logits ``[rows, topk_blocks * 8]`` of the published blocks: ``q_fp4`` + ``[rows, 1, heads, 64]`` int8 with ``q_sf`` ``[rows, 1, heads]`` int32 (packed + ue8m0), ``k_cache`` ``[pages, page_size, 1, 68]`` uint8 whose page stride is + a multiple of 512 bytes, ``weights`` ``[rows, heads]`` bf16.""" + import deep_gemm + + return deep_gemm.fp8_fp4_paged_sparse_mqa_logits( + (q_fp4, q_sf), + k_cache, + weights, + table.schedule, + table.blocks.shape[1], + CANDIDATE_BLOCK_SIZE, + ) + + +def topk_transform_sparse( + logits: torch.Tensor, + valid_lens: torch.Tensor, + table: SparseBlockTable, + page_indices: torch.Tensor, +) -> None: + """Top-``k`` (``k = page_indices.shape[1]``) of every row of the sparse + ``logits`` (bf16 ``[rows, topk_blocks * 8]``) within its first ``valid_lens[b]`` + columns, written as pool slots, ``-1`` where a row has fewer than ``k`` valid + columns, in no particular order. Column ``j`` is slot + ``phys_blocks[b, j // 8] * 8 + j % 8``: the bf16 top-k kernel's page-table + transform at page size 8 over the physical block table.""" + topk_transform_bf16_small( + logits, valid_lens, table.phys_blocks, page_indices, CANDIDATE_BLOCK_SIZE + ) + + +def warmup(rows: int, topk_blocks: int, page_size: int, q_dtype: torch.dtype, device): + """Build one schedule so DeepGEMM allocates its per-stream metadata workspace + outside any CUDA-graph capture.""" + seq_lens = torch.full( + (rows,), CANDIDATE_BLOCK_SIZE, dtype=torch.int32, device=device + ) + blocks = torch.zeros(rows, topk_blocks, dtype=torch.int32, device=device) + page_table = torch.zeros(rows, 1, dtype=torch.int32, device=device) + build_sparse_indexer_schedule(blocks, seq_lens, page_table, page_size, q_dtype) + + +class DeepGemmCandidateIndexer: + def __init__(self, topk_blocks: int, block_size: int): + assert block_size == CANDIDATE_BLOCK_SIZE, block_size + self.topk_blocks = topk_blocks + self.block_size = block_size + self.alt_stream = torch.cuda.Stream() + self.need_wait = False + + def publish_decode( + self, + inputs: IndexerInputs, + page_indices: torch.Tensor, + raw_indices: Optional[torch.Tensor] = None, + ) -> SparseBlockTable: + """Layer 20 end to end: dense logits, its own plain top-k into + ``page_indices`` (and ``raw_indices`` when given), the block table from + the same logits and DeepGEMM's schedule for it; the backend stores the + table on the forward metadata.""" + metadata = inputs.metadata + seq_lens = metadata.compressed_seq_lens.reshape(-1) + logits = fp4_paged_mqa_logits( + (inputs.q_fp4, inputs.q_sf), + inputs.k_cache, + inputs.weights, + metadata.compressed_seq_lens, + metadata.page_table, + metadata.deep_gemm_metadata, + metadata.max_compressed_seq_len, + ) + self.alt_stream.wait_stream(torch.cuda.current_stream()) + # The block-selection chain reads logits after the main stream moves on. + logits.record_stream(self.alt_stream) + self.need_wait = True + # TODO(candidate): one kernel for both selections below (dense logits read once) + topk_transform_paged_v2( + logits, + seq_lens, + metadata.page_table, + page_indices, + metadata.compressed_page_size, + metadata.topk_metadata, + out_raw_indices=raw_indices, + ) + # per row: block count for the block top-k, sparse-row length for the consumers + with torch.cuda.stream(self.alt_stream): + nblocks, row_valid_lens = candidate_row_lens(seq_lens, self.topk_blocks) + blocks = amax_topk_blocks(logits, seq_lens, nblocks, self.topk_blocks) + # in place: ascending, INT32_MAX padded, plus the blocks as pool slots / 8 + phys_blocks = sort_candidate_blocks( + blocks, + seq_lens, + metadata.page_table, + metadata.compressed_page_size, + ) + schedule = build_sparse_indexer_schedule( + blocks, + seq_lens, + metadata.page_table, + metadata.compressed_page_size, + inputs.q_fp4.dtype, + inputs.request_ids, + ) + return SparseBlockTable( + blocks=blocks, + schedule=schedule, + phys_blocks=phys_blocks, + valid_lens=row_valid_lens, + ) + + def scores(self, table: SparseBlockTable, inputs: IndexerInputs) -> torch.Tensor: + """A consumer: bf16 logits ``[rows, topk_blocks * 8]`` of the published + blocks only.""" + return sparse_logits( + inputs.q_fp4, + inputs.q_sf, + inputs.k_cache, + inputs.weights.to(torch.bfloat16), + table, + ) + + def select_decode( + self, + candidate_metadata: SparseBlockTable, + inputs: IndexerInputs, + page_indices: torch.Tensor, + raw_indices: Optional[torch.Tensor] = None, + ) -> None: + """A consumer: top-``k`` (``k = page_indices.shape[1]``) inside the + published blocks, written ascending as slots through the page table with + ``-1`` past the valid count (and as positions into ``raw_indices`` when + given).""" + assert raw_indices is None + table = candidate_metadata + if self.need_wait: + self.need_wait = False + torch.cuda.current_stream().wait_stream(self.alt_stream) + logits = self.scores(table, inputs) + # decode carries no raw_indices; the kernel writes slots only + topk_transform_sparse(logits, table.valid_lens, table, page_indices) diff --git a/python/sglang/srt/layers/attention/dsv4/candidate_indexer.py b/python/sglang/srt/layers/attention/dsv4/candidate_indexer.py new file mode 100644 index 000000000000..fea5e4b20acb --- /dev/null +++ b/python/sglang/srt/layers/attention/dsv4/candidate_indexer.py @@ -0,0 +1,72 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Generic, Optional, Protocol, TypeVar + +import torch + +from sglang.srt.layers.attention.dsv4.metadata import PagedIndexerMetadata + + +class CandidateMetadata: + """Base of an implementation's published state on + ``DSV4Metadata.candidate_metadata``.""" + + +@dataclass(frozen=True) +class IndexerInputs: + """One index-source layer's operands on the paged fp4 decode path (one query + row per request, or per draft token under verify).""" + + q_fp4: torch.Tensor # [rows, 1, heads, 64] int8, packed fp4 + q_sf: torch.Tensor # [rows, 1, heads] int32, packed ue8m0 + k_cache: torch.Tensor # [pages, page_size, 1, 68] uint8, the layer's index-K pool + weights: torch.Tensor # [rows, heads] bf16/fp32 head weights + metadata: PagedIndexerMetadata # this ratio's lengths, page table and plans + # [rows] int, one request id per query row, the rows of one request + # consecutive (verify: its draft tokens); None = every row its own request + request_ids: Optional[torch.Tensor] = None + + @property + def num_rows(self) -> int: + return self.q_fp4.shape[0] + + +T = TypeVar("T", bound=CandidateMetadata) + + +# TODO(dark): support publish prefill/select prefill +# TODO(dark): support fusion of publish + topk of publish layer +class CandidateIndexer(Protocol, Generic[T]): + def publish_decode( + self, + inputs: IndexerInputs, + page_indices: torch.Tensor, + raw_indices: Optional[torch.Tensor] = None, + ) -> T: ... + def select_decode( + self, + candidate_metadata: T, + inputs: IndexerInputs, + page_indices: torch.Tensor, + raw_indices: Optional[torch.Tensor] = None, + ) -> None: ... + + +def make_candidate_indexer(topk_blocks: int, block_size: int) -> CandidateIndexer: + """Use the sparse indexer when the installed DeepGEMM provides its APIs.""" + from sglang.srt.layers.deep_gemm_wrapper.configurer import ( + DEEPGEMM_SPARSE_INDEXER, + ) + + if DEEPGEMM_SPARSE_INDEXER and block_size == 8 and topk_blocks > 0: + from sglang.srt.layers.attention.dsv4.candidate_deep_gemm import ( + DeepGemmCandidateIndexer, + ) + + return DeepGemmCandidateIndexer(topk_blocks, block_size) + from sglang.srt.layers.attention.dsv4.candidate_torch import ( + TorchCandidateIndexer, + ) + + return TorchCandidateIndexer(topk_blocks, block_size) diff --git a/python/sglang/srt/layers/attention/dsv4/candidate_torch.py b/python/sglang/srt/layers/attention/dsv4/candidate_torch.py new file mode 100644 index 000000000000..a2bcb003c218 --- /dev/null +++ b/python/sglang/srt/layers/attention/dsv4/candidate_torch.py @@ -0,0 +1,208 @@ +# TODO(candidate): retire with the last path that still selects through masks +# (Hopper decode, prefill); the paged fp4 decode path already has the DeepGEMM one. +from __future__ import annotations + +from dataclasses import dataclass +from typing import List, Optional, Tuple + +import torch + +from sglang.srt.layers.attention.dsv4.candidate_indexer import ( + CandidateMetadata, + IndexerInputs, +) +from sglang.srt.layers.attention.dsv4.indexer import ( + fp4_paged_mqa_logits, + fp32_jit_paged_topk, + select_candidate_blocks, +) + + +@dataclass +class CandidateMasks(CandidateMetadata): + mask: Optional[torch.Tensor] = None # decode: [rows, width] bool + request_masks: Optional[List[torch.Tensor]] = None # prefill: [rows_b, lc_b] each + + +def published_masks(candidate) -> CandidateMasks: + """The forward's ``candidate_metadata`` as the masks the source published.""" + assert isinstance(candidate, CandidateMasks), "candidate masks missing" + return candidate + + +def two_level_decode_logits( + logits: torch.Tensor, + seq_lens: torch.Tensor, + *, + is_candidate_source: bool, + uses_candidates: bool, + topk_blocks: int, + block_size: int, + published: Optional[torch.Tensor], +) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: + """Apply candidate-block filtering and return logits plus an optional published mask. + + Mask columns past each sequence length to -inf before selection: the paged + logits kernel leaves that tail uninitialized, and an all -inf block means + unreachable. This path is graph-captured and must not synchronize with the host. + Work scales with allocated page-table capacity, not live sequence length. + """ + if not (is_candidate_source or uses_candidates): + return logits, None + + if ( + logits.is_cuda + and torch.version.cuda is not None + and logits.ndim == 2 + and logits.stride(1) == 1 + and seq_lens.device == logits.device + and seq_lens.dtype in (torch.int32, torch.int64) + and seq_lens.is_contiguous() + and seq_lens.shape in ((logits.shape[0],), (logits.shape[0], 1)) + and logits.numel() > 0 + and 0 < block_size <= 1024 + ): + from sglang.kernels.ops.attention.dsv4.candidate_blocks import ( + candidate_block_logits, + ) + + if not is_candidate_source: + assert ( + torch.is_tensor(published) + and published.shape[0] == logits.shape[0] + and published.shape[1] >= logits.shape[1] + ), "candidate mask missing for decode" + return candidate_block_logits( + logits, + seq_lens, + topk_blocks=topk_blocks, + block_size=block_size, + published=None if is_candidate_source else published, + ) + + lens_col = seq_lens if seq_lens.dim() > 1 else seq_lens.unsqueeze(-1) + reachable = torch.arange(logits.shape[-1], device=logits.device) < lens_col + logits = logits.float().masked_fill(~reachable, -torch.inf) + + if is_candidate_source: + # The source scores over every reachable position itself and only publishes, + # which is what the reference does. + return logits, select_candidate_blocks( + logits, lens_col, topk_blocks=topk_blocks, block_size=block_size + ) + + assert torch.is_tensor(published) and published.shape[0] == logits.shape[0], ( + "candidate mask missing for decode" + ) + return logits.masked_fill(~published[:, : logits.shape[-1]], -torch.inf), None + + +def mask_topk_scores( + scores: torch.Tensor, + indices: torch.Tensor, + offsets: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """Keep masked indexer scores out of attention even when top-k underfills.""" + columns = indices.to(torch.int64) + if offsets is not None: + columns = columns - offsets[:, None] + selected_scores = scores.gather(1, columns.clamp(0, scores.shape[1] - 1)) + valid = ( + (columns >= 0) & (columns < scores.shape[1]) & (selected_scores > -torch.inf) + ) + return indices.masked_fill(~valid, -1) + + +class TorchCandidateIndexer: + def __init__(self, topk_blocks: int, block_size: int): + self.topk_blocks = topk_blocks + self.block_size = block_size + + def publish_decode( + self, + inputs: IndexerInputs, + page_indices: torch.Tensor, + raw_indices: Optional[torch.Tensor] = None, + ) -> CandidateMasks: + """Layer 20 end to end: dense logits, its own plain top-k into + ``page_indices`` (and ``raw_indices`` when given), and the mask for the + layers after it; the backend stores the mask on the forward metadata.""" + metadata = inputs.metadata + logits = fp4_paged_mqa_logits( + (inputs.q_fp4, inputs.q_sf), + inputs.k_cache, + inputs.weights, + metadata.compressed_seq_lens, + metadata.page_table, + metadata.deep_gemm_metadata, + metadata.max_compressed_seq_len, + ) + logits, mask = two_level_decode_logits( + logits, + metadata.compressed_seq_lens, + is_candidate_source=True, + uses_candidates=False, + topk_blocks=self.topk_blocks, + block_size=self.block_size, + published=None, + ) + fp32_jit_paged_topk(logits, metadata, page_indices, raw_indices) + return CandidateMasks(mask=mask) + + def select_decode( + self, + candidate_metadata: CandidateMasks, + inputs: IndexerInputs, + page_indices: torch.Tensor, + raw_indices: Optional[torch.Tensor] = None, + ) -> None: + """A consumer: dense logits, masked, top-k, written as slots through the + page table with ``-1`` past the valid count (and as positions into + ``raw_indices`` when given).""" + metadata = inputs.metadata + page_size = metadata.compressed_page_size + assert isinstance(candidate_metadata, CandidateMasks) + logits = fp4_paged_mqa_logits( + (inputs.q_fp4, inputs.q_sf), + inputs.k_cache, + inputs.weights, + metadata.compressed_seq_lens, + metadata.page_table, + metadata.deep_gemm_metadata, + metadata.max_compressed_seq_len, + ) + logits, _ = two_level_decode_logits( + logits, + metadata.compressed_seq_lens, + is_candidate_source=False, + uses_candidates=True, + topk_blocks=self.topk_blocks, + block_size=self.block_size, + published=candidate_metadata.mask, + ) + # raw positions into `selected`; `page_indices` gets the unfiltered slots + # here and is rewritten with the masked selection just below + selected = torch.empty_like(page_indices) + fp32_jit_paged_topk(logits, metadata, page_indices, raw_indices=selected) + if logits.is_cuda and torch.version.cuda is not None: + # fused: drop the selections the mask zeroed, page-transform the rest + from sglang.kernels.ops.attention.dsv4.indexer_postprocess import ( + filter_topk_pages, + ) + + filter_topk_pages( + logits, + selected, + metadata.page_table, + page_indices, + page_size, + raw_indices, + ) + else: + selected = mask_topk_scores(logits, selected) + columns = selected.clamp_min(0).to(torch.int64) + slots = metadata.page_table.gather(1, columns // page_size) * page_size + slots = slots + columns % page_size + page_indices.copy_(torch.where(selected >= 0, slots, -1)) + if raw_indices is not None: + raw_indices.copy_(selected) diff --git a/python/sglang/srt/layers/attention/dsv4/indexer.py b/python/sglang/srt/layers/attention/dsv4/indexer.py index a778b2307d8e..6403dc4876ac 100644 --- a/python/sglang/srt/layers/attention/dsv4/indexer.py +++ b/python/sglang/srt/layers/attention/dsv4/indexer.py @@ -1127,3 +1127,88 @@ def forward( q_lora_ready=q_lora_ready, skip_compressor=skip_compressor, ) + + +def fp4_paged_mqa_logits( + q_fp4: Tuple[torch.Tensor, torch.Tensor], + k_cache: torch.Tensor, + weights: torch.Tensor, + seq_lens: torch.Tensor, + page_table: torch.Tensor, + deep_gemm_metadata, + max_seq_len: int, +) -> torch.Tensor: + """DeepGEMM paged fp4 logits for the low-ratio indexer. No hadamard: the + reference does not apply one.""" + from deep_gemm import fp8_fp4_paged_mqa_logits + + sl = seq_lens.to(torch.int32) + if sl.dim() == 1: + sl = sl.unsqueeze(-1) + return fp8_fp4_paged_mqa_logits( + q_fp4, + k_cache, + weights, + sl, + page_table, + deep_gemm_metadata, + max_seq_len, + False, + ) + + +def fp32_jit_paged_topk( + logits: torch.Tensor, + metadata, + page_indices: torch.Tensor, + raw_indices: Optional[torch.Tensor] = None, +) -> None: + """Plain top-k of the dense paged ``logits``: pool slots into ``page_indices`` + (``-1`` past the valid count) and, when given, positions into ``raw_indices``; + ``metadata`` is the ratio's ``PagedIndexerMetadata``.""" + if metadata.use_topk_v2: + topk_transform_paged_v2( + logits, + metadata.compressed_seq_lens, + metadata.page_table, + page_indices, + metadata.compressed_page_size, + metadata.topk_metadata, + raw_indices, + ) + else: + topk_transform_paged( + logits, + metadata.compressed_seq_lens, + metadata.page_table, + page_indices, + metadata.compressed_page_size, + raw_indices, + ) + + +def select_candidate_blocks( + logits: torch.Tensor, + compress_lens: torch.Tensor | int, + topk_blocks: int, + block_size: int, +) -> torch.Tensor: + """Level one of the two-level top-k: a bool mask over positions keeping the + topk_blocks best-scoring blocks per query. Unreachable positions are already -inf + in logits, so an all -inf block means not reachable yet; the block holding the + query's newest position is always kept.""" + width = logits.size(-1) + scores = F.pad(logits, (0, -width % block_size), value=-torch.inf) + scores = scores.unflatten(-1, (-1, block_size)).amax(dim=-1) + num_blocks = scores.size(-1) + + last = (compress_lens - 1) // block_size + scores = scores.masked_fill( + torch.arange(num_blocks, device=logits.device) == last, torch.inf + ) + + top = scores.topk(min(topk_blocks, num_blocks), dim=-1) + keep = torch.zeros_like(scores, dtype=torch.bool).scatter_( + -1, top.indices, top.values > -torch.inf + ) + return keep.repeat_interleave(block_size, dim=-1)[..., :width] diff --git a/python/sglang/srt/layers/deep_gemm_wrapper/configurer.py b/python/sglang/srt/layers/deep_gemm_wrapper/configurer.py index a21073f5cf83..4276b33359d6 100644 --- a/python/sglang/srt/layers/deep_gemm_wrapper/configurer.py +++ b/python/sglang/srt/layers/deep_gemm_wrapper/configurer.py @@ -43,3 +43,20 @@ def _compute_enable_deep_gemm(): get_platform().is_sm100 or get_device_sm() == 120 ) DEEPGEMM_NEED_TMA_ALIGNED_SCALES = not (DEEPGEMM_SCALE_UE8M0 or _is_musa) + + +def _supports_sparse_indexer() -> bool: + if not DEEPGEMM_BLACKWELL: + return False + import deep_gemm + + return all( + callable(getattr(deep_gemm, name, None)) + for name in ( + "get_paged_sparse_mqa_logits_metadata", + "fp8_fp4_paged_sparse_mqa_logits", + ) + ) + + +DEEPGEMM_SPARSE_INDEXER = _supports_sparse_indexer() diff --git a/test/registered/kernel/attention/dsv4/test_dsv41_sparse_indexer.py b/test/registered/kernel/attention/dsv4/test_dsv41_sparse_indexer.py new file mode 100644 index 000000000000..6ae103d0825e --- /dev/null +++ b/test/registered/kernel/attention/dsv4/test_dsv41_sparse_indexer.py @@ -0,0 +1,234 @@ +"""Decode two-level indexer on DeepGEMM's paged sparse MQA logits. + +The block table layer 20 publishes (``amax_topk_blocks``) is checked against +the model code's ``select_candidate_blocks``; a consumer's sparse logits +are checked against DeepGEMM's dense bf16 paged logits gathered at the published +positions, which the kernel is documented to match bitwise; its selection is +checked against a torch top-k of those. The kernel path needs a DeepGEMM with +``fp8_fp4_paged_sparse_mqa_logits`` on an SM100 device; it skips elsewhere. +""" + +import unittest + +import torch + +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci(est_time=120, stage="base-b-kernel-unit", runner_config="4-gpu-b200") + +HEADS = 32 +HEAD_DIM = 128 +PAGE = 128 # index pool page on SM100: 128 * 68 bytes = 17 * 512 +TOPK = 512 +BLOCKS = 2048 + + +def _sparse_indexer_available() -> bool: + if not torch.cuda.is_available() or torch.cuda.get_device_capability()[0] != 10: + return False + try: + import deep_gemm + except ImportError: + return False + return hasattr(deep_gemm, "fp8_fp4_paged_sparse_mqa_logits") + + +def _reference_blocks(logits, lens, block_size=8): + """Block ids the model code keeps, per row, ascending.""" + from sglang.srt.layers.attention.dsv4.indexer import select_candidate_blocks + + width = logits.shape[1] + reach = torch.arange(width, device=logits.device)[None, :] < lens[:, None] + mask = select_candidate_blocks( + logits.masked_fill(~reach, -torch.inf), + lens[:, None], + topk_blocks=BLOCKS, + block_size=block_size, + ) + return [m.view(-1, block_size).any(-1).nonzero().flatten() for m in mask] + + +class TestSparseIndexer(CustomTestCase): + def test_amax_topk_blocks_matches_reference(self): + # short rows: the block top-k skips its plan; 40 rows of 300K tokens: + # 37500 keys per row on a batch above the persistent pool, plan needed + self._check_amax_topk_blocks( + torch.tensor( + [1, 37, 16384, 16389, 40000, 131072], dtype=torch.int32, device="cuda" + ), + 131072, + ) + self._check_amax_topk_blocks( + torch.randint(200000, 300001, (40,), dtype=torch.int32, device="cuda"), + 300000, + ) + + def _check_amax_topk_blocks(self, lens, width): + from sglang.kernels.ops.attention.dsv4.candidate_blocks import ( + candidate_row_lens, + ) + from sglang.kernels.ops.attention.dsv4.topk import sort_candidate_blocks + from sglang.srt.layers.attention.dsv4.candidate_deep_gemm import ( + amax_topk_blocks, + valid_lens, + ) + + torch.manual_seed(0) + bs = lens.numel() + logits = torch.randn(bs, width, device="cuda") + # the tail past the length is garbage in production: make it loud + logits.masked_fill_( + torch.arange(width, device="cuda")[None, :] >= lens[:, None], 1e4 + ) + pages = (width + PAGE - 1) // PAGE + page_table = torch.stack( + [torch.randperm(pages, device="cuda") for _ in range(bs)] + ).to(torch.int32) + nblocks, valid = candidate_row_lens(lens, BLOCKS) + self.assertTrue(torch.equal(nblocks, (lens + 7) // 8)) + self.assertTrue(torch.equal(valid, valid_lens(lens, BLOCKS))) + blocks = amax_topk_blocks(logits, lens, nblocks, BLOCKS) + phys = sort_candidate_blocks(blocks, lens, page_table, PAGE) + keys = logits.view(bs, -1, 8).amax(-1) + bpp = PAGE // 8 + for b, ref in enumerate(_reference_blocks(logits, lens)): + nb = (int(lens[b]) + 7) // 8 + n = min(nb, BLOCKS) + got = blocks[b, :n].long() + self.assertTrue(torch.equal(got, got.sort().values), "not ascending") + self.assertEqual(got.unique().numel(), n) + self.assertIn(nb - 1, got.tolist(), "newest block not kept") + self.assertTrue(bool((got < nb).all())) + # equal keys may swap blocks: compare the key multiset (the forced + # block excluded, its key is arbitrary garbage) + keep = got != nb - 1 + keep_ref = ref != nb - 1 + self.assertTrue( + torch.equal( + keys[b][got[keep]].sort().values, + keys[b][ref[keep_ref]].sort().values, + ), + msg=f"row {b}", + ) + # past the valid count nothing looks like a block DeepGEMM could read + self.assertTrue(bool((blocks[b, n:] >= nb).all())) + # the same blocks as pool slots / 8 through the row's page table + ref_phys = page_table[b][got // bpp].long() * bpp + got % bpp + self.assertTrue(torch.equal(phys[b, :n].long(), ref_phys)) + self.assertTrue(bool((phys[b, n:] == torch.iinfo(torch.int32).max).all())) + expect_valid = 8 * (n - 1) + ((int(lens[b]) - 1) % 8 + 1) + self.assertEqual(int(valid[b]), expect_valid) + + @unittest.skipUnless( + _sparse_indexer_available(), "needs DeepGEMM's paged sparse MQA logits on SM100" + ) + def test_paired_verify_rows_match_unpaired(self): + """Verify shape: each request has 6 consecutive rows (its draft tokens) with + lengths L, L+1, ..., sharing one page-table row. With request ids DeepGEMM + pairs the rows on one KV pass; the sparse logits must equal the unpaired + (every row its own request) result bitwise, row layout unchanged.""" + import deep_gemm + + from sglang.kernels.ops.attention.dsv4.candidate_blocks import ( + candidate_row_lens, + ) + from sglang.kernels.ops.attention.dsv4.topk import ( + sort_candidate_blocks, + ) + from sglang.srt.layers.attention.dsv4.candidate_deep_gemm import ( + SparseBlockTable, + amax_topk_blocks, + build_sparse_indexer_schedule, + sparse_logits, + valid_lens, + ) + + torch.manual_seed(3) + draft = 6 + base = torch.tensor([20000, 16385, 70000], dtype=torch.int32, device="cuda") + lens = ( + base[:, None] + torch.arange(draft, device="cuda", dtype=torch.int32) + ).flatten() + request_ids = torch.repeat_interleave( + torch.tensor([7, 3, 11], dtype=torch.int64, device="cuda"), draft + ) + rows = lens.numel() + max_pages = (int(lens.max()) + PAGE - 1) // PAGE + num_pages = base.numel() * max_pages + pool = torch.randint( + 0, 255, (num_pages, PAGE * 68), dtype=torch.uint8, device="cuda" + ) + pool[:, PAGE * 64 :] = torch.randint( + 118, 123, (num_pages, PAGE * 4), dtype=torch.uint8, device="cuda" + ) + k_cache = pool.view(num_pages, PAGE, 1, 68) + per_request = ( + torch.randperm(num_pages, device="cuda") + .view(base.numel(), max_pages) + .to(torch.int32) + ) + page_table = per_request.repeat_interleave(draft, dim=0) # rows share theirs + q_fp4 = torch.randint( + 0, 255, (rows, 1, HEADS, HEAD_DIM // 2), dtype=torch.uint8, device="cuda" + ).view(torch.int8) + q_sf = ( + torch.randint( + 118, 123, (rows, 1, HEADS, 4), dtype=torch.uint8, device="cuda" + ) + .view(torch.int32) + .squeeze(-1) + ) + weights = (torch.rand(rows, HEADS, device="cuda") * 0.05).to(torch.bfloat16) + sched = deep_gemm.get_paged_mqa_logits_metadata( + lens.view(-1, 1), PAGE, deep_gemm.get_num_sms() + ) + dense = deep_gemm.fp8_fp4_paged_mqa_logits( + (q_fp4, q_sf), + k_cache, + weights.float(), + lens.view(-1, 1), + page_table, + sched, + int(lens.max()), + False, + torch.float32, + ) + nblocks, row_valid = candidate_row_lens(lens, BLOCKS) + blocks = amax_topk_blocks(dense, lens, nblocks, BLOCKS) + phys = sort_candidate_blocks(blocks, lens, page_table, PAGE) + out = {} + for name, ids in (("paired", request_ids), ("unpaired", None)): + schedule = build_sparse_indexer_schedule( + blocks, lens, page_table, PAGE, q_fp4.dtype, ids + ) + table = SparseBlockTable( + blocks=blocks, schedule=schedule, phys_blocks=phys, valid_lens=row_valid + ) + out[name] = sparse_logits(q_fp4, q_sf, k_cache, weights, table) + cols = torch.arange(BLOCKS * 8, device="cuda") + valid = cols[None, :] < row_valid[:, None].long() + self.assertTrue(torch.equal(row_valid, valid_lens(lens, BLOCKS))) + self.assertTrue( + torch.equal(out["paired"][valid], out["unpaired"][valid]), + "pairing changed the sparse logits", + ) + # and both equal the dense bf16 logits at the published positions + dense16 = deep_gemm.fp8_fp4_paged_mqa_logits( + (q_fp4, q_sf), + k_cache, + weights, + lens.view(-1, 1), + page_table, + sched, + int(lens.max()), + False, + torch.bfloat16, + ) + pos = blocks.long().repeat_interleave(8, dim=1) * 8 + (cols % 8)[None, :] + ref = dense16.gather(1, pos.clamp(max=dense16.shape[1] - 1)) + self.assertTrue(torch.equal(out["paired"][valid], ref[valid])) + + +if __name__ == "__main__": + unittest.main() From dee8e7d88c411d31537c1ea88e5d14e12991b1c2 Mon Sep 17 00:00:00 2001 From: hnyls2002 Date: Tue, 15 Sep 2026 15:53:42 -0700 Subject: [PATCH 17/30] dsv4.1: extract RoPE and FP4 packing kernels --- .../kernels/jit/csrc/deepseek_v4/fp4_rope.cuh | 458 ++++++++++++++++++ .../kernels/ops/attention/dsv4/fp4_indexer.py | 46 +- .../kernels/ops/attention/dsv4/fp4_rope.py | 166 +++++++ .../ops/attention/dsv4/rope_fake_quant_fp4.py | 130 +++++ .../ops/attention/dsv4/rope_pack_indexer.py | 138 ++++++ .../dsv4/test_compressed_kv_quant.py | 81 ++++ 6 files changed, 1015 insertions(+), 4 deletions(-) create mode 100644 python/sglang/kernels/jit/csrc/deepseek_v4/fp4_rope.cuh create mode 100644 python/sglang/kernels/ops/attention/dsv4/fp4_rope.py create mode 100644 python/sglang/kernels/ops/attention/dsv4/rope_fake_quant_fp4.py create mode 100644 python/sglang/kernels/ops/attention/dsv4/rope_pack_indexer.py create mode 100644 test/registered/kernel/attention/dsv4/test_compressed_kv_quant.py diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/fp4_rope.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/fp4_rope.cuh new file mode 100644 index 000000000000..704fb74cbef6 --- /dev/null +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/fp4_rope.cuh @@ -0,0 +1,458 @@ +#include +#include + +#include +#include +#include +#include +#include + +#include + +#include + +#include +#include +#include + +namespace sglang { + +/// \brief RMSNorm, RoPE, two FP4 stages and a 68-byte index-K cache store. +/// +/// `input` is `wk(latent)`, before `k_norm`. A ratio-r group's latent uses +/// its first position, `positions & ~(r - 1)`, for power-of-two ratios. +struct IndexKParams { + const bf16_t* __restrict__ input; // [num_tokens, kHeadDim] bf16, pre-norm + const bf16_t* __restrict__ norm_weight; // [kHeadDim] bf16 + const float* __restrict__ freqs_cis; // [max_pos, kRopeDim] fp32, real/imag interleaved + const void* __restrict__ positions; // [num_tokens] PosT + const int64_t* __restrict__ loc; // [num_tokens] index-K slot; 0 publishes nothing + uint8_t* __restrict__ cache; // [npages, kPageSize * 68] uint8 + uint32_t num_tokens; + float eps; +}; + +/// \brief RoPE and two-stage FP4 packing for index-Q, after `wq_b`. +/// +/// Each (token, head) row uses its token's own position, without RMSNorm +/// or a ratio mask. Payload and scales are contiguous outputs, not a paged cache. +struct IndexQParams { + const bf16_t* __restrict__ input; // [num_tokens, heads, kHeadDim] bf16 + const float* __restrict__ freqs_cis; // [max_pos, kRopeDim] fp32, real/imag interleaved + const void* __restrict__ positions; // [num_tokens] PosT + int8_t* __restrict__ payload; // [num_tokens * heads, kHeadDim / 2] int8 + int32_t* __restrict__ scale; // [num_tokens * heads] int32, four ue8m0 bytes + // Optional head-weight epilogue (kWeights): the raw `weights_proj` output for + // the same (token, head) rows, and where `float(bf16(w * weight_scale))` goes. + const bf16_t* __restrict__ head_weights; // [num_tokens * heads] bf16, or nullptr + float* __restrict__ weights_out; // [num_tokens * heads] fp32, or nullptr + float weight_scale; + uint32_t num_rows; + uint32_t heads; +}; + +/// Warps per CTA; one warp owns one row. Chosen from B200 decode measurements, +/// where occupancy has little effect; multiple warps avoid starving larger batches. +constexpr uint32_t kFp4RopeWarpsPerCTA = 4; + +/// \brief Indexer packer scale: `_ceil_ue8m0_exp(max(amax / 6, 1e-4))`. +/// +/// Unlike fake quantization, the floor follows the divide, division is not +/// replaced by multiplication by 1/6, and the exponent is clamped. +/// The two stages therefore require separate scales. +SGL_DEVICE uint32_t index_pack_exponent(float amax) { + const auto bits = __float_as_uint(fmaxf(amax / 6.0f, 1.0e-4f)); + const auto exponent = static_cast((bits >> 23) & 0xFF) + ((bits & 0x7FFFFF) != 0); + // Neither bound is reachable for finite fp32 inputs: the 1e-4 floor keeps + // the exponent above 1, and reaching 254 requires absmax > 6 * 2^126. + return static_cast(min(max(exponent, 1), 254)); +} + +/// \brief Clear the sign of every packed nibble whose magnitude rounded to zero. +/// +/// `cvt.rn.satfinite.e2m1x2.f32` keeps the sign of a small negative, giving the +/// `-0` code `0x8`; the reference packer drops it (`sign = (x < 0) & (idx != 0)`). +/// Both dequantize to zero, but the stored byte differs, so match the reference. +/// Branchless and independent of how many nibbles the word holds: bit 4k+3 of +/// each nibble survives only if one of bits 4k..4k+2 is set. +SGL_DEVICE uint32_t clear_negative_zero(uint32_t packed) { + const auto any_magnitude = (packed | (packed >> 1) | (packed >> 2)) & 0x11111111u; + return packed & ((any_magnitude << 3) | 0x77777777u); +} + +/// One lane's share of a packed 128-element row: two payload bytes and the two +/// block exponents its half of the warp owns. +struct IndexPacked { + uint32_t payload[2]; // the head pair's byte, then the tail pair's + uint32_t exponent[2]; // blocks {0, 1} then {2, 3}, by half of the warp +}; + +/// \brief The whole shared body of both directions: RoPE tail, both fp4 stages, +/// and the indexer's pack. +/// +/// `head` and `tail` are the lane's two bf16 pairs, widened and already rounded +/// to bf16 by whatever produced them (the caller's RMSNorm, or the load itself). +/// `tail` is pre-rotation and `freq` is its matching `(real, imag)`. +/// +/// A 32-element fp4 block is 16 lanes of *one* half -- blocks 0/1 are the head +/// on lanes 0-15 / 16-31 and blocks 2/3 the tail -- so each quantization stage +/// costs two `reduce_max<16>`, not four, and nothing here spans the row. +SGL_DEVICE IndexPacked index_rope_quant_pack(fp32x2_t head, fp32x2_t tail, fp32x2_t freq) { + using namespace device; + namespace fp4 = deepseek_v4::fp4; + + constexpr uint32_t kHalfLanes = kWarpThreads / 2; + static_assert(fp4::kBlockSize == kHalfLanes * 2, "an fp4 block must be half a warp of one half"); + + float data[4]; + data[0] = head.x; + data[1] = head.y; + // `rope_tail` ends in `.to(x.dtype)`, so the rotated pair is rounded again. + const auto rotated = + cast(cast(fp32x2_t{tail.x * freq.x - tail.y * freq.y, tail.x * freq.y + tail.y * freq.x})); + data[2] = rotated.x; + data[3] = rotated.y; + + // The fake-quant result is already exact in bf16: e2m1 needs at most three + // significant bits, and the power-of-two scales admitted by the amax floor fit bf16. +#pragma unroll + for (uint32_t half = 0; half < 2; ++half) { + const auto amax = warp::reduce_max(fmaxf(fabsf(data[half * 2]), fabsf(data[half * 2 + 1]))); + const auto [scale, inv_scale] = fp4::block_scale(amax); + const auto q = fp4::fake_quant_x2({data[half * 2], data[half * 2 + 1]}, scale, inv_scale); + data[half * 2 + 0] = q.x; + data[half * 2 + 1] = q.y; + } + + // The packer's scale floor differs from fake quantization; keep both stages. + // Each packed byte puts `.x` in the low nibble. + IndexPacked out; +#pragma unroll + for (uint32_t half = 0; half < 2; ++half) { + const auto amax = warp::reduce_max(fmaxf(fabsf(data[half * 2]), fabsf(data[half * 2 + 1]))); + out.exponent[half] = index_pack_exponent(amax); + // `inv_scale_ue8m0` instead of the reference's division: the scale is a + // power of two so both are exact, except at exponent 254, which needs a + // block absmax above `6 * 2^126` and so cannot come from a finite float. + const auto inv_scale = deepseek_v4::fp8::inv_scale_ue8m0(static_cast(out.exponent[half])); + const auto code = __nv_cvt_float2_to_fp4x2( + fp32x2_t{data[half * 2] * inv_scale, data[half * 2 + 1] * inv_scale}, __NV_E2M1, cudaRoundNearest); + out.payload[half] = clear_negative_zero(static_cast(code)); + } + return out; +} + +/// \brief The row's four block exponents, packed little-endian into one word. +/// +/// They live on two lanes -- 0 holds blocks 0 and 2, 16 holds 1 and 3 -- so this +/// costs two shuffles, and the result is the word only for the lower half of +/// the warp. Every lane must reach it: the shuffles are warp-wide. +SGL_DEVICE uint32_t index_scale_word(const uint32_t (&exponent)[2]) { + using namespace device; + const auto exp_1 = __shfl_sync(warp::kFullMask, exponent[0], kWarpThreads / 2); + const auto exp_3 = __shfl_sync(warp::kFullMask, exponent[1], kWarpThreads / 2); + return exponent[0] | (exp_1 << 8) | (exponent[1] << 16) | (exp_3 << 24); +} + +/// \brief One warp per token; grid = ceil(num_tokens / kFp4RopeWarpsPerCTA). +/// +/// Lane L owns head elements {2L, 2L+1} and tail elements {64+2L, 64+2L+1}, +/// so each lane carries one complex RoPE pair. Only RMSNorm spans the full +/// row; the remaining reductions use the FP4 block layout. +template +__global__ +__launch_bounds__(kFp4RopeWarpsPerCTA* device::kWarpThreads) void flash_index_k_kernel(const IndexKParams params) { + using namespace device; + namespace fp4 = deepseek_v4::fp4; + + constexpr uint32_t kPayloadBytes = kHeadDim / 2; + constexpr uint32_t kScaleBytes = kHeadDim / fp4::kBlockSize; + constexpr uint32_t kSlotBytes = kPayloadBytes + kScaleBytes; + + static_assert( + kHeadDim == 128 && kRopeDim == 64, + "the one-warp tiling is specific to a 128-wide row whose second half is the RoPE tail"); + static_assert(kHeadDim == 2 * kWarpThreads * 2, "a lane owns one bf16x2 of each half"); + static_assert(kScaleBytes == 4, "the four block exponents are packed into one uint32 store"); + static_assert(std::has_single_bit(kRatio), "group_pos is derived by masking, so the ratio must be a power of two"); + + using bf16_vec_t = AlignedVector; + using fp32_vec_t = AlignedVector; + + const auto lane = threadIdx.x % kWarpThreads; + const auto row = blockIdx.x * kFp4RopeWarpsPerCTA + threadIdx.x / kWarpThreads; + // Warp-uniform, so the reductions below still see a full warp. + if (row >= params.num_tokens) return; + + // Both come from the step's metadata rather than the predecessor, so + // prefetching them ahead of the PDL gate overlaps with the `wk` GEMM's tail. + // A row that publishes nothing is dropped at the store, not here. + const auto slot_id = params.loc[row]; + const auto position = static_cast(static_cast(params.positions)[row]); + PDLWaitPrimary(); + + bf16_vec_t head_in, tail_in, head_w, tail_w; + head_in.load(params.input + row * kHeadDim, lane); + tail_in.load(params.input + row * kHeadDim, lane + kWarpThreads); + head_w.load(params.norm_weight, lane); + tail_w.load(params.norm_weight, lane + kWarpThreads); + fp32_vec_t freq; + freq.load(params.freqs_cis + (position & ~static_cast(kRatio - 1)) * kRopeDim, lane); + + fp32x2_t head, tail; + { + const auto [h0, h1] = cast(head_in[0]); + const auto [t0, t1] = cast(tail_in[0]); + const auto sqrsum = warp::reduce_sum(h0 * h0 + h1 * h1 + t0 * t0 + t1 * t1); + const auto inv_rms = math::rsqrt(sqrsum * (1.0f / static_cast(kHeadDim)) + params.eps); + const auto [wh0, wh1] = cast(head_w[0]); + const auto [wt0, wt1] = cast(tail_w[0]); + // `k_norm` materializes a bf16 tensor, so the norm result is rounded before + // anything downstream sees it -- the RoPE below included. + head = cast(cast(fp32x2_t{wh0 * (h0 * inv_rms), wh1 * (h1 * inv_rms)})); + tail = cast(cast(fp32x2_t{wt0 * (t0 * inv_rms), wt1 * (t1 * inv_rms)})); + } + + const auto packed = index_rope_quant_pack(head, tail, fp32x2_t{freq[0], freq[1]}); + const auto scale_word = index_scale_word(packed.exponent); + + // A padded graph row, and at ratio > 1 a row completing no group, carry the + // reserved slot 0 and must publish nothing. + if (slot_id <= 0) return; + const auto page = slot_id / kPageSize; + const auto slot = slot_id % kPageSize; + const auto page_ptr = params.cache + page * (kPageSize * kSlotBytes); + + // Byte i of the payload covers elements (2i, 2i+1), so a lane's head pair is + // byte `lane` and its tail pair byte `lane + 32`: two coalesced 32-byte runs. + const auto payload_ptr = page_ptr + slot * kPayloadBytes; + payload_ptr[lane] = static_cast(packed.payload[0]); + payload_ptr[lane + kWarpThreads] = static_cast(packed.payload[1]); + + if (lane == 0) { + *reinterpret_cast(page_ptr + kPageSize * kPayloadBytes + slot * kScaleBytes) = scale_word; + } +} + +/// \brief One warp per (token, head); grid = ceil(num_rows / kFp4RopeWarpsPerCTA). +/// +/// Input is contiguous [num_tokens, heads, kHeadDim]; row r uses token r / heads +/// and head r % heads. Queries use their own position, without a ratio mask. +template +__global__ +__launch_bounds__(kFp4RopeWarpsPerCTA* device::kWarpThreads) void flash_index_q_kernel(const IndexQParams params) { + using namespace device; + + constexpr uint32_t kPayloadBytes = kHeadDim / 2; + + static_assert( + kHeadDim == 128 && kRopeDim == 64, + "the one-warp tiling is specific to a 128-wide row whose second half is the RoPE tail"); + static_assert(kHeadDim == 2 * kWarpThreads * 2, "a lane owns one bf16x2 of each half"); + + using bf16_vec_t = AlignedVector; + using fp32_vec_t = AlignedVector; + + const auto lane = threadIdx.x % kWarpThreads; + const auto row = blockIdx.x * kFp4RopeWarpsPerCTA + threadIdx.x / kWarpThreads; + // Warp-uniform, so the reductions below still see a full warp. + if (row >= params.num_rows) return; + + // The position lookup is independent of the PDL producer. + const auto position = static_cast(static_cast(params.positions)[row / params.heads]); + PDLWaitPrimary(); + + bf16_vec_t head_in, tail_in; + head_in.load(params.input + row * kHeadDim, lane); + tail_in.load(params.input + row * kHeadDim, lane + kWarpThreads); + fp32_vec_t freq; + freq.load(params.freqs_cis + position * kRopeDim, lane); + + const auto packed = + index_rope_quant_pack(cast(head_in[0]), cast(tail_in[0]), fp32x2_t{freq[0], freq[1]}); + const auto scale_word = index_scale_word(packed.exponent); + + const auto payload_ptr = params.payload + row * kPayloadBytes; + payload_ptr[lane] = static_cast(packed.payload[0]); + payload_ptr[lane + kWarpThreads] = static_cast(packed.payload[1]); + + if (lane == 0) params.scale[row] = static_cast(scale_word); + + if constexpr (kWeights) { + // Match head_weights(x).float(): multiply by the fp32-rounded scale, + // round to nearest-even bf16, then widen to fp32. + if (lane == 0) { + const auto w = cast(params.head_weights[row]) * params.weight_scale; + params.weights_out[row] = cast(cast(w)); + } + } +} + +/// \brief Host side of `flash_index_k_kernel`. +template +struct FlashIndexKKernel { + static constexpr uint32_t kBlockSize = kFp4RopeWarpsPerCTA * device::kWarpThreads; + static constexpr int64_t kSlotBytes = kHeadDim / 2 + kHeadDim / deepseek_v4::fp4::kBlockSize; + + template + static constexpr auto kernel = flash_index_k_kernel; + + /// \param input `[num_tokens, kHeadDim]` bf16, `wk(latent)` before `k_norm`. + /// \param norm_weight `[kHeadDim]` bf16, `k_norm.weight`. + /// \param freqs_cis `[max_pos, kRopeDim]` fp32, real/imag interleaved. + /// \param positions `[num_tokens]` int32 or int64, the *token* position; the + /// group position is derived from it and the ratio. + /// \param loc `[num_tokens]` int64, the index-K slot; `0` publishes nothing. + /// \param cache `[npages, kPageSize * 68]` uint8. + static void run_index_k( + const tvm::ffi::TensorView input, + const tvm::ffi::TensorView norm_weight, + const tvm::ffi::TensorView freqs_cis, + const tvm::ffi::TensorView positions, + const tvm::ffi::TensorView loc, + const tvm::ffi::TensorView cache, + const float eps) { + using namespace host; + + auto N = SymbolicSize{"num_tokens"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + + TensorMatcher({N, kHeadDim}).with_dtype().with_device(device_).verify(input); + TensorMatcher({kHeadDim}).with_dtype().with_device(device_).verify(norm_weight); + // Real/imag interleaved, so the trailing dim is kRopeDim, not kRopeDim / 2. + TensorMatcher({-1, kRopeDim}).with_dtype().with_device(device_).verify(freqs_cis); + auto pos_dtype = SymbolicDType{}; + TensorMatcher({N}).with_dtype(pos_dtype).with_device(device_).verify(positions); + TensorMatcher({N}).with_dtype().with_device(device_).verify(loc); + TensorMatcher({-1, kPageSize * kSlotBytes}).with_dtype().with_device(device_).verify(cache); + + const auto num_tokens = static_cast(N.unwrap()); + if (num_tokens == 0) return; + + const auto params = IndexKParams{ + .input = static_cast(input.data_ptr()), + .norm_weight = static_cast(norm_weight.data_ptr()), + .freqs_cis = static_cast(freqs_cis.data_ptr()), + .positions = positions.data_ptr(), + .loc = static_cast(loc.data_ptr()), + .cache = static_cast(cache.data_ptr()), + .num_tokens = num_tokens, + .eps = eps, + }; + const auto k_int32 = kernel; + const auto k_int64 = kernel; + const auto k = pos_dtype.is_type() ? k_int32 : k_int64; + LaunchKernel(div_ceil(num_tokens, kFp4RopeWarpsPerCTA), kBlockSize, device_.unwrap()) // + .enable_pdl(kUsePDL)(k, params); + } +}; + +/// \brief Host side of `flash_index_q_kernel`. +template +struct FlashIndexQKernel { + static constexpr uint32_t kBlockSize = kFp4RopeWarpsPerCTA * device::kWarpThreads; + + template + static constexpr auto kernel = flash_index_q_kernel; + + /// \param input `[num_tokens, heads, kHeadDim]` bf16, `wq_b(q_lora)`. + /// \param freqs_cis `[max_pos, kRopeDim]` fp32, real/imag interleaved. + /// \param positions `[num_tokens]` int32 or int64, the query's own position. + /// \param payload `[num_tokens * heads, kHeadDim / 2]` int8. + /// \param scale `[num_tokens * heads]` int32, the four ue8m0 block exponents + /// packed little-endian. + static void run_index_q( + const tvm::ffi::TensorView input, + const tvm::ffi::TensorView freqs_cis, + const tvm::ffi::TensorView positions, + const tvm::ffi::TensorView payload, + const tvm::ffi::TensorView scale) { + launch(input, freqs_cis, positions, payload, scale, std::nullopt, std::nullopt, 0.0f); + } + + /// \brief `run_index_q` plus the indexer's head-weight epilogue. + /// + /// \param head_weights `[num_tokens, heads]` bf16, the raw `weights_proj(x)`. + /// \param weights_out `[num_tokens, heads]` fp32, receives + /// `float(bf16(head_weights * weight_scale))`, i.e. `head_weights(x).float()`. + /// \param weight_scale `softmax_scale * heads^-0.5`, applied in fp32. + static void run_index_q_weights( + const tvm::ffi::TensorView input, + const tvm::ffi::TensorView freqs_cis, + const tvm::ffi::TensorView positions, + const tvm::ffi::TensorView payload, + const tvm::ffi::TensorView scale, + const tvm::ffi::TensorView head_weights, + const tvm::ffi::TensorView weights_out, + const double weight_scale) { + launch(input, freqs_cis, positions, payload, scale, head_weights, weights_out, static_cast(weight_scale)); + } + + private: + using MaybeTensor = std::optional; + + static void launch( + const tvm::ffi::TensorView input, + const tvm::ffi::TensorView freqs_cis, + const tvm::ffi::TensorView positions, + const tvm::ffi::TensorView payload, + const tvm::ffi::TensorView scale, + const MaybeTensor head_weights, + const MaybeTensor weights_out, + const float weight_scale) { + using namespace host; + + auto N = SymbolicSize{"num_tokens"}; + auto H = SymbolicSize{"heads"}; + auto R = SymbolicSize{"num_rows"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + + TensorMatcher({N, H, kHeadDim}).with_dtype().with_device(device_).verify(input); + // Real/imag interleaved, so the trailing dim is kRopeDim, not kRopeDim / 2. + TensorMatcher({-1, kRopeDim}).with_dtype().with_device(device_).verify(freqs_cis); + auto pos_dtype = SymbolicDType{}; + TensorMatcher({N}).with_dtype(pos_dtype).with_device(device_).verify(positions); + TensorMatcher({R, kHeadDim / 2}).with_dtype().with_device(device_).verify(payload); + TensorMatcher({R}).with_dtype().with_device(device_).verify(scale); + const auto weights = head_weights.has_value(); + RuntimeCheck(weights == weights_out.has_value(), "head_weights and weights_out come together"); + if (weights) { + TensorMatcher({N, H}).with_dtype().with_device(device_).verify(*head_weights); + TensorMatcher({N, H}).with_dtype().with_device(device_).verify(*weights_out); + } + RuntimeCheck( + R.unwrap() == N.unwrap() * H.unwrap(), + "payload holds ", + R.unwrap(), + " rows, but the input is ", + N.unwrap(), + " tokens x ", + H.unwrap(), + " heads"); + + const auto num_rows = static_cast(R.unwrap()); + if (num_rows == 0) return; + + const auto params = IndexQParams{ + .input = static_cast(input.data_ptr()), + .freqs_cis = static_cast(freqs_cis.data_ptr()), + .positions = positions.data_ptr(), + .payload = static_cast(payload.data_ptr()), + .scale = static_cast(scale.data_ptr()), + .head_weights = weights ? static_cast(head_weights->data_ptr()) : nullptr, + .weights_out = weights ? static_cast(weights_out->data_ptr()) : nullptr, + .weight_scale = weight_scale, + .num_rows = num_rows, + .heads = static_cast(H.unwrap()), + }; + const auto i32 = pos_dtype.is_type(); + const auto k = weights ? (i32 ? kernel : kernel) + : (i32 ? kernel : kernel); + LaunchKernel(div_ceil(num_rows, kFp4RopeWarpsPerCTA), kBlockSize, device_.unwrap()) // + .enable_pdl(kUsePDL)(k, params); + } +}; + +} // namespace sglang diff --git a/python/sglang/kernels/ops/attention/dsv4/fp4_indexer.py b/python/sglang/kernels/ops/attention/dsv4/fp4_indexer.py index 64cb6bd1d679..847cf4bac1e1 100644 --- a/python/sglang/kernels/ops/attention/dsv4/fp4_indexer.py +++ b/python/sglang/kernels/ops/attention/dsv4/fp4_indexer.py @@ -37,6 +37,33 @@ def _fp4_e2m1_code(x): return idx | (sign << 3) +@triton.jit +def _fp4_e2m1_code_rne(x): + """Round-to-nearest-even e2m1 code, matching the reference rounding: at an + exact half-way value the even grid index wins.""" + ax = tl.minimum(tl.abs(x), 6.0) + idx = (ax >= 0.25).to(tl.uint8) + idx += (ax >= 0.75).to(tl.uint8) + idx += (ax >= 1.25).to(tl.uint8) + idx += (ax >= 1.75).to(tl.uint8) + idx += (ax >= 2.5).to(tl.uint8) + idx += (ax >= 3.5).to(tl.uint8) + idx += (ax >= 5.0).to(tl.uint8) + # Round-half-to-even: an odd index at an exact boundary drops to the even one. + is_boundary = ( + (ax == 0.25) + | (ax == 0.75) + | (ax == 1.25) + | (ax == 1.75) + | (ax == 2.5) + | (ax == 3.5) + | (ax == 5.0) + ) + idx = tl.where(is_boundary & ((idx & 1) == 1), idx - 1, idx) + sign = ((x < 0) & (idx != 0)).to(tl.uint8) + return idx | (sign << 3) + + @triton.jit def _quantize_fp4_indexer_kernel( x, @@ -44,6 +71,7 @@ def _quantize_fp4_indexer_kernel( x_sf, BLOCK_N: tl.constexpr, GROUP_N: tl.constexpr, + RNE: tl.constexpr, ): token_id = tl.program_id(0) offs = tl.arange(0, BLOCK_N) @@ -86,8 +114,12 @@ def _quantize_fp4_indexer_kernel( v0 = tl.load(x + token_id * BLOCK_N + offs0).to(tl.float32) / scale0 v1 = tl.load(x + token_id * BLOCK_N + offs1).to(tl.float32) / scale1 - code0 = _fp4_e2m1_code(v0) - code1 = _fp4_e2m1_code(v1) + if RNE: + code0 = _fp4_e2m1_code_rne(v0) + code1 = _fp4_e2m1_code_rne(v1) + else: + code0 = _fp4_e2m1_code(v0) + code1 = _fp4_e2m1_code(v1) packed = (code0 & 0x0F) | ((code1 & 0x0F) << 4) tl.store(x_fp4 + token_id * (BLOCK_N // 2) + pair_offsets, packed) @@ -120,7 +152,11 @@ def _store_fp4_index_k_cache_kernel( ) -def quantize_fp4_indexer_tensor(x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: +def quantize_fp4_indexer_tensor( + x: torch.Tensor, rne: bool = False +) -> tuple[torch.Tensor, torch.Tensor]: + """Per-32 ue8m0 fp4 quantize. rne=True uses round-to-nearest-even (the dsv41 + reference rounding); the default keeps the c4 threshold behavior.""" assert x.shape[-1] == 128 x = x.contiguous().view(-1, x.shape[-1]) x_fp4 = torch.empty((x.shape[0], 64), device=x.device, dtype=torch.int8) @@ -132,6 +168,7 @@ def quantize_fp4_indexer_tensor(x: torch.Tensor) -> tuple[torch.Tensor, torch.Te x_sf, BLOCK_N=128, GROUP_N=32, + RNE=rne, ) return x_fp4, x_sf @@ -142,9 +179,10 @@ def store_fp4_index_k_cache( loc: torch.Tensor, *, page_size: int, + rne: bool = False, ) -> None: assert input.shape[-1] == 128 - k_fp4, k_sf = quantize_fp4_indexer_tensor(input.contiguous()) + k_fp4, k_sf = quantize_fp4_indexer_tensor(input.contiguous(), rne=rne) n_tokens = input.numel() // input.shape[-1] assert k_fp4.shape == (n_tokens, 64) assert k_sf.shape == (n_tokens,) diff --git a/python/sglang/kernels/ops/attention/dsv4/fp4_rope.py b/python/sglang/kernels/ops/attention/dsv4/fp4_rope.py new file mode 100644 index 000000000000..9ed52c016a0d --- /dev/null +++ b/python/sglang/kernels/ops/attention/dsv4/fp4_rope.py @@ -0,0 +1,166 @@ +"""Fused RoPE and two-stage fp4 packing for low-ratio index keys and queries. + +Keys include RMSNorm and a 68-byte cache store, using the group's first position. +Queries use each token's own position, without RMSNorm or a cache store. +The decode backend uses index_q_rope_pack_weights to also produce FP32 head weights. +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +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: + from tvm_ffi.module import Module + +# Payload bytes plus one ue8m0 exponent per 32 elements, per compressed token. +SLOT_BYTES = 68 +INDEX_PAGE_SIZE = 64 + + +@cache_once +def _jit_index_k_module( + head_dim: int, rope_dim: int, page_size: int, ratio: int +) -> Module: + args = make_cpp_args(head_dim, rope_dim, page_size, ratio, is_arch_support_pdl()) + return load_jit( + make_name("fp4_rope"), + *args, + cuda_files=["deepseek_v4/fp4_rope.cuh"], + cuda_wrappers=[("index_k", f"FlashIndexKKernel<{args}>::run_index_k")], + ) + + +@cache_once +def _jit_index_q_module(head_dim: int, rope_dim: int) -> Module: + args = make_cpp_args(head_dim, rope_dim, is_arch_support_pdl()) + return load_jit( + make_name("fp4_rope"), + *args, + cuda_files=["deepseek_v4/fp4_rope.cuh"], + cuda_wrappers=[ + ("index_q", f"FlashIndexQKernel<{args}>::run_index_q"), + ("index_q_weights", f"FlashIndexQKernel<{args}>::run_index_q_weights"), + ], + ) + + +def index_k_norm_rope_pack_store( + input: torch.Tensor, + norm_weight: torch.Tensor, + eps: float, + freqs_cis: torch.Tensor, + positions: torch.Tensor, + loc: torch.Tensor, + cache: torch.Tensor, + *, + ratio: int, +) -> None: + """Normalize, rotate, quantize twice and store one index-K slot per token. + + :param input: ``[num_tokens, index_head_dim]`` bf16 -- ``wk(latent)``, + *before* ``k_norm``. + :param norm_weight: ``[index_head_dim]`` bf16, ``k_norm.weight``. + :param eps: ``k_norm.eps``. + :param freqs_cis: ``[max_pos, rope_head_dim]`` fp32, real/imag interleaved -- + ``torch.view_as_real(freqs).flatten(-2)``. Indexed + in-kernel, so pass the whole table rather than a gather. + :param positions: ``[num_tokens]`` int32 or int64, the token position. The + group position is masked out of it in-kernel. + :param loc: ``[num_tokens]`` int64, the index-K slot. ``0`` is the reserved + dummy: those rows publish nothing, which covers both padded + graph rows and, at ratio > 1, rows completing no group. + :param cache: the layer's index-K buffer, ``[npages, page_size * 68]`` uint8. + :param ratio: the layer's compress ratio. A power of two. + + .. note:: Two quantization stages, not one. The fake-quant's amax floor is + ``6 * 2**-126`` and the packer's is ``1e-4``, applied on opposite sides + of the divide by 6, so the packer can recover an exponent the fake-quant + gave away and collapsing them is not equivalent. + """ + head_dim = input.shape[-1] + _jit_index_k_module( + head_dim, freqs_cis.shape[-1], cache.shape[1] // SLOT_BYTES, ratio + ).index_k(input, norm_weight, freqs_cis, positions, loc, cache, float(eps)) + + +def index_q_rope_pack( + input: torch.Tensor, + freqs_cis: torch.Tensor, + positions: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor]: + """Rotate, quantize twice and pack one indexer query per (token, head). + + :param input: ``[num_tokens, heads, index_head_dim]`` bf16 contiguous -- + ``wq_b(q_lora)`` viewed per head. + :param freqs_cis: ``[max_pos, rope_head_dim]`` fp32, real/imag interleaved -- + ``torch.view_as_real(freqs).flatten(-2)``. Indexed + in-kernel, so pass the whole table rather than a gather. + :param positions: ``[num_tokens]`` int32 or int64. A query rotates by its + own position, so this is used unmasked. + :return: ``(payload, scale)`` -- ``[num_tokens * heads, index_head_dim // 2]`` + int8 and ``[num_tokens * heads]`` int32, the four ue8m0 block + exponents packed little-endian. Exactly what the paged MQA logits + kernel takes and what the Triton path returns without a cache. + + .. note:: Two quantization stages, not one -- see + :func:`index_k_norm_rope_pack_store`. Neither can be dropped. + """ + num_tokens, heads, head_dim = input.shape + rows = num_tokens * heads + payload = input.new_empty((rows, head_dim // 2), dtype=torch.int8) + scale = input.new_empty((rows,), dtype=torch.int32) + + _jit_index_q_module(head_dim, freqs_cis.shape[-1]).index_q( + input, freqs_cis, positions, payload, scale + ) + return payload, scale + + +def index_q_rope_pack_weights( + input: torch.Tensor, + freqs_cis: torch.Tensor, + positions: torch.Tensor, + head_weights: torch.Tensor, + weight_scale: float, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """:func:`index_q_rope_pack` plus the indexer's head weights, one launch. + + Head weights match ``head_weights(x).float()``: multiply in fp32, + round to nearest-even bf16, then widen to fp32. + + :param head_weights: ``[num_tokens, heads]`` bf16, the raw ``weights_proj`` + output (before the scale). + :param weight_scale: ``softmax_scale * heads**-0.5``; rounded to fp32 in the + kernel exactly as torch rounds a Python scalar for a + bf16 tensor multiply. + :return: ``(payload, scale, weights)`` -- the first two as + :func:`index_q_rope_pack`, ``weights`` ``[num_tokens, heads]`` fp32. + """ + num_tokens, heads, head_dim = input.shape + rows = num_tokens * heads + payload = input.new_empty((rows, head_dim // 2), dtype=torch.int8) + scale = input.new_empty((rows,), dtype=torch.int32) + weights = input.new_empty((num_tokens, heads), dtype=torch.float32) + + _jit_index_q_module(head_dim, freqs_cis.shape[-1]).index_q_weights( + input, + freqs_cis, + positions, + payload, + scale, + head_weights, + weights, + float(weight_scale), + ) + return payload, scale, weights diff --git a/python/sglang/kernels/ops/attention/dsv4/rope_fake_quant_fp4.py b/python/sglang/kernels/ops/attention/dsv4/rope_fake_quant_fp4.py new file mode 100644 index 000000000000..3f67f1f37f1e --- /dev/null +++ b/python/sglang/kernels/ops/attention/dsv4/rope_fake_quant_fp4.py @@ -0,0 +1,130 @@ +"""Fused RoPE tail and fp4 fake-quant for the DeepSeek-V4.1 low-ratio path. + +Preserves the bf16 round-trip after RoPE and before quantization, and the +round-half-to-even behavior of torch.round. +""" + +from __future__ import annotations + +import torch +import triton +import triton.language as tl +from triton.language.extra import libdevice + +FP4_MAX = 6.0 +FP4_AMAX_FLOOR = 6 * (2.0**-126) + + +@triton.jit +def _rope_tail_fake_quant_fp4_kernel( + x_ptr, + f_ptr, + out_ptr, + x_stride_r, + out_stride_r, + f_stride_t, + rows_per_token, + D: tl.constexpr, + RD: tl.constexpr, + BLK: tl.constexpr, + AMAX_FLOOR: tl.constexpr, + INVERSE: tl.constexpr, + COMPRESSED_KV: tl.constexpr, +): + r = tl.program_id(0) + t = r // rows_per_token + offs = tl.arange(0, D) + v = tl.load(x_ptr + r * x_stride_r + offs).to(tl.float32) + + # ---- rope_tail: adjacent pairs of the last RD features as one complex number + head_len = D - RD + in_tail = offs >= head_len + pos = offs - head_len + j = pos // 2 + is_im = (pos % 2) == 1 + re = tl.load(x_ptr + x_stride_r * r + head_len + 2 * j, mask=in_tail, other=0.0).to( + tl.float32 + ) + im = tl.load( + x_ptr + x_stride_r * r + head_len + 2 * j + 1, mask=in_tail, other=0.0 + ).to(tl.float32) + # freqs is a real/imag-interleaved view with stride 2 between complex pairs; + # index 2*j / 2*j+1 to preserve that layout. + fr = tl.load(f_ptr + t * f_stride_t + 2 * j, mask=in_tail, other=1.0) + fi = tl.load(f_ptr + t * f_stride_t + 2 * j + 1, mask=in_tail, other=0.0) + if INVERSE: + fi = -fi + rot = tl.where(is_im, re * fi + im * fr, re * fr - im * fi) + # rope_tail casts the rotated tail back to x.dtype before the cat; the head + # never leaves it. Reproduce that rounding or the quant sees different input. + rot = rot.to(tl.bfloat16).to(tl.float32) + v = tl.where(in_tail, rot, v) + + # ---- FP4 round-trip, with a separate scale format for compressed KV. + vb = tl.reshape(v, (D // BLK, BLK)) + amax = tl.max(tl.abs(vb), axis=1) + if COMPRESSED_KV: + scale = tl.minimum(tl.maximum(amax * (1.0 / 6.0), 2.0**-9), 448.0) + scale = scale.to(tl.float8e4nv).to(tl.float32) + s = tl.div_rn(vb, scale[:, None]) + else: + amax = tl.maximum(amax, AMAX_FLOOR) * (1.0 / 6.0) + # ceil_pow2 on the IEEE bits, exact at powers of two + bits = amax.to(tl.int32, bitcast=True) + expo = ((bits >> 23) & 0xFF) - 127 + expo = expo + ((bits & 0x7FFFFF) != 0).to(tl.int32) + scale = ((expo + 127) << 23).to(tl.float32, bitcast=True) + s = vb / scale[:, None] + s = tl.minimum(tl.maximum(s, -6.0), 6.0) + mag = tl.abs(s) + step = tl.where(mag < 2.0, 0.5, tl.where(mag < 4.0, 1.0, 2.0)) + # torch.round is round-half-to-even; torch.sign(0) is 0 + sgn = tl.where(s > 0, 1.0, tl.where(s < 0, -1.0, 0.0)) + q = libdevice.rint(mag / step) * step * sgn + out = tl.reshape(q * scale[:, None], (D,)) + tl.store(out_ptr + r * out_stride_r + offs, out.to(out_ptr.dtype.element_ty)) + + +def rope_tail_fake_quant_fp4( + x: torch.Tensor, + freqs: torch.Tensor, + rope_dim: int, + inverse: bool = False, + block_size: int = 32, + *, + compressed_kv: bool = False, +) -> torch.Tensor: + """RoPE and FP4 round-trip: per-16 E4M3 for compressed KV, per-32 UE8M0 otherwise. + + x: [T, ..., D] contiguous in the last dim; freqs: complex [T, rope_dim // 2]. + """ + x = x.contiguous() + if compressed_kv: + block_size = 16 + assert x.shape[-1] % block_size == 0 + assert rope_dim % 2 == 0 and rope_dim <= x.shape[-1] + d = x.shape[-1] + x2 = x.reshape(-1, d) + rows = x2.shape[0] + out = torch.empty_like(x) + if rows == 0: + return out + rows_per_token = rows // x.shape[0] + f_real = torch.view_as_real(freqs.contiguous()).contiguous() + _rope_tail_fake_quant_fp4_kernel[(rows,)]( + x2, + f_real, + out.reshape(-1, d), + x2.stride(0), + d, + f_real.stride(0), + rows_per_token, + D=d, + RD=rope_dim, + BLK=block_size, + AMAX_FLOOR=FP4_AMAX_FLOOR, + INVERSE=inverse, + COMPRESSED_KV=compressed_kv, + num_warps=4, + ) + return out diff --git a/python/sglang/kernels/ops/attention/dsv4/rope_pack_indexer.py b/python/sglang/kernels/ops/attention/dsv4/rope_pack_indexer.py new file mode 100644 index 000000000000..9f4770ca80fb --- /dev/null +++ b/python/sglang/kernels/ops/attention/dsv4/rope_pack_indexer.py @@ -0,0 +1,138 @@ +"""Fuse low-ratio RoPE, fake FP4 quantization and indexer packing/cache store. + +Keep both quantization stages: the indexer packer has a different scale floor +from fake_quant_fp4, so directly packing the first stage is not equivalent. +""" + +from __future__ import annotations + +import torch +import triton +import triton.language as tl +from triton.language.extra import libdevice + +from sglang.kernels.ops.attention.dsv4.fp4_indexer import ( + _ceil_ue8m0_exp, + _fp4_e2m1_code_rne, +) +from sglang.kernels.ops.attention.dsv4.rope_fake_quant_fp4 import FP4_AMAX_FLOOR + + +@triton.jit +def _rope_fake_quant_pack_indexer_kernel( + X, + F, + Pos, + Payload, + Scale, + Cache, + Loc, + F_STRIDE: tl.constexpr, + HEADS: tl.constexpr, + RD: tl.constexpr, + INDEXED: tl.constexpr, + STORE_CACHE: tl.constexpr, + PAGE_SIZE: tl.constexpr, + CACHE_STRIDE: tl.constexpr, + AMAX_FLOOR: tl.constexpr, +): + row = tl.program_id(0) + token = row // HEADS + frow = tl.load(Pos + token) if INDEXED else token + offsets = tl.arange(0, 128) + value = tl.load(X + row * 128 + offsets).to(tl.float32) + tail = offsets >= 128 - RD + pair = (offsets - (128 - RD)) // 2 + imaginary = (offsets % 2) == 1 + real = tl.load(X + row * 128 + 128 - RD + 2 * pair, tail, 0).to(tl.float32) + imag = tl.load(X + row * 128 + 128 - RD + 2 * pair + 1, tail, 0).to(tl.float32) + fr = tl.load(F + frow * F_STRIDE + 2 * pair, tail, 1.0) + fi = tl.load(F + frow * F_STRIDE + 2 * pair + 1, tail, 0.0) + rotated = tl.where(imaginary, real * fi + imag * fr, real * fr - imag * fi) + value = tl.where(tail, rotated.to(tl.bfloat16).to(tl.float32), value) + + blocks = tl.reshape(value, (4, 32)) + amax = tl.maximum(tl.max(tl.abs(blocks), 1), AMAX_FLOOR) * (1.0 / 6.0) + bits = amax.to(tl.int32, bitcast=True) + exponent = ((bits >> 23) & 0xFF) + ((bits & 0x7FFFFF) != 0).to(tl.int32) + fake_scale = (exponent << 23).to(tl.float32, bitcast=True) + scaled = tl.minimum(tl.maximum(blocks / fake_scale[:, None], -6.0), 6.0) + magnitude = tl.abs(scaled) + step = tl.where(magnitude < 2.0, 0.5, tl.where(magnitude < 4.0, 1.0, 2.0)) + sign = tl.where(scaled > 0, 1.0, tl.where(scaled < 0, -1.0, 0.0)) + rounded = libdevice.rint(magnitude / step) * step * sign + # Preserve the BF16 intermediate before recomputing the indexer scale, + # including its 1e-4 lower bound. + dequantized = (rounded * fake_scale[:, None]).to(tl.bfloat16).to(tl.float32) + pack_amax = tl.max(tl.abs(dequantized), 1) + pack_exponent = _ceil_ue8m0_exp(tl.maximum(pack_amax / 6.0, 1.0e-4)) + pack_scale = (pack_exponent << 23).to(tl.float32, bitcast=True) + codes = _fp4_e2m1_code_rne(dequantized / pack_scale[:, None]) + low, high = tl.split(tl.reshape(codes, (64, 2))) + payload = low | (high << 4) + sf = tl.sum(pack_exponent.to(tl.uint32) << (tl.arange(0, 4) * 8), 0) + byte_offsets = tl.arange(0, 64) + if STORE_CACHE: + location = tl.load(Loc + token) + page = location // PAGE_SIZE + slot = location % PAGE_SIZE + tl.store(Cache + page * CACHE_STRIDE + slot * 64 + byte_offsets, payload) + scale_bytes = (sf >> (tl.arange(0, 4) * 8)) & 0xFF + tl.store( + Cache + page * CACHE_STRIDE + PAGE_SIZE * 64 + slot * 4 + tl.arange(0, 4), + scale_bytes, + ) + else: + tl.store(Payload + row * 64 + byte_offsets, payload) + tl.store(Scale + row, sf) + + +def rope_fake_quant_pack_indexer( + x: torch.Tensor, + freqs: torch.Tensor, + rope_dim: int, + *, + positions: torch.Tensor | None = None, + cache: torch.Tensor | None = None, + loc: torch.Tensor | None = None, +) -> tuple[torch.Tensor, torch.Tensor] | None: + """Return packed [T*heads,64] / [T*heads], or write paged key cache. + + Without positions, freqs is already gathered per token; otherwise the + kernel reads freqs[positions] directly, removing the gather launch. + """ + assert x.dtype == torch.bfloat16 and x.shape[-1] == 128 + assert 0 <= rope_dim <= 128 and rope_dim % 2 == 0 + x = x.contiguous() + rows = x.numel() // 128 + heads = x[0].numel() // 128 if x.shape[0] else 1 + f = torch.view_as_real(freqs.contiguous()) + if cache is None: + payload = torch.empty((rows, 64), dtype=torch.int8, device=x.device) + scale = torch.empty((rows,), dtype=torch.int32, device=x.device) + page_size = cache_stride = 0 + else: + assert heads == 1 and loc is not None and loc.numel() == rows + assert cache.ndim == 2 and cache.shape[1] % 68 == 0 + payload = scale = None + page_size, cache_stride = cache.shape[1] // 68, cache.stride(0) + if rows: + _rope_fake_quant_pack_indexer_kernel[(rows,)]( + x, + f, + positions, + payload, + scale, + cache, + loc, + F_STRIDE=f.stride(0), + HEADS=heads, + RD=rope_dim, + INDEXED=positions is not None, + STORE_CACHE=cache is not None, + PAGE_SIZE=page_size, + CACHE_STRIDE=cache_stride, + AMAX_FLOOR=FP4_AMAX_FLOOR, + num_warps=4, + ) + return (payload, scale) if cache is None else None diff --git a/test/registered/kernel/attention/dsv4/test_compressed_kv_quant.py b/test/registered/kernel/attention/dsv4/test_compressed_kv_quant.py new file mode 100644 index 000000000000..2bc8407e905a --- /dev/null +++ b/test/registered/kernel/attention/dsv4/test_compressed_kv_quant.py @@ -0,0 +1,81 @@ +import unittest + +import torch + +from sglang.srt.layers.attention.dsv4.torch_quant import ( + fake_quant_compressed_kv, + fake_quant_fp4, +) +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci(est_time=20, stage="base-b-kernel-unit", runner_config="1-gpu-large") +register_cuda_ci(est_time=20, stage="base-b-kernel-unit", runner_config="4-gpu-b200") + + +def _rope_fq4(x, freqs, rope_dim, *, compressed_kv=False): + """RoPE plus fake FP4 quantization, fused for CUDA BF16 inputs.""" + if x.is_cuda and torch.version.cuda is not None and x.dtype == torch.bfloat16: + from sglang.kernels.ops.attention.dsv4.rope_fake_quant_fp4 import ( + rope_tail_fake_quant_fp4, + ) + + return rope_tail_fake_quant_fp4(x, freqs, rope_dim, compressed_kv=compressed_kv) + quant = fake_quant_compressed_kv if compressed_kv else fake_quant_fp4 + return quant(rope_tail(x, freqs, rope_dim)) + + +def rope_tail( + x: torch.Tensor, freqs: torch.Tensor, rope_dim: int, inverse: bool = False +) -> torch.Tensor: + """Rotate the last rope_dim features of x [T, ..., D] with complex freqs [T, rope_dim // 2].""" + head, tail = x[..., :-rope_dim], x[..., -rope_dim:] + tc = torch.view_as_complex(tail.float().unflatten(-1, (-1, 2)).contiguous()) + f = freqs.conj() if inverse else freqs + f = f.view(x.shape[0], *([1] * (x.ndim - 2)), rope_dim // 2) + rotated = torch.view_as_real(tc * f).flatten(-2).to(x.dtype) + return torch.cat([head, rotated], dim=-1) + + +class TestCompressedKVQuant(CustomTestCase): + @unittest.skipUnless(torch.cuda.is_available(), "requires CUDA") + def test_triton_matches_torch_for_both_quantization_rules(self): + from sglang.kernels.ops.attention.dsv4.rope_fake_quant_fp4 import ( + rope_tail_fake_quant_fp4, + ) + + generator = torch.Generator(device="cuda").manual_seed(17) + for rows in (0, 1, 33, 129): + x = torch.randn( + rows, 512, generator=generator, device="cuda", dtype=torch.bfloat16 + ) + angles = torch.randn(rows, 32, generator=generator, device="cuda") + freqs = torch.polar(torch.ones_like(angles), angles) + for compressed_kv in (False, True): + with self.subTest(rows=rows, compressed_kv=compressed_kv): + quant = ( + fake_quant_compressed_kv if compressed_kv else fake_quant_fp4 + ) + expected = quant(rope_tail(x, freqs, 64)) + actual = _rope_fq4(x, freqs, 64, compressed_kv=compressed_kv) + self.assertTrue(torch.equal(actual, expected)) + + # Identity RoPE isolates quantization boundaries from trigonometric rounding. + maxima = torch.tensor( + [0, 2**-12, 6 * 2**-9, 6 * 1.0625, 6 * 1.1875, 6 * 448, 1e6], + device="cuda", + dtype=torch.bfloat16, + ) + x = maxima[:, None].expand(-1, 512).contiguous() + freqs = torch.ones(x.shape[0], 32, device="cuda", dtype=torch.complex64) + actual = rope_tail_fake_quant_fp4(x, freqs, 64, compressed_kv=True) + expected = torch.tensor( + [0, 0, 6 * 2**-9, 6, 7.5, 2688, 2688], + device="cuda", + dtype=torch.bfloat16, + )[:, None].expand_as(x) + self.assertTrue(torch.equal(actual, expected)) + + +if __name__ == "__main__": + unittest.main() From e95bed5d72d2f8b18e8c5ef0cb923b82eb0f988b Mon Sep 17 00:00:00 2001 From: hnyls2002 Date: Tue, 15 Sep 2026 15:53:43 -0700 Subject: [PATCH 18/30] dsv4.1: extract Hopper FP8 matmul kernels and tuning --- ...0,dtype=fp8_w8a8,block_shape=[32, 32].json | 20 +++ ...0,dtype=fp8_w8a8,block_shape=[32, 32].json | 20 +++ ...0,dtype=fp8_w8a8,block_shape=[32, 32].json | 19 +++ ...0,dtype=fp8_w8a8,block_shape=[32, 32].json | 20 +++ ...0,dtype=fp8_w8a8,block_shape=[32, 32].json | 20 +++ ...0,dtype=fp8_w8a8,block_shape=[32, 32].json | 20 +++ ...0,dtype=fp8_w8a8,block_shape=[32, 32].json | 20 +++ .../kernels/ops/quantization/fp8_kernel.py | 152 +++++++++++++++++- .../quantization/test_fp8_hopper_swapab.py | 113 +++++++++++++ 9 files changed, 401 insertions(+), 3 deletions(-) create mode 100644 python/sglang/kernels/ops/quantization/configs/N=1152,K=5120,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json create mode 100644 python/sglang/kernels/ops/quantization/configs/N=1792,K=5120,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json create mode 100644 python/sglang/kernels/ops/quantization/configs/N=25600,K=6144,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json create mode 100644 python/sglang/kernels/ops/quantization/configs/N=4096,K=1280,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json create mode 100644 python/sglang/kernels/ops/quantization/configs/N=5120,K=2048,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json create mode 100644 python/sglang/kernels/ops/quantization/configs/N=5120,K=576,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json create mode 100644 python/sglang/kernels/ops/quantization/configs/N=8192,K=1280,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json create mode 100644 test/registered/kernel/quantization/test_fp8_hopper_swapab.py diff --git a/python/sglang/kernels/ops/quantization/configs/N=1152,K=5120,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json b/python/sglang/kernels/ops/quantization/configs/N=1152,K=5120,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json new file mode 100644 index 000000000000..80ed75f24e63 --- /dev/null +++ b/python/sglang/kernels/ops/quantization/configs/N=1152,K=5120,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json @@ -0,0 +1,20 @@ +{ + "1": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 32, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 4, + "SWAP_AB": true, + "SPLIT_K": 16 + }, + "2": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 32, + "BLOCK_SIZE_K": 32, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3 + } +} diff --git a/python/sglang/kernels/ops/quantization/configs/N=1792,K=5120,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json b/python/sglang/kernels/ops/quantization/configs/N=1792,K=5120,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json new file mode 100644 index 000000000000..fb5b60cdfb87 --- /dev/null +++ b/python/sglang/kernels/ops/quantization/configs/N=1792,K=5120,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json @@ -0,0 +1,20 @@ +{ + "1": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 32, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 4, + "SWAP_AB": true, + "SPLIT_K": 8 + }, + "2": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 32, + "BLOCK_SIZE_K": 32, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3 + } +} diff --git a/python/sglang/kernels/ops/quantization/configs/N=25600,K=6144,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json b/python/sglang/kernels/ops/quantization/configs/N=25600,K=6144,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json new file mode 100644 index 000000000000..91b3ff0dff75 --- /dev/null +++ b/python/sglang/kernels/ops/quantization/configs/N=25600,K=6144,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json @@ -0,0 +1,19 @@ +{ + "1": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 32, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 4, + "SWAP_AB": true + }, + "2": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 32, + "BLOCK_SIZE_K": 32, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3 + } +} diff --git a/python/sglang/kernels/ops/quantization/configs/N=4096,K=1280,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json b/python/sglang/kernels/ops/quantization/configs/N=4096,K=1280,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json new file mode 100644 index 000000000000..5d34539e4f05 --- /dev/null +++ b/python/sglang/kernels/ops/quantization/configs/N=4096,K=1280,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json @@ -0,0 +1,20 @@ +{ + "1": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 32, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 4, + "SWAP_AB": true, + "SPLIT_K": 4 + }, + "2": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 32, + "BLOCK_SIZE_K": 32, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3 + } +} diff --git a/python/sglang/kernels/ops/quantization/configs/N=5120,K=2048,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json b/python/sglang/kernels/ops/quantization/configs/N=5120,K=2048,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json new file mode 100644 index 000000000000..fb5b60cdfb87 --- /dev/null +++ b/python/sglang/kernels/ops/quantization/configs/N=5120,K=2048,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json @@ -0,0 +1,20 @@ +{ + "1": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 32, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 4, + "SWAP_AB": true, + "SPLIT_K": 8 + }, + "2": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 32, + "BLOCK_SIZE_K": 32, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3 + } +} diff --git a/python/sglang/kernels/ops/quantization/configs/N=5120,K=576,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json b/python/sglang/kernels/ops/quantization/configs/N=5120,K=576,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json new file mode 100644 index 000000000000..5d34539e4f05 --- /dev/null +++ b/python/sglang/kernels/ops/quantization/configs/N=5120,K=576,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json @@ -0,0 +1,20 @@ +{ + "1": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 32, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 4, + "SWAP_AB": true, + "SPLIT_K": 4 + }, + "2": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 32, + "BLOCK_SIZE_K": 32, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3 + } +} diff --git a/python/sglang/kernels/ops/quantization/configs/N=8192,K=1280,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json b/python/sglang/kernels/ops/quantization/configs/N=8192,K=1280,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json new file mode 100644 index 000000000000..b6c1003553ff --- /dev/null +++ b/python/sglang/kernels/ops/quantization/configs/N=8192,K=1280,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[32, 32].json @@ -0,0 +1,20 @@ +{ + "1": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 32, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 4, + "SWAP_AB": true, + "SPLIT_K": 2 + }, + "2": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 32, + "BLOCK_SIZE_K": 32, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3 + } +} diff --git a/python/sglang/kernels/ops/quantization/fp8_kernel.py b/python/sglang/kernels/ops/quantization/fp8_kernel.py index 5cfe48e09177..da92c781bc2c 100644 --- a/python/sglang/kernels/ops/quantization/fp8_kernel.py +++ b/python/sglang/kernels/ops/quantization/fp8_kernel.py @@ -26,6 +26,7 @@ from sglang.kernels.jit.utils import is_arch_support_pdl from sglang.kernels.ops.quantization.fp8_utils import fp8_dtype_to_triton from sglang.srt.layers import deep_gemm_wrapper +from sglang.srt.runtime_context import get_platform from sglang.srt.utils import ( ceil_align, get_bool_env_var, @@ -1029,6 +1030,133 @@ def _w8a8_block_fp8_matmul( tl.store(c_ptrs, c, mask=c_mask) +@triton.jit +def _w8a8_block_fp8_matmul_hopper( + # Pointers to inputs and output + A, + B, + C, + As, + Bs, + # Shape for matmul + M, + N, + K, + # Block size for block-wise quantization + group_n, + group_k, + # Stride for inputs and output + stride_am, + stride_ak, + stride_bk, + stride_bn, + stride_cm, + stride_cn, + stride_As_m, + stride_As_k, + stride_Bs_k, + stride_Bs_n, + # Meta-parameters + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, + BLOCK_SIZE_K: tl.constexpr, + GROUP_SIZE_M: tl.constexpr, + needs_masking: tl.constexpr, + SWAP_AB: tl.constexpr = False, + SPLIT_K: tl.constexpr = 1, +): + + pid = tl.program_id(axis=0) + split = tl.program_id(axis=1) + tiles_per_split = tl.cdiv(tl.cdiv(K, BLOCK_SIZE_K), SPLIT_K) + first_tile = split * tiles_per_split + C += split * M * N + num_pid_m = tl.cdiv(M, BLOCK_SIZE_M) + num_pid_n = tl.cdiv(N, BLOCK_SIZE_N) + num_pid_in_group = GROUP_SIZE_M * num_pid_n + group_id = pid // num_pid_in_group + first_pid_m = group_id * GROUP_SIZE_M + group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M) + pid_m = first_pid_m + (pid % group_size_m) + pid_n = (pid % num_pid_in_group) // group_size_m + + offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M + offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N + offs_k = tl.arange(0, BLOCK_SIZE_K) + a_ptrs = A + (offs_am[:, None] * stride_am + offs_k[None, :] * stride_ak) + b_ptrs = B + (offs_k[:, None] * stride_bk + offs_bn[None, :] * stride_bn) + + As_ptrs = As + offs_am * stride_As_m + offs_bsn = offs_bn // group_n + Bs_ptrs = Bs + offs_bsn * stride_Bs_n + n_tiles_k_per_group_k = group_k // BLOCK_SIZE_K + + a_ptrs += first_tile * BLOCK_SIZE_K * stride_ak + b_ptrs += first_tile * BLOCK_SIZE_K * stride_bk + As_ptrs += (first_tile // n_tiles_k_per_group_k) * stride_As_k + Bs_ptrs += (first_tile // n_tiles_k_per_group_k) * stride_Bs_k + + # Small-M Hopper configs transpose the MMA so the weight tile occupies M. + if SWAP_AB: + accumulator = tl.zeros((BLOCK_SIZE_N, BLOCK_SIZE_M), dtype=tl.float32) + else: + accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + for k in range( + first_tile, tl.minimum(first_tile + tiles_per_split, tl.cdiv(K, BLOCK_SIZE_K)) + ): + if needs_masking: + a = tl.load(a_ptrs, mask=offs_k[None, :] < K - k * BLOCK_SIZE_K, other=0.0) + b = tl.load(b_ptrs, mask=offs_k[:, None] < K - k * BLOCK_SIZE_K, other=0.0) + else: + a = tl.load(a_ptrs) + b = tl.load(b_ptrs) + + a_s = tl.load(As_ptrs) + b_s = tl.load(Bs_ptrs) + + scale_step_k = tl.where((k + 1) % n_tiles_k_per_group_k == 0, 1, 0) + if SWAP_AB: + accumulator += ( + tl.dot(tl.trans(b), tl.trans(a)) * b_s[:, None] * a_s[None, :] + ) + else: + accumulator += tl.dot(a, b) * a_s[:, None] * b_s[None, :] + a_ptrs += BLOCK_SIZE_K * stride_ak + b_ptrs += BLOCK_SIZE_K * stride_bk + As_ptrs += scale_step_k * stride_As_k + Bs_ptrs += scale_step_k * stride_Bs_k + + if SWAP_AB: + accumulator = tl.trans(accumulator) + + if C.dtype.element_ty == tl.bfloat16: + c = accumulator.to(tl.bfloat16) + elif C.dtype.element_ty == tl.float16: + c = accumulator.to(tl.float16) + else: + c = accumulator.to(tl.float32) + + offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + c_ptrs = C + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :] + c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N) + tl.store(c_ptrs, c, mask=c_mask) + + +@triton.jit +def _reduce_block_fp8_split_k( + Parts, Out, ELEMENTS: tl.constexpr, SPLITS: tl.constexpr, BLOCK: tl.constexpr +): + offsets = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + splits = tl.arange(0, SPLITS) + values = tl.load( + Parts + splits[:, None] * ELEMENTS + offsets[None, :], + offsets[None, :] < ELEMENTS, + 0.0, + ) + tl.store(Out + offsets, tl.sum(values, axis=0), offsets < ELEMENTS) + + @triton.jit def _w8a8_block_fp8_matmul_gfx1250( # Pointers to inputs and output @@ -1585,17 +1713,30 @@ def w8a8_block_fp8_matmul_triton( else: kernel = select_w8a8_block_fp8_matmul_kernel(M, N, config) + hopper_tuned = get_platform().is_sm90 and ( + config.get("SWAP_AB", False) or config.get("SPLIT_K", 1) > 1 + ) + if hopper_tuned: + kernel = _w8a8_block_fp8_matmul_hopper + split_k = config.get("SPLIT_K", 1) if hopper_tuned else 1 + if split_k > 1: + assert split_k & (split_k - 1) == 0 + partials = torch.empty((split_k, M, N), device=A.device, dtype=torch.float32) + else: + partials = C + needs_masking = bool(K % config["BLOCK_SIZE_K"] != 0) def grid(META): - return ( - triton.cdiv(M, META["BLOCK_SIZE_M"]) * triton.cdiv(N, META["BLOCK_SIZE_N"]), + blocks = triton.cdiv(M, META["BLOCK_SIZE_M"]) * triton.cdiv( + N, META["BLOCK_SIZE_N"] ) + return (blocks, split_k) if hopper_tuned else (blocks,) kernel[grid]( A, B, - C, + partials, As, Bs, M, @@ -1617,6 +1758,11 @@ def grid(META): needs_masking=needs_masking, ) + if split_k > 1: + _reduce_block_fp8_split_k[(triton.cdiv(M * N, 256),)]( + partials, C, M * N, split_k, 256 + ) + return C diff --git a/test/registered/kernel/quantization/test_fp8_hopper_swapab.py b/test/registered/kernel/quantization/test_fp8_hopper_swapab.py new file mode 100644 index 000000000000..aff33bd9c4fd --- /dev/null +++ b/test/registered/kernel/quantization/test_fp8_hopper_swapab.py @@ -0,0 +1,113 @@ +"""Small-M Hopper block-FP8 MMA transpose with nonuniform UE8M0 scales.""" + +import unittest +from unittest.mock import patch + +import torch + +from sglang.kernels.ops.quantization import fp8_kernel +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="1-gpu-large") + + +@unittest.skipUnless( + torch.cuda.is_available() and torch.cuda.get_device_capability() == (9, 0), + "Hopper required", +) +class TestHopperBlockFP8SwapAB(CustomTestCase): + def test_graph_replay_matches_original(self): + self._check_graph_replay(split_k=False) + + def test_split_k_replay_matches_original(self): + self._check_graph_replay(split_k=True) + + def _check_graph_replay(self, split_k): + original = dict( + BLOCK_SIZE_M=64, + BLOCK_SIZE_N=32, + BLOCK_SIZE_K=32, + GROUP_SIZE_M=32, + num_warps=4, + num_stages=3, + ) + for n, k in [ + (1152, 5120), + (1792, 5120), + (25600, 6144), + (4096, 1280), + (5120, 2048), + (5120, 576), + (8192, 1280), + ]: + with self.subTest(n=n, k=k): + torch.manual_seed(17) + a = torch.randn(1, k, device="cuda").to(torch.float8_e4m3fn) + b = torch.randn(n, k, device="cuda").to(torch.float8_e4m3fn) + sa = torch.exp2( + torch.randint(-5, 3, (1, k // 32), device="cuda").float() + ) + sb = torch.exp2( + torch.randint(-5, 3, (n // 32, k // 32), device="cuda").float() + ) + config = { + **original, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128 if n == 25600 or k == 576 else 64, + "num_stages": 4, + "SWAP_AB": True, + } + + if split_k and n != 25600: + splits = { + (1152, 5120): 16, + (1792, 5120): 8, + (4096, 1280): 4, + (5120, 2048): 8, + (5120, 576): 4, + (8192, 1280): 2, + } + config.update(SPLIT_K=splits[n, k], BLOCK_SIZE_N=64) + + def run(): + return fp8_kernel.w8a8_block_fp8_matmul_triton( + a, b, sa, sb, [32, 32], torch.bfloat16 + ) + + with patch.object( + fp8_kernel, "get_w8a8_block_fp8_configs", return_value={1: original} + ): + expected = run() + with patch.object( + fp8_kernel, "get_w8a8_block_fp8_configs", return_value={1: config} + ): + run() # Compile outside graph capture. + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + actual = run() + graph.replay() + torch.testing.assert_close( + actual, + expected, + rtol=0.008 if split_k else 0, + atol=0.0001 if split_k else 0, + ) + # Replay must read fresh activation payload and scales. + a.copy_(torch.randn(1, k, device="cuda").to(a.dtype)) + sa.mul_(2) + with patch.object( + fp8_kernel, "get_w8a8_block_fp8_configs", return_value={1: original} + ): + expected = run() + graph.replay() + torch.testing.assert_close( + actual, + expected, + rtol=0.008 if split_k else 0, + atol=0.0001 if split_k else 0, + ) + + +if __name__ == "__main__": + unittest.main() From c4bb0d63bc9cfac9fd0b4c3b4ac1f4dac7c45da5 Mon Sep 17 00:00:00 2001 From: hnyls2002 Date: Tue, 15 Sep 2026 15:53:44 -0700 Subject: [PATCH 19/30] dsv4.1: extract mHC computation and compensated projections --- .../ops/layernorm/hc_mix_stats_bf16x3.py | 100 +++++ .../ops/layernorm/hc_mix_stats_deepgemm.py | 67 ++++ python/sglang/kernels/ops/layernorm/mhc.py | 343 +++++++++++++++++- .../hyperconnection/test_compensated_mhc.py | 176 +++++++++ 4 files changed, 685 insertions(+), 1 deletion(-) create mode 100644 python/sglang/kernels/ops/layernorm/hc_mix_stats_bf16x3.py create mode 100644 python/sglang/kernels/ops/layernorm/hc_mix_stats_deepgemm.py create mode 100644 test/registered/kernel/hyperconnection/test_compensated_mhc.py diff --git a/python/sglang/kernels/ops/layernorm/hc_mix_stats_bf16x3.py b/python/sglang/kernels/ops/layernorm/hc_mix_stats_bf16x3.py new file mode 100644 index 000000000000..f6fe6000cd13 --- /dev/null +++ b/python/sglang/kernels/ops/layernorm/hc_mix_stats_bf16x3.py @@ -0,0 +1,100 @@ +"""Compensated mHC prefill projection with a shared activation load. + +Keep three BF16 components of the FP32 weights and accumulate their products +separately. The 16 fixed K slices bound FP32 accumulation error, as in the +compensated DeepGEMM path, while avoiding its second activation read/reduction. +""" + +import torch +import triton +import triton.language as tl + + +def split_bf16_hc_weight(weight: torch.Tensor): + assert weight.dtype == torch.float32 and weight.is_contiguous() + high = weight.bfloat16() + residual = weight - high.float() + middle = residual.bfloat16() + low = (residual - middle.float()).bfloat16() + return high, middle, low + + +@triton.jit +def _hc_mix_stats_bf16x3(X, W_HI, W_MID, W_LO, MIX, SQ, M, BLOCK_M: tl.constexpr): + # M stays runtime-valued so variable prefill lengths reuse the same binary. + rows = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M) + cols = tl.arange(0, 32) + # 20480 input features / 16 independent slices. + start = tl.program_id(1) * 1280 + ks = start + tl.arange(0, 64) + hi = tl.zeros((BLOCK_M, 32), tl.float32) + mid = tl.zeros((BLOCK_M, 32), tl.float32) + lo = tl.zeros((BLOCK_M, 32), tl.float32) + sq = tl.zeros((BLOCK_M,), tl.float32) + for block in range(20): + k = ks + block * 64 + x = tl.load( + X + rows[:, None].to(tl.int64) * 20480 + k[None, :], + rows[:, None] < M, + 0, + ) + offsets = cols[None, :] * 20480 + k[:, None] + w_hi = tl.load(W_HI + offsets, cols[None, :] < 24, 0) + w_mid = tl.load(W_MID + offsets, cols[None, :] < 24, 0) + w_lo = tl.load(W_LO + offsets, cols[None, :] < 24, 0) + hi = tl.dot(x, w_hi, hi) + mid = tl.dot(x, w_mid, mid) + lo = tl.dot(x, w_lo, lo) + xf = x.to(tl.float32) + sq += tl.sum(xf * xf, 1) + offsets = (tl.program_id(1) * M + rows[:, None]) * 24 + cols[None, :] + tl.store(MIX + offsets, (hi + mid) + lo, (rows[:, None] < M) & (cols[None, :] < 24)) + tl.store(SQ + tl.program_id(1) * M + rows, sq, rows < M) + + +def hc_mix_stats_sinkhorn_bf16x3( + x: torch.Tensor, + weight_parts, + scale: torch.Tensor, + base: torch.Tensor, + sinkhorn_iters: int, + rms_eps: float, + hc_eps: float, +): + from sglang.kernels.ops.layernorm.mhc import _hc_mix_reduce_sinkhorn_kernel + + m = x.shape[0] + assert x.shape == (m, 20480) and x.is_contiguous() + assert x.dtype == torch.bfloat16 and 4096 <= m <= 65536 + assert len(weight_parts) == 3 + assert all( + w.shape == (24, 20480) and w.dtype == torch.bfloat16 and w.is_contiguous() + for w in weight_parts + ) + mix = torch.empty((16, m, 24), device=x.device, dtype=torch.float32) + sq = torch.empty((16, m), device=x.device, dtype=torch.float32) + pre = torch.empty((m, 4), device=x.device, dtype=torch.float32) + post = torch.empty_like(pre) + comb = torch.empty((m, 4, 4), device=x.device, dtype=torch.float32) + _hc_mix_stats_bf16x3[(triton.cdiv(m, 128), 16)]( + x, *weight_parts, mix, sq, m, 128, num_warps=4, num_stages=3 + ) + _hc_mix_reduce_sinkhorn_kernel[(m,)]( + mix, + sq, + scale, + base, + pre, + post, + comb, + m, + 1.0 / 20480, + rms_eps, + MIX=24, + HC=4, + NUM_SLICES=16, + ITERS=sinkhorn_iters, + EPS=hc_eps, + num_warps=1, + ) + return pre, post, comb diff --git a/python/sglang/kernels/ops/layernorm/hc_mix_stats_deepgemm.py b/python/sglang/kernels/ops/layernorm/hc_mix_stats_deepgemm.py new file mode 100644 index 000000000000..6cb245e65469 --- /dev/null +++ b/python/sglang/kernels/ops/layernorm/hc_mix_stats_deepgemm.py @@ -0,0 +1,67 @@ +"""Compensated FP32 mHC projections for SM100 batches with at least 128 rows. + +The small-row and batch-invariant paths remain in mhc.py. Native TF32 discards +too much of the FP32 projection weights, so evaluate their high and residual +components separately and bound accumulation length with a fixed split count. +""" + +import torch + +_NUM_SPLITS = 16 + + +def split_tf32_hc_weight(weight: torch.Tensor): + assert weight.dtype == torch.float32 and weight.is_contiguous() + high = (weight.view(torch.int32) & -8192).view(torch.float32) + return high, weight - high + + +def hc_mix_stats_sinkhorn_deepgemm( + x_flat: torch.Tensor, + weight_parts, + hc_scale: torch.Tensor, + hc_base: torch.Tensor, + sinkhorn_iters: int, + rms_eps: float, + hc_eps: float, +): + from sglang.kernels.ops.layernorm.mhc import _hc_mix_reduce_sinkhorn_kernel + from sglang.srt.layers.deep_gemm_wrapper.entrypoint import tf32_hc_prenorm_gemm + + assert x_flat.dtype == torch.bfloat16 and x_flat.is_contiguous() + m, k = x_flat.shape + high, low = weight_parts + assert k == 20480 and high.shape == low.shape == (24, k) + dev = x_flat.device + pre = torch.empty((m, 4), dtype=torch.float32, device=dev) + post = torch.empty_like(pre) + comb = torch.empty((m, 4, 4), dtype=torch.float32, device=dev) + if m == 0: + return pre, post, comb + + mix_hi = torch.empty((_NUM_SPLITS, m, 24), dtype=torch.float32, device=dev) + mix_lo = torch.empty_like(mix_hi) + sq = torch.empty((_NUM_SPLITS, m), dtype=torch.float32, device=dev) + unused_sq = torch.empty_like(sq) + tf32_hc_prenorm_gemm(x_flat, high, mix_hi, sq, _NUM_SPLITS) + tf32_hc_prenorm_gemm(x_flat, low, mix_lo, unused_sq, _NUM_SPLITS) + _hc_mix_reduce_sinkhorn_kernel[(m,)]( + mix_hi, + sq, + hc_scale, + hc_base, + pre, + post, + comb, + m, + 1.0 / k, + rms_eps, + MIX=24, + HC=4, + NUM_SLICES=_NUM_SPLITS, + ITERS=sinkhorn_iters, + EPS=hc_eps, + part_mix_residual_ptr=mix_lo, + num_warps=1, + ) + return pre, post, comb diff --git a/python/sglang/kernels/ops/layernorm/mhc.py b/python/sglang/kernels/ops/layernorm/mhc.py index 42446c8379bc..2c7a89dc1025 100644 --- a/python/sglang/kernels/ops/layernorm/mhc.py +++ b/python/sglang/kernels/ops/layernorm/mhc.py @@ -18,6 +18,7 @@ from sglang.srt.layers.attention.dsa.utils import is_dsa_prefill_cp_interleave from sglang.srt.layers.dp_attention import is_allocation_symmetric from sglang.srt.layers.utils.common import strict_contiguous +from sglang.srt.runtime_context import get_platform from sglang.srt.utils.common import is_gfx1250_supported logger = logging.getLogger(__name__) @@ -362,13 +363,22 @@ def hc_split_sinkhorn( sinkhorn_iters: int = 20, eps: float = 1e-6, ): + b, s, _ = mixes.size() + if b * s == 0: + # DP attention's idle forward carries no tokens. Every backend below + # derives its grid from the token count, and CUDA rejects a launch with + # a zero-sized grid, so answer the empty batch directly. + return ( + mixes.new_empty(b, s, hc_mult), + mixes.new_empty(b, s, hc_mult), + mixes.new_empty(b, s, hc_mult, hc_mult), + ) if is_gfx1250_supported(): # TileLang's CK-backed addressing doesn't compile on gfx1250; use the # Triton port. _hc_split_sinkhorn_torch is kept as a reference fallback. return _hc_split_sinkhorn_triton( mixes, hc_scale, hc_base, hc_mult, sinkhorn_iters, eps ) - b, s, _ = mixes.size() pre = mixes.new_empty(b, s, hc_mult) post = mixes.new_empty(b, s, hc_mult) comb = mixes.new_empty(b, s, hc_mult, hc_mult) @@ -2027,6 +2037,337 @@ def _hc_combine_kernel( tl.store(y_ptr + pid_m * y_stride_m + offs_h, acc, mask=mask) +@triton.jit +def _hc_mix_stats_partial_kernel( + x_ptr, + w_ptr, + part_mix_ptr, + part_sq_ptr, + M, + K, + x_stride_m, + w_stride_n, + MIX: tl.constexpr, + MIX_PAD: tl.constexpr, + NUM_SLICES: tl.constexpr, + BLOCK_M: tl.constexpr, + BLOCK_K: tl.constexpr, + DOT_PRECISION: tl.constexpr, +): + """Mixing dot products and row sum of squares over one K slice; the slicing + and tiles are compile-time constants, so a row's fp32 operation sequence + does not depend on the batch size.""" + pid_m = tl.program_id(0) + pid_s = tl.program_id(1) + offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M) + offs_n = tl.arange(0, MIX_PAD) + mask_m = offs_m < M + mask_n = offs_n < MIX + k_per_slice = K // NUM_SLICES + k_start = pid_s * k_per_slice + acc = tl.zeros([BLOCK_M, MIX_PAD], dtype=tl.float32) + sq = tl.zeros([BLOCK_M], dtype=tl.float32) + for kb in range(0, k_per_slice, BLOCK_K): + offs_k = k_start + kb + tl.arange(0, BLOCK_K) + mask_k = offs_k < k_start + k_per_slice + x_tile = tl.load( + x_ptr + offs_m[:, None] * x_stride_m + offs_k[None, :], + mask=mask_m[:, None] & mask_k[None, :], + other=0.0, + ).to(tl.float32) + w_tile = tl.load( + w_ptr + offs_n[None, :] * w_stride_n + offs_k[:, None], + mask=mask_n[None, :] & mask_k[:, None], + other=0.0, + ).to(tl.float32) + acc += tl.dot(x_tile, w_tile, input_precision=DOT_PRECISION) + sq += tl.sum(x_tile * x_tile, axis=1) + tl.store( + part_mix_ptr + (pid_s * M + offs_m[:, None]) * MIX + offs_n[None, :], + acc, + mask=mask_m[:, None] & mask_n[None, :], + ) + tl.store(part_sq_ptr + pid_s * M + offs_m, sq, mask=mask_m) + + +@triton.jit +def _hc_mix_stats_reduce_kernel( + part_mix_ptr, + part_sq_ptr, + mixes_ptr, + M, + inv_k, + eps, + MIX: tl.constexpr, + MIX_PAD: tl.constexpr, + NUM_SLICES: tl.constexpr, + BLOCK_M: tl.constexpr, +): + """Sum the NUM_SLICES partials in slice order and apply the rms scaling.""" + pid_m = tl.program_id(0) + offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M) + offs_n = tl.arange(0, MIX_PAD) + mask_m = offs_m < M + mask_n = offs_n < MIX + acc = tl.zeros([BLOCK_M, MIX_PAD], dtype=tl.float32) + sq = tl.zeros([BLOCK_M], dtype=tl.float32) + for s in tl.static_range(NUM_SLICES): + acc += tl.load( + part_mix_ptr + (s * M + offs_m[:, None]) * MIX + offs_n[None, :], + mask=mask_m[:, None] & mask_n[None, :], + other=0.0, + ) + sq += tl.load(part_sq_ptr + s * M + offs_m, mask=mask_m, other=0.0) + rsqrt = 1.0 / tl.sqrt(sq * inv_k + eps) + tl.store( + mixes_ptr + offs_m[:, None] * MIX + offs_n[None, :], + acc * rsqrt[:, None], + mask=mask_m[:, None] & mask_n[None, :], + ) + + +# K slicing, BLOCK_K and dot precision must stay independent of M; +# a row must produce the same bits alone and in any batch. +_HC_MIX_SLICE_CHOICES = (80, 64, 40, 32, 16, 8, 4, 2, 1) +_HC_MIX_BLOCK_M = 32 +_HC_MIX_BLOCK_K = 64 +_HC_MIX_NUM_WARPS = 4 +_HC_MIX_DOT_PRECISION = "tf32x3" +# num_stages only reorders memory issue, not arithmetic; 2 is enough to cover the +# short k_per_slice loop (K=20480 gives 80 slices, i.e. 4 BLOCK_K tiles per CTA). +_HC_MIX_NUM_STAGES = 2 + +# BLOCK_M 8/16/32 preserve each row's K reduction order; thresholds were measured on GB300. +# Keep BLOCK_M below 64, where Triton lowers tf32x3 to plain TF32 and changes rounding. +_HC_MIX_BLOCK_M_SMALL = 8 +_HC_MIX_BLOCK_M_MID = 16 +_HC_MIX_MID_MAX_M = 2048 + + +def _block_m_for(m: int) -> int: + """Row-tile choices preserve each row's arithmetic and may depend on M.""" + if m <= _HC_MIX_BLOCK_M_SMALL: + return _HC_MIX_BLOCK_M_SMALL + if m <= _HC_MIX_MID_MAX_M: + return _HC_MIX_BLOCK_M_MID + return _HC_MIX_BLOCK_M + + +def _num_stages_for(m: int, k: int) -> int: + # GB300 verify batches benefit from a smaller shared-memory footprint. + # This changes memory scheduling only; K tiles and reduction order stay fixed. + if get_platform().is_blackwell and k == 20480 and 64 <= m <= 384: + return 1 + return _HC_MIX_NUM_STAGES + + +def _num_slices_for(k: int) -> int: + """Slice count depends only on K, never on batch size M.""" + blocks = k // _HC_MIX_BLOCK_K + assert k % _HC_MIX_BLOCK_K == 0, k + for n in _HC_MIX_SLICE_CHOICES: + if blocks % n == 0: + return n + return 1 + + +def hc_mix_stats(x_flat: torch.Tensor, hc_fn: torch.Tensor, eps: float) -> torch.Tensor: + """Batch-invariant F.linear(x_flat.float(), hc_fn) * rsqrt(mean(x_flat^2) + eps). + + x_flat is [M, K] in any float dtype; hc_fn is [MIX, K] fp32; returns [M, MIX] fp32. + K slicing and reduction order are independent of M, so each row is bitwise + identical whether computed alone or in a batch. + """ + assert x_flat.dim() == 2 and hc_fn.dim() == 2 + assert x_flat.stride(1) == 1 and hc_fn.stride(1) == 1 + assert hc_fn.dtype == torch.float32 + m, k = x_flat.shape + mix = hc_fn.shape[0] + assert hc_fn.shape[1] == k + num_slices = _num_slices_for(k) + mix_pad = max(16, triton.next_power_of_2(mix)) + part_mix = torch.empty( + (num_slices, m, mix), dtype=torch.float32, device=x_flat.device + ) + part_sq = torch.empty((num_slices, m), dtype=torch.float32, device=x_flat.device) + mixes = torch.empty((m, mix), dtype=torch.float32, device=x_flat.device) + if m == 0: + return mixes + block_m = _block_m_for(m) + grid_m = triton.cdiv(m, block_m) + _hc_mix_stats_partial_kernel[(grid_m, num_slices)]( + x_flat, + hc_fn, + part_mix, + part_sq, + m, + k, + x_flat.stride(0), + hc_fn.stride(0), + MIX=mix, + MIX_PAD=mix_pad, + NUM_SLICES=num_slices, + BLOCK_M=block_m, + BLOCK_K=_HC_MIX_BLOCK_K, + DOT_PRECISION=_HC_MIX_DOT_PRECISION, + num_warps=_HC_MIX_NUM_WARPS, + num_stages=_num_stages_for(m, k), + ) + _hc_mix_stats_reduce_kernel[(grid_m,)]( + part_mix, + part_sq, + mixes, + m, + 1.0 / k, + eps, + MIX=mix, + MIX_PAD=mix_pad, + NUM_SLICES=num_slices, + BLOCK_M=block_m, + num_warps=4, + ) + return mixes + + +@triton.jit +def _hc_mix_reduce_sinkhorn_kernel( + part_mix_ptr, + part_sq_ptr, + scale_ptr, + base_ptr, + pre_ptr, + post_ptr, + comb_ptr, + m, + inv_k, + rms_eps, + MIX: tl.constexpr, + HC: tl.constexpr, + NUM_SLICES: tl.constexpr, + ITERS: tl.constexpr, + EPS: tl.constexpr, + part_mix_residual_ptr=None, +): + """One CTA per row keeps the sinkhorn reductions two-dimensional. + Per-row arithmetic follows the slice reduction, then the Triton sinkhorn. + """ + row = tl.program_id(0) + if row >= m: + return + j = tl.arange(0, HC) + jj = j[:, None] + kk = j[None, :] + + a_pre = tl.zeros([HC], dtype=tl.float32) + a_post = tl.zeros([HC], dtype=tl.float32) + a_comb = tl.zeros([HC, HC], dtype=tl.float32) + sq = tl.zeros([], dtype=tl.float32) + for s in tl.static_range(NUM_SLICES): + off = (s * m + row) * MIX + v_pre = tl.load(part_mix_ptr + off + j) + v_post = tl.load(part_mix_ptr + off + HC + j) + v_comb = tl.load(part_mix_ptr + off + 2 * HC + jj * HC + kk) + if part_mix_residual_ptr is not None: + v_pre += tl.load(part_mix_residual_ptr + off + j) + v_post += tl.load(part_mix_residual_ptr + off + HC + j) + v_comb += tl.load(part_mix_residual_ptr + off + 2 * HC + jj * HC + kk) + a_pre += v_pre + a_post += v_post + a_comb += v_comb + sq += tl.load(part_sq_ptr + s * m + row) + rsqrt = 1.0 / tl.sqrt(sq * inv_k + rms_eps) + + s0 = tl.load(scale_ptr + 0) + s1 = tl.load(scale_ptr + 1) + s2 = tl.load(scale_ptr + 2) + + pre = tl.sigmoid(a_pre * rsqrt * s0 + tl.load(base_ptr + j)) + EPS + tl.store(pre_ptr + row * HC + j, pre) + post = 2.0 * tl.sigmoid(a_post * rsqrt * s1 + tl.load(base_ptr + HC + j)) + tl.store(post_ptr + row * HC + j, post) + + comb = a_comb * rsqrt * s2 + tl.load(base_ptr + 2 * HC + jj * HC + kk) + comb = tl.exp(comb - tl.max(comb, axis=1)[:, None]) + comb = comb / tl.sum(comb, axis=1)[:, None] + EPS + comb = comb / (tl.sum(comb, axis=0)[None, :] + EPS) + for _ in tl.static_range(ITERS - 1): + comb = comb / (tl.sum(comb, axis=1)[:, None] + EPS) + comb = comb / (tl.sum(comb, axis=0)[None, :] + EPS) + tl.store(comb_ptr + row * HC * HC + jj * HC + kk, comb) + + +def hc_mix_stats_sinkhorn( + x_flat: torch.Tensor, + hc_fn: torch.Tensor, + hc_scale: torch.Tensor, + hc_base: torch.Tensor, + hc_mult: int, + sinkhorn_iters: int, + rms_eps: float, + hc_eps: float, +): + """Fuse the reduce and sinkhorn stages of hc_mix_stats followed by hc_split_sinkhorn. + + The split-K kernel fixes the reduction order and preserves batch invariance. + Sinkhorn uses the Triton port's transcendental lowering, which differs from TileLang. + """ + assert x_flat.dim() == 2 and hc_fn.dim() == 2 + assert x_flat.stride(1) == 1 and hc_fn.stride(1) == 1 + assert hc_fn.dtype == torch.float32 + m, k = x_flat.shape + mix = hc_fn.shape[0] + assert mix == (2 + hc_mult) * hc_mult and hc_fn.shape[1] == k + dev = x_flat.device + pre = torch.empty(m, hc_mult, dtype=torch.float32, device=dev) + post = torch.empty(m, hc_mult, dtype=torch.float32, device=dev) + comb = torch.empty(m, hc_mult, hc_mult, dtype=torch.float32, device=dev) + if m == 0: + return pre, post, comb + + num_slices = _num_slices_for(k) + mix_pad = max(16, triton.next_power_of_2(mix)) + part_mix = torch.empty((num_slices, m, mix), dtype=torch.float32, device=dev) + part_sq = torch.empty((num_slices, m), dtype=torch.float32, device=dev) + block_m = _block_m_for(m) + _hc_mix_stats_partial_kernel[(triton.cdiv(m, block_m), num_slices)]( + x_flat, + hc_fn, + part_mix, + part_sq, + m, + k, + x_flat.stride(0), + hc_fn.stride(0), + MIX=mix, + MIX_PAD=mix_pad, + NUM_SLICES=num_slices, + BLOCK_M=block_m, + BLOCK_K=_HC_MIX_BLOCK_K, + DOT_PRECISION=_HC_MIX_DOT_PRECISION, + num_warps=_HC_MIX_NUM_WARPS, + num_stages=_num_stages_for(m, k), + ) + _hc_mix_reduce_sinkhorn_kernel[(m,)]( + part_mix, + part_sq, + hc_scale.float().contiguous(), + hc_base.float().contiguous(), + pre, + post, + comb, + m, + 1.0 / k, + rms_eps, + MIX=mix, + HC=hc_mult, + NUM_SLICES=num_slices, + ITERS=sinkhorn_iters, + EPS=hc_eps, + num_warps=1, + ) + return pre, post, comb + + def hc_combine( x_flat: torch.Tensor, pre: torch.Tensor, hc: int, out_dtype: torch.dtype ) -> torch.Tensor: diff --git a/test/registered/kernel/hyperconnection/test_compensated_mhc.py b/test/registered/kernel/hyperconnection/test_compensated_mhc.py new file mode 100644 index 000000000000..d0884c840678 --- /dev/null +++ b/test/registered/kernel/hyperconnection/test_compensated_mhc.py @@ -0,0 +1,176 @@ +import sys + +import pytest +import torch + +from sglang.kernels.ops.layernorm.hc_mix_stats_deepgemm import ( + hc_mix_stats_sinkhorn_deepgemm, + split_tf32_hc_weight, +) +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=120, stage="base-b-kernel-unit", runner_config="4-gpu-b200") +pytestmark = pytest.mark.skipif( + not torch.cuda.is_available() or torch.cuda.get_device_capability()[0] != 10, + reason="Compensated mHC path targets datacenter Blackwell", +) +EPS = 1e-6 + + +def inputs(m, seed): + torch.manual_seed(seed) + x = torch.randn((m, 20480), device="cuda", dtype=torch.bfloat16) + w = torch.randn((24, 20480), device="cuda", dtype=torch.float32) * 0.02 + scale = torch.tensor([0.1, 0.2, 0.3], device="cuda") + base = torch.randn(24, device="cuda", dtype=torch.float32) * 0.2 + return x, w, scale, base + + +def reference(x, w, scale, base): + x, w, scale, base = (v.double() for v in (x, w, scale, base)) + mixes = (x @ w.T) * torch.rsqrt(x.square().mean(-1, keepdim=True) + EPS) + pre = torch.sigmoid(mixes[:, :4] * scale[0] + base[:4]) + EPS + post = 2 * torch.sigmoid(mixes[:, 4:8] * scale[1] + base[4:8]) + comb = (mixes[:, 8:] * scale[2] + base[8:]).view(-1, 4, 4) + comb = torch.softmax(comb, dim=-1) + EPS + comb = comb / (comb.sum(-2, keepdim=True) + EPS) + for _ in range(19): + comb = comb / (comb.sum(-1, keepdim=True) + EPS) + comb = comb / (comb.sum(-2, keepdim=True) + EPS) + return pre, post, comb + + +@pytest.mark.parametrize("m", [0, 128, 384, 2049, 4096, 16384, 32768, 65536]) +@pytest.mark.parametrize("seed", [0, 42]) +def test_compensated_coefficients_match_fp64(m, seed): + x, w, scale, base = inputs(m, seed) + parts = split_tf32_hc_weight(w) + assert torch.equal(parts[0] + parts[1], w) + got = hc_mix_stats_sinkhorn_deepgemm(x, parts, scale, base, 20, EPS, EPS) + expected = reference(x, w, scale, base) + for actual, ref in zip(got, expected): + assert actual.dtype == torch.float32 + torch.testing.assert_close(actual.double(), ref, rtol=2e-5, atol=2e-6) + + +@pytest.mark.parametrize("m", [384, 4096]) +def test_graph_replay_reads_updated_input(m): + x, w, scale, base = inputs(m, 13) + parts = split_tf32_hc_weight(w) + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(stream): + for _ in range(3): + hc_mix_stats_sinkhorn_deepgemm(x, parts, scale, base, 20, EPS, EPS) + torch.cuda.current_stream().wait_stream(stream) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + captured = hc_mix_stats_sinkhorn_deepgemm(x, parts, scale, base, 20, EPS, EPS) + # Scaling alone almost cancels under RMS normalization and would not expose + # a replay that accidentally kept using the capture-time activations. + x.normal_() + graph.replay() + expected = reference(x, w, scale, base) + for actual, ref in zip(captured, expected): + torch.testing.assert_close(actual.double(), ref, rtol=2e-5, atol=2e-6) + + +@pytest.mark.parametrize("m", [1, 6, 64, 127]) +def test_original_sinkhorn_path_without_residual_matches_fp64(m): + from sglang.kernels.ops.layernorm.mhc import hc_mix_stats_sinkhorn + + x, w, scale, base = inputs(m, 7) + got = hc_mix_stats_sinkhorn(x, w, scale, base, 4, 20, EPS, EPS) + expected = reference(x, w, scale, base) + for actual, ref in zip(got, expected): + torch.testing.assert_close(actual.double(), ref, rtol=2e-5, atol=2e-6) + + +@pytest.mark.parametrize("m", [128, 384, 4096]) +def test_fused_compensation_preserves_epilogue_bits(m): + from sglang.kernels.ops.layernorm.mhc import _hc_mix_reduce_sinkhorn_kernel + + torch.manual_seed(1) + hi = torch.randn(16, m, 24, device="cuda") + lo = torch.randn_like(hi) * 0.001 + sq = torch.rand(16, m, device="cuda") * 1280 + scale = torch.tensor([0.1, 0.2, 0.3], device="cuda") + base = torch.randn(24, device="cuda") + + def reduce(partial, residual=None): + pre = torch.empty(m, 4, device="cuda") + post = torch.empty_like(pre) + comb = torch.empty(m, 4, 4, device="cuda") + _hc_mix_reduce_sinkhorn_kernel[(m,)]( + partial, + sq, + scale, + base, + pre, + post, + comb, + m, + 1.0 / 20480, + EPS, + MIX=24, + HC=4, + NUM_SLICES=16, + ITERS=20, + EPS=EPS, + part_mix_residual_ptr=residual, + num_warps=1, + ) + return pre, post, comb + + expected = reduce(hi + lo) + actual = reduce(hi, lo) + for got, ref in zip(actual, expected): + torch.testing.assert_close(got, ref, rtol=0, atol=0) + + +@pytest.mark.parametrize("m", [4096, 4097, 16384, 65536]) +@pytest.mark.parametrize("seed", [0, 42]) +def test_bf16x3_matches_compensated_and_fp64(m, seed): + from sglang.kernels.ops.layernorm.hc_mix_stats_bf16x3 import ( + hc_mix_stats_sinkhorn_bf16x3, + split_bf16_hc_weight, + ) + + x, w, scale, base = inputs(m, seed) + parts = split_bf16_hc_weight(w) + torch.testing.assert_close(sum(p.float() for p in parts), w, rtol=2e-7, atol=0) + actual = hc_mix_stats_sinkhorn_bf16x3(x, parts, scale, base, 20, EPS, EPS) + original = hc_mix_stats_sinkhorn_deepgemm( + x, split_tf32_hc_weight(w), scale, base, 20, EPS, EPS + ) + indices = torch.cat( + (torch.arange(32, device="cuda"), torch.arange(m - 32, m, device="cuda")) + ) + expected = reference(x[indices], w, scale, base) + for a, b, ref in zip(actual, original, expected): + assert torch.isfinite(a).all() + torch.testing.assert_close(a, b, rtol=2e-5, atol=2e-6) + torch.testing.assert_close(a[indices].double(), ref, rtol=2e-5, atol=2e-6) + + +def test_bf16x3_graph_replay(): + from sglang.kernels.ops.layernorm.hc_mix_stats_bf16x3 import ( + hc_mix_stats_sinkhorn_bf16x3, + split_bf16_hc_weight, + ) + + x, w, scale, base = inputs(4097, 13) + parts = split_bf16_hc_weight(w) + hc_mix_stats_sinkhorn_bf16x3(x, parts, scale, base, 20, EPS, EPS) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + actual = hc_mix_stats_sinkhorn_bf16x3(x, parts, scale, base, 20, EPS, EPS) + for _ in range(3): + x.normal_() + graph.replay() + for a, b in zip(actual, reference(x[-32:], w, scale, base)): + torch.testing.assert_close(a[-32:].double(), b, rtol=2e-5, atol=2e-6) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__])) From 0a66478130fc36a4bd0a08ff0e149730cac9c9ab Mon Sep 17 00:00:00 2001 From: BBuf <1182563586@qq.com> Date: Wed, 16 Sep 2026 15:31:57 +0800 Subject: [PATCH 20/30] remove standalone sparse indexer test --- .../dsv4/test_dsv41_sparse_indexer.py | 234 ------------------ 1 file changed, 234 deletions(-) delete mode 100644 test/registered/kernel/attention/dsv4/test_dsv41_sparse_indexer.py diff --git a/test/registered/kernel/attention/dsv4/test_dsv41_sparse_indexer.py b/test/registered/kernel/attention/dsv4/test_dsv41_sparse_indexer.py deleted file mode 100644 index 6ae103d0825e..000000000000 --- a/test/registered/kernel/attention/dsv4/test_dsv41_sparse_indexer.py +++ /dev/null @@ -1,234 +0,0 @@ -"""Decode two-level indexer on DeepGEMM's paged sparse MQA logits. - -The block table layer 20 publishes (``amax_topk_blocks``) is checked against -the model code's ``select_candidate_blocks``; a consumer's sparse logits -are checked against DeepGEMM's dense bf16 paged logits gathered at the published -positions, which the kernel is documented to match bitwise; its selection is -checked against a torch top-k of those. The kernel path needs a DeepGEMM with -``fp8_fp4_paged_sparse_mqa_logits`` on an SM100 device; it skips elsewhere. -""" - -import unittest - -import torch - -from sglang.test.ci.ci_register import register_cuda_ci -from sglang.test.test_utils import CustomTestCase - -register_cuda_ci(est_time=120, stage="base-b-kernel-unit", runner_config="4-gpu-b200") - -HEADS = 32 -HEAD_DIM = 128 -PAGE = 128 # index pool page on SM100: 128 * 68 bytes = 17 * 512 -TOPK = 512 -BLOCKS = 2048 - - -def _sparse_indexer_available() -> bool: - if not torch.cuda.is_available() or torch.cuda.get_device_capability()[0] != 10: - return False - try: - import deep_gemm - except ImportError: - return False - return hasattr(deep_gemm, "fp8_fp4_paged_sparse_mqa_logits") - - -def _reference_blocks(logits, lens, block_size=8): - """Block ids the model code keeps, per row, ascending.""" - from sglang.srt.layers.attention.dsv4.indexer import select_candidate_blocks - - width = logits.shape[1] - reach = torch.arange(width, device=logits.device)[None, :] < lens[:, None] - mask = select_candidate_blocks( - logits.masked_fill(~reach, -torch.inf), - lens[:, None], - topk_blocks=BLOCKS, - block_size=block_size, - ) - return [m.view(-1, block_size).any(-1).nonzero().flatten() for m in mask] - - -class TestSparseIndexer(CustomTestCase): - def test_amax_topk_blocks_matches_reference(self): - # short rows: the block top-k skips its plan; 40 rows of 300K tokens: - # 37500 keys per row on a batch above the persistent pool, plan needed - self._check_amax_topk_blocks( - torch.tensor( - [1, 37, 16384, 16389, 40000, 131072], dtype=torch.int32, device="cuda" - ), - 131072, - ) - self._check_amax_topk_blocks( - torch.randint(200000, 300001, (40,), dtype=torch.int32, device="cuda"), - 300000, - ) - - def _check_amax_topk_blocks(self, lens, width): - from sglang.kernels.ops.attention.dsv4.candidate_blocks import ( - candidate_row_lens, - ) - from sglang.kernels.ops.attention.dsv4.topk import sort_candidate_blocks - from sglang.srt.layers.attention.dsv4.candidate_deep_gemm import ( - amax_topk_blocks, - valid_lens, - ) - - torch.manual_seed(0) - bs = lens.numel() - logits = torch.randn(bs, width, device="cuda") - # the tail past the length is garbage in production: make it loud - logits.masked_fill_( - torch.arange(width, device="cuda")[None, :] >= lens[:, None], 1e4 - ) - pages = (width + PAGE - 1) // PAGE - page_table = torch.stack( - [torch.randperm(pages, device="cuda") for _ in range(bs)] - ).to(torch.int32) - nblocks, valid = candidate_row_lens(lens, BLOCKS) - self.assertTrue(torch.equal(nblocks, (lens + 7) // 8)) - self.assertTrue(torch.equal(valid, valid_lens(lens, BLOCKS))) - blocks = amax_topk_blocks(logits, lens, nblocks, BLOCKS) - phys = sort_candidate_blocks(blocks, lens, page_table, PAGE) - keys = logits.view(bs, -1, 8).amax(-1) - bpp = PAGE // 8 - for b, ref in enumerate(_reference_blocks(logits, lens)): - nb = (int(lens[b]) + 7) // 8 - n = min(nb, BLOCKS) - got = blocks[b, :n].long() - self.assertTrue(torch.equal(got, got.sort().values), "not ascending") - self.assertEqual(got.unique().numel(), n) - self.assertIn(nb - 1, got.tolist(), "newest block not kept") - self.assertTrue(bool((got < nb).all())) - # equal keys may swap blocks: compare the key multiset (the forced - # block excluded, its key is arbitrary garbage) - keep = got != nb - 1 - keep_ref = ref != nb - 1 - self.assertTrue( - torch.equal( - keys[b][got[keep]].sort().values, - keys[b][ref[keep_ref]].sort().values, - ), - msg=f"row {b}", - ) - # past the valid count nothing looks like a block DeepGEMM could read - self.assertTrue(bool((blocks[b, n:] >= nb).all())) - # the same blocks as pool slots / 8 through the row's page table - ref_phys = page_table[b][got // bpp].long() * bpp + got % bpp - self.assertTrue(torch.equal(phys[b, :n].long(), ref_phys)) - self.assertTrue(bool((phys[b, n:] == torch.iinfo(torch.int32).max).all())) - expect_valid = 8 * (n - 1) + ((int(lens[b]) - 1) % 8 + 1) - self.assertEqual(int(valid[b]), expect_valid) - - @unittest.skipUnless( - _sparse_indexer_available(), "needs DeepGEMM's paged sparse MQA logits on SM100" - ) - def test_paired_verify_rows_match_unpaired(self): - """Verify shape: each request has 6 consecutive rows (its draft tokens) with - lengths L, L+1, ..., sharing one page-table row. With request ids DeepGEMM - pairs the rows on one KV pass; the sparse logits must equal the unpaired - (every row its own request) result bitwise, row layout unchanged.""" - import deep_gemm - - from sglang.kernels.ops.attention.dsv4.candidate_blocks import ( - candidate_row_lens, - ) - from sglang.kernels.ops.attention.dsv4.topk import ( - sort_candidate_blocks, - ) - from sglang.srt.layers.attention.dsv4.candidate_deep_gemm import ( - SparseBlockTable, - amax_topk_blocks, - build_sparse_indexer_schedule, - sparse_logits, - valid_lens, - ) - - torch.manual_seed(3) - draft = 6 - base = torch.tensor([20000, 16385, 70000], dtype=torch.int32, device="cuda") - lens = ( - base[:, None] + torch.arange(draft, device="cuda", dtype=torch.int32) - ).flatten() - request_ids = torch.repeat_interleave( - torch.tensor([7, 3, 11], dtype=torch.int64, device="cuda"), draft - ) - rows = lens.numel() - max_pages = (int(lens.max()) + PAGE - 1) // PAGE - num_pages = base.numel() * max_pages - pool = torch.randint( - 0, 255, (num_pages, PAGE * 68), dtype=torch.uint8, device="cuda" - ) - pool[:, PAGE * 64 :] = torch.randint( - 118, 123, (num_pages, PAGE * 4), dtype=torch.uint8, device="cuda" - ) - k_cache = pool.view(num_pages, PAGE, 1, 68) - per_request = ( - torch.randperm(num_pages, device="cuda") - .view(base.numel(), max_pages) - .to(torch.int32) - ) - page_table = per_request.repeat_interleave(draft, dim=0) # rows share theirs - q_fp4 = torch.randint( - 0, 255, (rows, 1, HEADS, HEAD_DIM // 2), dtype=torch.uint8, device="cuda" - ).view(torch.int8) - q_sf = ( - torch.randint( - 118, 123, (rows, 1, HEADS, 4), dtype=torch.uint8, device="cuda" - ) - .view(torch.int32) - .squeeze(-1) - ) - weights = (torch.rand(rows, HEADS, device="cuda") * 0.05).to(torch.bfloat16) - sched = deep_gemm.get_paged_mqa_logits_metadata( - lens.view(-1, 1), PAGE, deep_gemm.get_num_sms() - ) - dense = deep_gemm.fp8_fp4_paged_mqa_logits( - (q_fp4, q_sf), - k_cache, - weights.float(), - lens.view(-1, 1), - page_table, - sched, - int(lens.max()), - False, - torch.float32, - ) - nblocks, row_valid = candidate_row_lens(lens, BLOCKS) - blocks = amax_topk_blocks(dense, lens, nblocks, BLOCKS) - phys = sort_candidate_blocks(blocks, lens, page_table, PAGE) - out = {} - for name, ids in (("paired", request_ids), ("unpaired", None)): - schedule = build_sparse_indexer_schedule( - blocks, lens, page_table, PAGE, q_fp4.dtype, ids - ) - table = SparseBlockTable( - blocks=blocks, schedule=schedule, phys_blocks=phys, valid_lens=row_valid - ) - out[name] = sparse_logits(q_fp4, q_sf, k_cache, weights, table) - cols = torch.arange(BLOCKS * 8, device="cuda") - valid = cols[None, :] < row_valid[:, None].long() - self.assertTrue(torch.equal(row_valid, valid_lens(lens, BLOCKS))) - self.assertTrue( - torch.equal(out["paired"][valid], out["unpaired"][valid]), - "pairing changed the sparse logits", - ) - # and both equal the dense bf16 logits at the published positions - dense16 = deep_gemm.fp8_fp4_paged_mqa_logits( - (q_fp4, q_sf), - k_cache, - weights, - lens.view(-1, 1), - page_table, - sched, - int(lens.max()), - False, - torch.bfloat16, - ) - pos = blocks.long().repeat_interleave(8, dim=1) * 8 + (cols % 8)[None, :] - ref = dense16.gather(1, pos.clamp(max=dense16.shape[1] - 1)) - self.assertTrue(torch.equal(out["paired"][valid], ref[valid])) - - -if __name__ == "__main__": - unittest.main() From 294d3e841e19055ef5c9411af06e45725a542a1f Mon Sep 17 00:00:00 2001 From: BBuf <1182563586@qq.com> Date: Wed, 16 Sep 2026 15:32:25 +0800 Subject: [PATCH 21/30] remove standalone compressed KV quant test --- .../dsv4/test_compressed_kv_quant.py | 81 ------------------- 1 file changed, 81 deletions(-) delete mode 100644 test/registered/kernel/attention/dsv4/test_compressed_kv_quant.py diff --git a/test/registered/kernel/attention/dsv4/test_compressed_kv_quant.py b/test/registered/kernel/attention/dsv4/test_compressed_kv_quant.py deleted file mode 100644 index 2bc8407e905a..000000000000 --- a/test/registered/kernel/attention/dsv4/test_compressed_kv_quant.py +++ /dev/null @@ -1,81 +0,0 @@ -import unittest - -import torch - -from sglang.srt.layers.attention.dsv4.torch_quant import ( - fake_quant_compressed_kv, - fake_quant_fp4, -) -from sglang.test.ci.ci_register import register_cuda_ci -from sglang.test.test_utils import CustomTestCase - -register_cuda_ci(est_time=20, stage="base-b-kernel-unit", runner_config="1-gpu-large") -register_cuda_ci(est_time=20, stage="base-b-kernel-unit", runner_config="4-gpu-b200") - - -def _rope_fq4(x, freqs, rope_dim, *, compressed_kv=False): - """RoPE plus fake FP4 quantization, fused for CUDA BF16 inputs.""" - if x.is_cuda and torch.version.cuda is not None and x.dtype == torch.bfloat16: - from sglang.kernels.ops.attention.dsv4.rope_fake_quant_fp4 import ( - rope_tail_fake_quant_fp4, - ) - - return rope_tail_fake_quant_fp4(x, freqs, rope_dim, compressed_kv=compressed_kv) - quant = fake_quant_compressed_kv if compressed_kv else fake_quant_fp4 - return quant(rope_tail(x, freqs, rope_dim)) - - -def rope_tail( - x: torch.Tensor, freqs: torch.Tensor, rope_dim: int, inverse: bool = False -) -> torch.Tensor: - """Rotate the last rope_dim features of x [T, ..., D] with complex freqs [T, rope_dim // 2].""" - head, tail = x[..., :-rope_dim], x[..., -rope_dim:] - tc = torch.view_as_complex(tail.float().unflatten(-1, (-1, 2)).contiguous()) - f = freqs.conj() if inverse else freqs - f = f.view(x.shape[0], *([1] * (x.ndim - 2)), rope_dim // 2) - rotated = torch.view_as_real(tc * f).flatten(-2).to(x.dtype) - return torch.cat([head, rotated], dim=-1) - - -class TestCompressedKVQuant(CustomTestCase): - @unittest.skipUnless(torch.cuda.is_available(), "requires CUDA") - def test_triton_matches_torch_for_both_quantization_rules(self): - from sglang.kernels.ops.attention.dsv4.rope_fake_quant_fp4 import ( - rope_tail_fake_quant_fp4, - ) - - generator = torch.Generator(device="cuda").manual_seed(17) - for rows in (0, 1, 33, 129): - x = torch.randn( - rows, 512, generator=generator, device="cuda", dtype=torch.bfloat16 - ) - angles = torch.randn(rows, 32, generator=generator, device="cuda") - freqs = torch.polar(torch.ones_like(angles), angles) - for compressed_kv in (False, True): - with self.subTest(rows=rows, compressed_kv=compressed_kv): - quant = ( - fake_quant_compressed_kv if compressed_kv else fake_quant_fp4 - ) - expected = quant(rope_tail(x, freqs, 64)) - actual = _rope_fq4(x, freqs, 64, compressed_kv=compressed_kv) - self.assertTrue(torch.equal(actual, expected)) - - # Identity RoPE isolates quantization boundaries from trigonometric rounding. - maxima = torch.tensor( - [0, 2**-12, 6 * 2**-9, 6 * 1.0625, 6 * 1.1875, 6 * 448, 1e6], - device="cuda", - dtype=torch.bfloat16, - ) - x = maxima[:, None].expand(-1, 512).contiguous() - freqs = torch.ones(x.shape[0], 32, device="cuda", dtype=torch.complex64) - actual = rope_tail_fake_quant_fp4(x, freqs, 64, compressed_kv=True) - expected = torch.tensor( - [0, 0, 6 * 2**-9, 6, 7.5, 2688, 2688], - device="cuda", - dtype=torch.bfloat16, - )[:, None].expand_as(x) - self.assertTrue(torch.equal(actual, expected)) - - -if __name__ == "__main__": - unittest.main() From f4cff50294f9b1df19f181c168d659aeb01a6749 Mon Sep 17 00:00:00 2001 From: BBuf <1182563586@qq.com> Date: Wed, 16 Sep 2026 15:32:52 +0800 Subject: [PATCH 22/30] remove standalone Hopper FP8 test --- .../quantization/test_fp8_hopper_swapab.py | 113 ------------------ 1 file changed, 113 deletions(-) delete mode 100644 test/registered/kernel/quantization/test_fp8_hopper_swapab.py diff --git a/test/registered/kernel/quantization/test_fp8_hopper_swapab.py b/test/registered/kernel/quantization/test_fp8_hopper_swapab.py deleted file mode 100644 index aff33bd9c4fd..000000000000 --- a/test/registered/kernel/quantization/test_fp8_hopper_swapab.py +++ /dev/null @@ -1,113 +0,0 @@ -"""Small-M Hopper block-FP8 MMA transpose with nonuniform UE8M0 scales.""" - -import unittest -from unittest.mock import patch - -import torch - -from sglang.kernels.ops.quantization import fp8_kernel -from sglang.test.ci.ci_register import register_cuda_ci -from sglang.test.test_utils import CustomTestCase - -register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="1-gpu-large") - - -@unittest.skipUnless( - torch.cuda.is_available() and torch.cuda.get_device_capability() == (9, 0), - "Hopper required", -) -class TestHopperBlockFP8SwapAB(CustomTestCase): - def test_graph_replay_matches_original(self): - self._check_graph_replay(split_k=False) - - def test_split_k_replay_matches_original(self): - self._check_graph_replay(split_k=True) - - def _check_graph_replay(self, split_k): - original = dict( - BLOCK_SIZE_M=64, - BLOCK_SIZE_N=32, - BLOCK_SIZE_K=32, - GROUP_SIZE_M=32, - num_warps=4, - num_stages=3, - ) - for n, k in [ - (1152, 5120), - (1792, 5120), - (25600, 6144), - (4096, 1280), - (5120, 2048), - (5120, 576), - (8192, 1280), - ]: - with self.subTest(n=n, k=k): - torch.manual_seed(17) - a = torch.randn(1, k, device="cuda").to(torch.float8_e4m3fn) - b = torch.randn(n, k, device="cuda").to(torch.float8_e4m3fn) - sa = torch.exp2( - torch.randint(-5, 3, (1, k // 32), device="cuda").float() - ) - sb = torch.exp2( - torch.randint(-5, 3, (n // 32, k // 32), device="cuda").float() - ) - config = { - **original, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 128 if n == 25600 or k == 576 else 64, - "num_stages": 4, - "SWAP_AB": True, - } - - if split_k and n != 25600: - splits = { - (1152, 5120): 16, - (1792, 5120): 8, - (4096, 1280): 4, - (5120, 2048): 8, - (5120, 576): 4, - (8192, 1280): 2, - } - config.update(SPLIT_K=splits[n, k], BLOCK_SIZE_N=64) - - def run(): - return fp8_kernel.w8a8_block_fp8_matmul_triton( - a, b, sa, sb, [32, 32], torch.bfloat16 - ) - - with patch.object( - fp8_kernel, "get_w8a8_block_fp8_configs", return_value={1: original} - ): - expected = run() - with patch.object( - fp8_kernel, "get_w8a8_block_fp8_configs", return_value={1: config} - ): - run() # Compile outside graph capture. - graph = torch.cuda.CUDAGraph() - with torch.cuda.graph(graph): - actual = run() - graph.replay() - torch.testing.assert_close( - actual, - expected, - rtol=0.008 if split_k else 0, - atol=0.0001 if split_k else 0, - ) - # Replay must read fresh activation payload and scales. - a.copy_(torch.randn(1, k, device="cuda").to(a.dtype)) - sa.mul_(2) - with patch.object( - fp8_kernel, "get_w8a8_block_fp8_configs", return_value={1: original} - ): - expected = run() - graph.replay() - torch.testing.assert_close( - actual, - expected, - rtol=0.008 if split_k else 0, - atol=0.0001 if split_k else 0, - ) - - -if __name__ == "__main__": - unittest.main() From e4a3c13147f683042de1ee81a55f3c7371023014 Mon Sep 17 00:00:00 2001 From: BBuf <1182563586@qq.com> Date: Wed, 16 Sep 2026 15:33:21 +0800 Subject: [PATCH 23/30] remove standalone compensated mHC test --- .../hyperconnection/test_compensated_mhc.py | 176 ------------------ 1 file changed, 176 deletions(-) delete mode 100644 test/registered/kernel/hyperconnection/test_compensated_mhc.py diff --git a/test/registered/kernel/hyperconnection/test_compensated_mhc.py b/test/registered/kernel/hyperconnection/test_compensated_mhc.py deleted file mode 100644 index d0884c840678..000000000000 --- a/test/registered/kernel/hyperconnection/test_compensated_mhc.py +++ /dev/null @@ -1,176 +0,0 @@ -import sys - -import pytest -import torch - -from sglang.kernels.ops.layernorm.hc_mix_stats_deepgemm import ( - hc_mix_stats_sinkhorn_deepgemm, - split_tf32_hc_weight, -) -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=120, stage="base-b-kernel-unit", runner_config="4-gpu-b200") -pytestmark = pytest.mark.skipif( - not torch.cuda.is_available() or torch.cuda.get_device_capability()[0] != 10, - reason="Compensated mHC path targets datacenter Blackwell", -) -EPS = 1e-6 - - -def inputs(m, seed): - torch.manual_seed(seed) - x = torch.randn((m, 20480), device="cuda", dtype=torch.bfloat16) - w = torch.randn((24, 20480), device="cuda", dtype=torch.float32) * 0.02 - scale = torch.tensor([0.1, 0.2, 0.3], device="cuda") - base = torch.randn(24, device="cuda", dtype=torch.float32) * 0.2 - return x, w, scale, base - - -def reference(x, w, scale, base): - x, w, scale, base = (v.double() for v in (x, w, scale, base)) - mixes = (x @ w.T) * torch.rsqrt(x.square().mean(-1, keepdim=True) + EPS) - pre = torch.sigmoid(mixes[:, :4] * scale[0] + base[:4]) + EPS - post = 2 * torch.sigmoid(mixes[:, 4:8] * scale[1] + base[4:8]) - comb = (mixes[:, 8:] * scale[2] + base[8:]).view(-1, 4, 4) - comb = torch.softmax(comb, dim=-1) + EPS - comb = comb / (comb.sum(-2, keepdim=True) + EPS) - for _ in range(19): - comb = comb / (comb.sum(-1, keepdim=True) + EPS) - comb = comb / (comb.sum(-2, keepdim=True) + EPS) - return pre, post, comb - - -@pytest.mark.parametrize("m", [0, 128, 384, 2049, 4096, 16384, 32768, 65536]) -@pytest.mark.parametrize("seed", [0, 42]) -def test_compensated_coefficients_match_fp64(m, seed): - x, w, scale, base = inputs(m, seed) - parts = split_tf32_hc_weight(w) - assert torch.equal(parts[0] + parts[1], w) - got = hc_mix_stats_sinkhorn_deepgemm(x, parts, scale, base, 20, EPS, EPS) - expected = reference(x, w, scale, base) - for actual, ref in zip(got, expected): - assert actual.dtype == torch.float32 - torch.testing.assert_close(actual.double(), ref, rtol=2e-5, atol=2e-6) - - -@pytest.mark.parametrize("m", [384, 4096]) -def test_graph_replay_reads_updated_input(m): - x, w, scale, base = inputs(m, 13) - parts = split_tf32_hc_weight(w) - stream = torch.cuda.Stream() - stream.wait_stream(torch.cuda.current_stream()) - with torch.cuda.stream(stream): - for _ in range(3): - hc_mix_stats_sinkhorn_deepgemm(x, parts, scale, base, 20, EPS, EPS) - torch.cuda.current_stream().wait_stream(stream) - graph = torch.cuda.CUDAGraph() - with torch.cuda.graph(graph): - captured = hc_mix_stats_sinkhorn_deepgemm(x, parts, scale, base, 20, EPS, EPS) - # Scaling alone almost cancels under RMS normalization and would not expose - # a replay that accidentally kept using the capture-time activations. - x.normal_() - graph.replay() - expected = reference(x, w, scale, base) - for actual, ref in zip(captured, expected): - torch.testing.assert_close(actual.double(), ref, rtol=2e-5, atol=2e-6) - - -@pytest.mark.parametrize("m", [1, 6, 64, 127]) -def test_original_sinkhorn_path_without_residual_matches_fp64(m): - from sglang.kernels.ops.layernorm.mhc import hc_mix_stats_sinkhorn - - x, w, scale, base = inputs(m, 7) - got = hc_mix_stats_sinkhorn(x, w, scale, base, 4, 20, EPS, EPS) - expected = reference(x, w, scale, base) - for actual, ref in zip(got, expected): - torch.testing.assert_close(actual.double(), ref, rtol=2e-5, atol=2e-6) - - -@pytest.mark.parametrize("m", [128, 384, 4096]) -def test_fused_compensation_preserves_epilogue_bits(m): - from sglang.kernels.ops.layernorm.mhc import _hc_mix_reduce_sinkhorn_kernel - - torch.manual_seed(1) - hi = torch.randn(16, m, 24, device="cuda") - lo = torch.randn_like(hi) * 0.001 - sq = torch.rand(16, m, device="cuda") * 1280 - scale = torch.tensor([0.1, 0.2, 0.3], device="cuda") - base = torch.randn(24, device="cuda") - - def reduce(partial, residual=None): - pre = torch.empty(m, 4, device="cuda") - post = torch.empty_like(pre) - comb = torch.empty(m, 4, 4, device="cuda") - _hc_mix_reduce_sinkhorn_kernel[(m,)]( - partial, - sq, - scale, - base, - pre, - post, - comb, - m, - 1.0 / 20480, - EPS, - MIX=24, - HC=4, - NUM_SLICES=16, - ITERS=20, - EPS=EPS, - part_mix_residual_ptr=residual, - num_warps=1, - ) - return pre, post, comb - - expected = reduce(hi + lo) - actual = reduce(hi, lo) - for got, ref in zip(actual, expected): - torch.testing.assert_close(got, ref, rtol=0, atol=0) - - -@pytest.mark.parametrize("m", [4096, 4097, 16384, 65536]) -@pytest.mark.parametrize("seed", [0, 42]) -def test_bf16x3_matches_compensated_and_fp64(m, seed): - from sglang.kernels.ops.layernorm.hc_mix_stats_bf16x3 import ( - hc_mix_stats_sinkhorn_bf16x3, - split_bf16_hc_weight, - ) - - x, w, scale, base = inputs(m, seed) - parts = split_bf16_hc_weight(w) - torch.testing.assert_close(sum(p.float() for p in parts), w, rtol=2e-7, atol=0) - actual = hc_mix_stats_sinkhorn_bf16x3(x, parts, scale, base, 20, EPS, EPS) - original = hc_mix_stats_sinkhorn_deepgemm( - x, split_tf32_hc_weight(w), scale, base, 20, EPS, EPS - ) - indices = torch.cat( - (torch.arange(32, device="cuda"), torch.arange(m - 32, m, device="cuda")) - ) - expected = reference(x[indices], w, scale, base) - for a, b, ref in zip(actual, original, expected): - assert torch.isfinite(a).all() - torch.testing.assert_close(a, b, rtol=2e-5, atol=2e-6) - torch.testing.assert_close(a[indices].double(), ref, rtol=2e-5, atol=2e-6) - - -def test_bf16x3_graph_replay(): - from sglang.kernels.ops.layernorm.hc_mix_stats_bf16x3 import ( - hc_mix_stats_sinkhorn_bf16x3, - split_bf16_hc_weight, - ) - - x, w, scale, base = inputs(4097, 13) - parts = split_bf16_hc_weight(w) - hc_mix_stats_sinkhorn_bf16x3(x, parts, scale, base, 20, EPS, EPS) - graph = torch.cuda.CUDAGraph() - with torch.cuda.graph(graph): - actual = hc_mix_stats_sinkhorn_bf16x3(x, parts, scale, base, 20, EPS, EPS) - for _ in range(3): - x.normal_() - graph.replay() - for a, b in zip(actual, reference(x[-32:], w, scale, base)): - torch.testing.assert_close(a[-32:].double(), b, rtol=2e-5, atol=2e-6) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__])) From 56efc4a2d75815ce1fa40e07d90d66a5fab2ab30 Mon Sep 17 00:00:00 2001 From: hnyls2002 Date: Wed, 16 Sep 2026 03:03:02 -0700 Subject: [PATCH 24/30] drop norm-only c2 entry and aliases; move small metadata into dsv4; name low-ratio metadata fields --- .../kernels/jit/csrc/deepseek_v4/c1.cuh | 2 +- .../kernels/jit/csrc/deepseek_v4/c2.cuh | 214 +++++++----------- .../csrc/deepseek_v4/fused_norm_rope_v2.cuh | 2 +- .../jit/csrc/deepseek_v4/main_norm_rope.cuh | 2 +- .../kernels/jit/csrc/deepseek_v4/store.cuh | 2 +- .../sglang/kernels/ops/attention/dsv4/c2.py | 77 +------ .../kernels/ops/attention/dsv4/kv_layout.py | 3 +- .../small_metadata.py} | 19 +- 8 files changed, 111 insertions(+), 210 deletions(-) rename python/sglang/kernels/ops/attention/{dsv41_small_metadata.py => dsv4/small_metadata.py} (86%) diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/c1.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/c1.cuh index 6dac3e3c1f47..ea76c8630448 100644 --- a/python/sglang/kernels/jit/csrc/deepseek_v4/c1.cuh +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/c1.cuh @@ -316,7 +316,7 @@ struct FlashC1DecodeKernel { } }; -// ensure that C++ wrapper can work +// The JIT module names and wrappers spell the layouts as bare enumerators. using enum deepseek_v4::KVLayout; } // namespace sglang diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/c2.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/c2.cuh index c50e1283e143..6f34c6be3eb1 100644 --- a/python/sglang/kernels/jit/csrc/deepseek_v4/c2.cuh +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/c2.cuh @@ -15,7 +15,6 @@ #include #include -#include namespace sglang { @@ -66,7 +65,6 @@ constexpr uint32_t kC2VecSize = 2; /// stores the e2m1 codes and their e4m3 scales directly, so the fp4 rounding /// happens once and no fp8 rounding follows it. template < - bool kStore, bool kVerify, int64_t kHeadDim, int64_t kRopeDim, @@ -140,9 +138,7 @@ __global__ __launch_bounds__(kHeadDim / kC2VecSize) void flash_c2_decode_kernel( fp32_vec_t staged, freq; bf16_vec_t weight, out; weight.load(params.norm_weight, tx); - if constexpr (kStore) { - if (tx >= kNopeThreads) freq.load(params.freqs_cis + (pos - 1) * kRopeDim, tx - kNopeThreads); - } + if (tx >= kNopeThreads) freq.load(params.freqs_cis + (pos - 1) * kRopeDim, tx - kNopeThreads); // With two scores `exp(-|s0 - s1|)` is the whole softmax: one exp, argument // always <= 0, so no max-subtraction pass and no overflow. @@ -191,94 +187,91 @@ __global__ __launch_bounds__(kHeadDim / kC2VecSize) void flash_c2_decode_kernel( out.store(params.kv_output, static_cast(row) * kCTASize + tx); PDLTriggerSecondary(); - if constexpr (kStore) { - // ---- main-KV branch: RoPE tail, fp4 fake-quant, 584-byte store ---- - // Match finish()'s bf16 rounding before RoPE. + // ---- main-KV branch: RoPE tail, fp4 fake-quant, 584-byte store ---- + // Match finish()'s bf16 rounding before RoPE. #pragma unroll - for (uint32_t i = 0; i < kVecSize / 2; ++i) { - const auto [x, y] = cast(out[i]); - staged[i * 2 + 0] = x; - staged[i * 2 + 1] = y; - } + for (uint32_t i = 0; i < kVecSize / 2; ++i) { + const auto [x, y] = cast(out[i]); + staged[i * 2 + 0] = x; + staged[i * 2 + 1] = y; + } - if (tx >= kNopeThreads) { - // Match rope_tail()'s bf16 rounding before fake quantization. - // Only odd positions reach here; the latent represents `pos - 1`. - freq.load(params.freqs_cis + (pos - 1) * kRopeDim, tx - kNopeThreads); + if (tx >= kNopeThreads) { + // Match rope_tail()'s bf16 rounding before fake quantization. + // Only odd positions reach here; the latent represents `pos - 1`. + freq.load(params.freqs_cis + (pos - 1) * kRopeDim, tx - kNopeThreads); #pragma unroll - for (uint32_t i = 0; i < kVecSize / 2; ++i) { - const auto x_real = staged[i * 2 + 0]; - const auto x_imag = staged[i * 2 + 1]; - const auto f_real = x_real * freq[i * 2 + 0] - x_imag * freq[i * 2 + 1]; - const auto f_imag = x_real * freq[i * 2 + 1] + x_imag * freq[i * 2 + 0]; - const auto rotated = cast(fp32x2_t{f_real, f_imag}); - const auto [r0, r1] = cast(rotated); - staged[i * 2 + 0] = r0; - staged[i * 2 + 1] = r1; - } + for (uint32_t i = 0; i < kVecSize / 2; ++i) { + const auto x_real = staged[i * 2 + 0]; + const auto x_imag = staged[i * 2 + 1]; + const auto f_real = x_real * freq[i * 2 + 0] - x_imag * freq[i * 2 + 1]; + const auto f_imag = x_real * freq[i * 2 + 1] + x_imag * freq[i * 2 + 0]; + const auto rotated = cast(fp32x2_t{f_real, f_imag}); + const auto [r0, r1] = cast(rotated); + staged[i * 2 + 0] = r0; + staged[i * 2 + 1] = r1; } + } - if constexpr (kLayout == KVLayout::V41_FP4) { - // The fp4 cache takes the rotated bf16 value as is: its row quantizer is the - // fake quantization, minus the dequantization. - const int32_t out_loc = raw_out_loc >> 1; - const auto kv_row = Paged::row(params.kvcache, out_loc); - return deepseek_v4::v41::store_row(kv_row.data, kv_row.scale, tx, staged); - } + if constexpr (kLayout == KVLayout::V41_FP4) { + // The fp4 cache takes the rotated bf16 value as is: its row quantizer is the + // fake quantization, minus the dequantization. + const int32_t out_loc = raw_out_loc >> 1; + const auto kv_row = Paged::row(params.kvcache, out_loc); + return deepseek_v4::v41::store_row(kv_row.data, kv_row.scale, tx, staged); + } - // FP4/E4M3 fake-quant over 16 elements, i.e. kFp4Lanes threads. - { - float amax = fabsf(staged[0]); + // FP4/E4M3 fake-quant over 16 elements, i.e. kFp4Lanes threads. + { + float amax = fabsf(staged[0]); #pragma unroll - for (uint32_t i = 1; i < kVecSize; ++i) { - amax = fmaxf(amax, fabsf(staged[i])); - } - amax = warp::reduce_max(amax); - const auto scale = deepseek_v4::fp4::compressed_kv_scale(amax); + for (uint32_t i = 1; i < kVecSize; ++i) { + amax = fmaxf(amax, fabsf(staged[i])); + } + amax = warp::reduce_max(amax); + const auto scale = deepseek_v4::fp4::compressed_kv_scale(amax); #pragma unroll - for (uint32_t i = 0; i < kVecSize / 2; ++i) { - const auto [x, y] = - deepseek_v4::fp4::fake_quant_compressed_kv_x2({staged[i * 2 + 0], staged[i * 2 + 1]}, scale); - staged[i * 2 + 0] = x; - staged[i * 2 + 1] = y; - } + for (uint32_t i = 0; i < kVecSize / 2; ++i) { + const auto [x, y] = deepseek_v4::fp4::fake_quant_compressed_kv_x2({staged[i * 2 + 0], staged[i * 2 + 1]}, scale); + staged[i * 2 + 0] = x; + staged[i * 2 + 1] = y; } + } - // `raw_out_loc / ratio`; ratio 2 makes it a shift. - const int32_t out_loc = raw_out_loc >> 1; - const auto kv_row = Paged::row(params.kvcache, out_loc); + // `raw_out_loc / ratio`; ratio 2 makes it a shift. + const int32_t out_loc = raw_out_loc >> 1; + const auto kv_row = Paged::row(params.kvcache, out_loc); - if constexpr (kLayout == KVLayout::V41) { - // fp8 with one ue8m0 scale per 32 elements over the whole row, RoPE included. - return deepseek_v4::v41::store_row(kv_row.data, kv_row.scale, tx, staged); - } + if constexpr (kLayout == KVLayout::V41) { + // fp8 with one ue8m0 scale per 32 elements over the whole row, RoPE included. + return deepseek_v4::v41::store_row(kv_row.data, kv_row.scale, tx, staged); + } - const auto value_ptr = kv_row.data; + const auto value_ptr = kv_row.data; - if (tx >= kNopeThreads) { - bf16_vec_t rope_out; + if (tx >= kNopeThreads) { + bf16_vec_t rope_out; #pragma unroll - for (uint32_t i = 0; i < kVecSize / 2; ++i) { - rope_out[i] = cast(fp32x2_t{staged[i * 2 + 0], staged[i * 2 + 1]}); - } - rope_out.store(value_ptr + (kHeadDim - kRopeDim), tx - kNopeThreads); - } else { - // fp8 e4m3 with one ue8m0 scale per 64 elements. - auto abs_max = fabsf(staged[0]); + for (uint32_t i = 0; i < kVecSize / 2; ++i) { + rope_out[i] = cast(fp32x2_t{staged[i * 2 + 0], staged[i * 2 + 1]}); + } + rope_out.store(value_ptr + (kHeadDim - kRopeDim), tx - kNopeThreads); + } else { + // fp8 e4m3 with one ue8m0 scale per 64 elements. + auto abs_max = fabsf(staged[0]); #pragma unroll - for (uint32_t i = 1; i < kVecSize; ++i) { - abs_max = fmaxf(abs_max, fabsf(staged[i])); - } - abs_max = warp::reduce_max(abs_max); - const auto scale_ue8m0 = cast_to_ue8m0(fmaxf(1e-4f, abs_max) / math::FP8_E4M3_MAX); - const auto inv_scale = inv_scale_ue8m0(scale_ue8m0); + for (uint32_t i = 1; i < kVecSize; ++i) { + abs_max = fmaxf(abs_max, fabsf(staged[i])); + } + abs_max = warp::reduce_max(abs_max); + const auto scale_ue8m0 = cast_to_ue8m0(fmaxf(1e-4f, abs_max) / math::FP8_E4M3_MAX); + const auto inv_scale = inv_scale_ue8m0(scale_ue8m0); #pragma unroll - for (uint32_t i = 0; i < kVecSize / 2; ++i) { - reinterpret_cast(value_ptr)[tx * (kVecSize / 2) + i] = - pack_fp8(staged[i * 2 + 0] * inv_scale, staged[i * 2 + 1] * inv_scale); - } - kv_row.scale[tx / kFp8Lanes] = scale_ue8m0; + for (uint32_t i = 0; i < kVecSize / 2; ++i) { + reinterpret_cast(value_ptr)[tx * (kVecSize / 2) + i] = + pack_fp8(staged[i * 2 + 0] * inv_scale, staged[i * 2 + 1] * inv_scale); } + kv_row.scale[tx / kFp8Lanes] = scale_ue8m0; } } @@ -288,15 +281,15 @@ struct FlashC2DecodeKernel { static constexpr int32_t kPageBits = std::bit_width(kPageSize) - 1; static constexpr int64_t kPageBytes = deepseek_v4::kv_page_bytes(kPageSize); static_assert(kLayout != deepseek_v4::KVLayout::V4 || kPageBytes == host::div_ceil(584ll * kPageSize, 576) * 576); - template + template static constexpr auto kernel = - flash_c2_decode_kernel; + flash_c2_decode_kernel; /// \brief The (`positions`, `raw_out_loc`) dtype pair, resolved at run time. - template + template static auto select(const bool pos_i32, const bool loc_i32) { - if (pos_i32) return loc_i32 ? kernel : kernel; - return loc_i32 ? kernel : kernel; + if (pos_i32) return loc_i32 ? kernel : kernel; + return loc_i32 ? kernel : kernel; } // The sum of squares is reduced through a fixed-size shared array, so the CTA @@ -304,32 +297,6 @@ struct FlashC2DecodeKernel { static_assert(kHeadDim % (4 * device::kWarpThreads) == 0, "head_dim must be a multiple of 128"); static_assert(std::has_single_bit(kPageSize), "the page/slot split needs a power-of-two page"); - /// \brief Pool + norm only. The main-KV write stays with the caller. - static void run_decode( - const tvm::ffi::TensorView kv_input, - const tvm::ffi::TensorView kv_state, - const tvm::ffi::TensorView kv_output, - const tvm::ffi::TensorView norm_weight, - const tvm::ffi::TensorView positions, - const tvm::ffi::TensorView req, - const tvm::ffi::TensorView raw_out_loc, - const float eps, - const int64_t ring_size) { - launch( - kv_input, - kv_state, - kv_output, - norm_weight, - positions, - req, - raw_out_loc, - eps, - ring_size, - std::nullopt, - std::nullopt, - /*draft_len=*/1); - } - /// \brief `run_decode_fusion` for a target-verify block. /// /// `draft_len` consecutive positions per request, request-major, which the @@ -363,8 +330,6 @@ struct FlashC2DecodeKernel { } private: - using MaybeTensor = std::optional; - static void launch( const tvm::ffi::TensorView kv_input, const tvm::ffi::TensorView kv_state, @@ -375,8 +340,8 @@ struct FlashC2DecodeKernel { const tvm::ffi::TensorView raw_out_loc, const float eps, const int64_t ring_size, - const MaybeTensor freqs_cis, - const MaybeTensor kvcache, + const tvm::ffi::TensorView freqs_cis, + const tvm::ffi::TensorView kvcache, const int64_t draft_len) { using namespace host; @@ -397,14 +362,11 @@ struct FlashC2DecodeKernel { TensorMatcher({N}).with_dtype().with_device(device_).verify(req); TensorMatcher({N}).with_dtype(loc_dtype).with_device(device_).verify(raw_out_loc); - const auto store = freqs_cis.has_value(); - if (store) { - // Real/imag interleaved, so the trailing dim is kRopeDim, not kRopeDim / 2. - TensorMatcher({-1, kRopeDim}).with_dtype().with_device(device_).verify(*freqs_cis); - // The pool allocates the buffer as uint8 and hands it out viewed as its - // fp8 dtype (`get_extra_key_buffer`); both are one byte per element. - TensorMatcher({-1, kPageBytes}).with_dtype().with_device(device_).verify(*kvcache); - } + // Real/imag interleaved, so the trailing dim is kRopeDim, not kRopeDim / 2. + TensorMatcher({-1, kRopeDim}).with_dtype().with_device(device_).verify(freqs_cis); + // The pool allocates the buffer as uint8 and hands it out viewed as its + // fp8 dtype (`get_extra_key_buffer`); both are one byte per element. + TensorMatcher({-1, kPageBytes}).with_dtype().with_device(device_).verify(kvcache); const auto num_tokens = static_cast(N.unwrap()); if (num_tokens == 0) return; @@ -418,11 +380,11 @@ struct FlashC2DecodeKernel { .kv_state = static_cast(kv_state.data_ptr()), .kv_output = static_cast(kv_output.data_ptr()), .norm_weight = static_cast(norm_weight.data_ptr()), - .freqs_cis = store ? static_cast(freqs_cis->data_ptr()) : nullptr, + .freqs_cis = static_cast(freqs_cis.data_ptr()), .positions = positions.data_ptr(), .req = static_cast(req.data_ptr()), .raw_out_loc = raw_out_loc.data_ptr(), - .kvcache = store ? static_cast(kvcache->data_ptr()) : nullptr, + .kvcache = static_cast(kvcache.data_ptr()), .ring_size = static_cast(ring_size), .eps = eps, }; @@ -431,22 +393,18 @@ struct FlashC2DecodeKernel { const auto loc_i32 = loc_dtype.is_type(); if (is_verify) { const auto block = static_cast(draft_len); - const auto k = select(pos_i32, loc_i32); + const auto k = select(pos_i32, loc_i32); LaunchKernel(dim3{block, num_tokens / block}, kBlockSize, device_.unwrap()) // .enable_pdl(kUsePDL)(k, params); - } else if (store) { - const auto k = select(pos_i32, loc_i32); - LaunchKernel(num_tokens, kBlockSize, device_.unwrap()) // - .enable_pdl(kUsePDL)(k, params); } else { - const auto k = select(pos_i32, loc_i32); + const auto k = select(pos_i32, loc_i32); LaunchKernel(num_tokens, kBlockSize, device_.unwrap()) // .enable_pdl(kUsePDL)(k, params); } } }; -// ensure that C++ wrapper can work +// The JIT module names and wrappers spell the layouts as bare enumerators. using enum deepseek_v4::KVLayout; } // namespace sglang diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh index 8b4772f2a30d..1ea85603db8d 100644 --- a/python/sglang/kernels/jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh @@ -829,7 +829,7 @@ struct FusedNormRopeKernel { } }; -// The JIT wrappers name the layouts as plain enumerators. +// The JIT module names and wrappers spell the layouts as bare enumerators. using enum deepseek_v4::KVLayout; } // namespace sglang diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/main_norm_rope.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/main_norm_rope.cuh index 3afa029a6c42..4655d6b71d9f 100644 --- a/python/sglang/kernels/jit/csrc/deepseek_v4/main_norm_rope.cuh +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/main_norm_rope.cuh @@ -917,7 +917,7 @@ struct FusedQIndexerRopeHadamardFp4QuantKernel { } }; -// The JIT wrappers name the layouts as plain enumerators. +// The JIT module names and wrappers spell the layouts as bare enumerators. using enum deepseek_v4::KVLayout; } // namespace sglang diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/store.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/store.cuh index 9d91f34e36c4..9968277a4966 100644 --- a/python/sglang/kernels/jit/csrc/deepseek_v4/store.cuh +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/store.cuh @@ -318,7 +318,7 @@ struct FusedStoreCacheIndexerKernel { } }; -// The JIT wrappers name the layouts as plain enumerators. +// The JIT module names and wrappers spell the layouts as bare enumerators. using enum deepseek_v4::KVLayout; } // namespace sglang diff --git a/python/sglang/kernels/ops/attention/dsv4/c2.py b/python/sglang/kernels/ops/attention/dsv4/c2.py index f3fdb5acebc1..65b9e13f168e 100644 --- a/python/sglang/kernels/ops/attention/dsv4/c2.py +++ b/python/sglang/kernels/ops/attention/dsv4/c2.py @@ -8,7 +8,7 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Final, Optional, Union +from typing import TYPE_CHECKING, Optional, Union import torch @@ -33,8 +33,6 @@ def _jit_c2_module( page_size: int, layout: KVLayout, ) -> Module: - # rope_dim / page_size / layout only shape the store half; the norm-only - # entry point ignores them. args = make_cpp_args( head_dim, rope_dim, @@ -47,76 +45,11 @@ def _jit_c2_module( *args, cuda_files=["deepseek_v4/c2.cuh"], cuda_wrappers=[ - ("decode", f"FlashC2DecodeKernel<{args}>::run_decode"), ("decode_fusion", f"FlashC2DecodeKernel<{args}>::run_decode_fusion"), ], ) -def c2_decode_norm( - kv_input: torch.Tensor, - kv_state: torch.Tensor, - norm_weight: torch.Tensor, - positions: torch.Tensor, - req: torch.Tensor, - raw_out_loc: torch.Tensor, - eps: float, - *, - ring_size: int, - out: Optional[torch.Tensor] = None, -) -> torch.Tensor: - """Pair-pool ``kv_input`` against ``kv_state`` and RMSNorm the result. - - :param kv_input: ``[num_tokens, 2 * head_dim]`` fp32, ``| kv | score |``. - :param kv_state: ``CompressStatePool``'s flat ``KVAndScore`` buffer, - ``[size, 2 * head_dim]`` fp32, same ``| kv | score |`` - layout. A request's pending pair lives at - ``req * ring_size + pos % ring_size``, so a completing row - reads what ``pos - 1`` left and a pending one writes its - own slot -- read and write never touch the same row. - :param ring_size: ``CompressStatePool.ring_size``, positions per request. - :param norm_weight: ``[head_dim]`` bf16 -- ``DeepseekV41Compressor.norm`` - holds its weight in the model dtype, and the multiply is - done in fp32 by promoting it, exactly as the module does. - :param positions: ``[num_tokens]`` int32 or int64. An odd position completes a group - with its even predecessor. - :param req: ``[num_tokens]`` int64 ``req_pool_idx``, the ``kv_state`` row - this token pairs through. - :param raw_out_loc: ``[num_tokens]`` int32 or int64, the token's FULL-pool slot - (the scheduler's ``out_cache_loc`` is int64). - ``0`` marks a padded graph row, which the kernel skips - entirely -- reading nothing and writing nothing, so no - spare pair-state row is needed. - :param eps: RMSNorm epsilon. - :param out: ``[num_tokens, head_dim]`` bf16 destination. Pass a persistent - buffer under CUDA graphs. - :return: ``out``, the pre-RoPE post-norm latent. - - .. note:: **Rows at an even position, and padded rows, are not written.** - They complete no group, so the kernel skips their output row rather than - paying for a store the caller discards. Anything already in ``out`` on - those rows survives the call. - """ - num_tokens, fused_dim = kv_input.shape - head_dim = fused_dim // 2 - if out is None: - out = kv_input.new_empty((num_tokens, head_dim), dtype=torch.bfloat16) - - # Norm only: the store half's arguments are irrelevant, fixed to share a build. - _jit_c2_module(head_dim, 64, 128, KVLayout.V4).decode( - kv_input, - kv_state, - out, - norm_weight, - positions, - req, - raw_out_loc, - float(eps), - int(ring_size), - ) - return out - - def c2_decode_or_verify_norm_rope_store( kv_input: torch.Tensor, kv_state: torch.Tensor, @@ -134,7 +67,7 @@ def c2_decode_or_verify_norm_rope_store( layout: Union[KVLayout, str] = KVLayout.V4, out: Optional[torch.Tensor] = None, ) -> torch.Tensor: - """``c2_decode_norm`` plus the whole main-KV write, in the same launch. + """Pair-pool ``kv_input`` against ``kv_state``, RMSNorm, and write the main KV slot. ``out`` contains the pre-RoPE latent for the index-K branch's ``wk`` projection. The cache store uses ``raw_out_loc // 2`` as its slot. @@ -171,9 +104,3 @@ def c2_decode_or_verify_norm_rope_store( draft_len, ) return out - - -# DO NOT try to modify the alias - -c2_decode_norm_rope_store: Final = c2_decode_or_verify_norm_rope_store -c2_verify_norm_rope_store: Final = c2_decode_or_verify_norm_rope_store diff --git a/python/sglang/kernels/ops/attention/dsv4/kv_layout.py b/python/sglang/kernels/ops/attention/dsv4/kv_layout.py index 5edd618a629d..bb51e4361136 100644 --- a/python/sglang/kernels/ops/attention/dsv4/kv_layout.py +++ b/python/sglang/kernels/ops/attention/dsv4/kv_layout.py @@ -57,7 +57,8 @@ def scale_offset(self, page_size: int) -> int: @property def cpp_name(self) -> str: - """The C++ enumerator, for JIT template arguments.""" + """The C++ enumerator, for JIT template arguments. Bare, because it is + also part of the JIT module name; the headers `using enum` it in.""" return self.name @classmethod diff --git a/python/sglang/kernels/ops/attention/dsv41_small_metadata.py b/python/sglang/kernels/ops/attention/dsv4/small_metadata.py similarity index 86% rename from python/sglang/kernels/ops/attention/dsv41_small_metadata.py rename to python/sglang/kernels/ops/attention/dsv4/small_metadata.py index d4783d2128d0..b15d958adc56 100644 --- a/python/sglang/kernels/ops/attention/dsv41_small_metadata.py +++ b/python/sglang/kernels/ops/attention/dsv4/small_metadata.py @@ -1,5 +1,7 @@ """Small-batch V4.1 page and compression metadata.""" +from typing import NamedTuple + import torch import triton import triton.language as tl @@ -108,7 +110,20 @@ def _low_ratio_metadata( tl.store(PAGE2 + row * PADDED + cols, -1, cols < PADDED) -def low_ratio_metadata(seq_lens, out_loc, topk): +class LowRatioMetadata(NamedTuple): + """Per-request slots and lengths of the ratio-1 and ratio-2 compressed caches.""" + + c1_out_loc: torch.Tensor + c1_seq_lens: torch.Tensor + c1_sparse_lens: torch.Tensor + c1_page_indices: torch.Tensor + c2_out_loc: torch.Tensor + c2_seq_lens: torch.Tensor + c2_sparse_lens: torch.Tensor + c2_page_indices: torch.Tensor + + +def low_ratio_metadata(seq_lens, out_loc, topk) -> LowRatioMetadata: assert seq_lens.numel() == out_loc.numel() rows = seq_lens.numel() kw = dict(device=seq_lens.device, dtype=torch.int32) @@ -126,4 +141,4 @@ def low_ratio_metadata(seq_lens, out_loc, topk): _low_ratio_metadata[(rows,)]( seq_lens, out_loc, *outputs, topk, padded, triton.next_power_of_2(padded) ) - return outputs + return LowRatioMetadata(*outputs) From c888f9ee763d5a8f71a7f6db7493ca6653e55d8e Mon Sep 17 00:00:00 2001 From: hnyls2002 Date: Wed, 16 Sep 2026 03:09:44 -0700 Subject: [PATCH 25/30] drop unused fp4 torch reference quantizer --- .../sgl_kernel/deepseek_v4/kv_layout.cuh | 2 +- .../srt/layers/attention/dsv4/torch_quant.py | 81 +------------------ 2 files changed, 4 insertions(+), 79 deletions(-) diff --git a/python/sglang/kernels/jit/include/sgl_kernel/deepseek_v4/kv_layout.cuh b/python/sglang/kernels/jit/include/sgl_kernel/deepseek_v4/kv_layout.cuh index cd34309d5be2..5b201ee8b6db 100644 --- a/python/sglang/kernels/jit/include/sgl_kernel/deepseek_v4/kv_layout.cuh +++ b/python/sglang/kernels/jit/include/sgl_kernel/deepseek_v4/kv_layout.cuh @@ -26,7 +26,7 @@ // // The reader requires the rows of a page to be contiguous and the page stride // to be a multiple of kPageAlign (its TMA row stride), which is what -// kv_page_bytes pads to. The pure-torch reference of the V4.1 quantizers is +// kv_page_bytes pads to. The pure-torch reference of the V41 (fp8) quantizer is // `sglang.srt.layers.attention.dsv4.torch_quant`. namespace sglang { diff --git a/python/sglang/srt/layers/attention/dsv4/torch_quant.py b/python/sglang/srt/layers/attention/dsv4/torch_quant.py index b8cb0e4dda3e..0d043b58bd2d 100644 --- a/python/sglang/srt/layers/attention/dsv4/torch_quant.py +++ b/python/sglang/srt/layers/attention/dsv4/torch_quant.py @@ -65,14 +65,11 @@ def fake_quant_compressed_kv(x: torch.Tensor) -> torch.Tensor: # --------------------------------------------------------------------------- -# Pure-torch references of the paged V4.1 KV cache formats read by the sparse -# decode kernel (528 B/token fp8 "V41", 288 B/token fp4 "V41_FP4"). They follow -# the kernel's own reference quantizer and are what the store / dequant kernels -# and the tests are checked against, byte for byte. +# Pure-torch reference of the paged V4.1 fp8 KV cache format read by the sparse +# decode kernel (528 B/token "V41"), byte for byte what the store / dequant +# kernels are checked against. # --------------------------------------------------------------------------- -_E2M1_MAGNITUDES = (0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0) - def cast_scale_inv_to_ue8m0(scale_inv: torch.Tensor) -> torch.Tensor: """``2 ** ceil(log2(max(scale_inv, 1e-4)))`` as fp32, computed on the IEEE @@ -85,29 +82,6 @@ def cast_scale_inv_to_ue8m0(scale_inv: torch.Tensor) -> torch.Tensor: return torch.where(torch.isfinite(scale_inv), scale, scale_inv) -def quantize_to_e2m1_codes(x: torch.Tensor) -> torch.Tensor: - """Round to the nearest e2m1 value with the semantics of - ``cvt.rn.satfinite.e2m1x2.f32`` (ties to even, saturating to +-6) and return - the 4-bit codes as uint8. The sign is kept for values that round to zero - (``-0.0`` and small negatives give the code ``0x8``); NaN maps to code 0.""" - x = x.float() - mags = torch.tensor(_E2M1_MAGNITUDES, dtype=torch.float32, device=x.device) - sign = torch.signbit(x).to(torch.uint8) << 3 - a = torch.nan_to_num(x.abs(), nan=0.0, posinf=6.0).clamp_max(6.0) - mids = (mags[:-1] + mags[1:]) / 2 - code = torch.bucketize(a, mids, right=True) - on_tie = (a.unsqueeze(-1) == mids).any(dim=-1) - tie_code = torch.bucketize(a, mids, right=False) - code = torch.where(on_tie, tie_code + (tie_code & 1), code) - return sign | code.to(torch.uint8) - - -def dequantize_e2m1_codes(codes: torch.Tensor) -> torch.Tensor: - mags = torch.tensor(_E2M1_MAGNITUDES, dtype=torch.float32, device=codes.device) - val = mags[(codes & 7).long()] - return torch.where((codes & 8) != 0, -val, val) - - def quantize_k_cache_v41( k: torch.Tensor, page_bytes: Optional[int] = None ) -> torch.Tensor: @@ -143,52 +117,3 @@ def dequantize_k_cache_v41(pages: torch.Tensor, page_size: int) -> torch.Tensor: return (values.view(num_pages, page_size, 16, 32) * scale_bf16.unsqueeze(-1)).view( num_pages, page_size, 512 ) - - -def quantize_k_cache_v41_fp4( - k: torch.Tensor, page_bytes: Optional[int] = None -) -> torch.Tensor: - """``k`` ``[num_pages, page_size, 512]`` -> uint8 ``[num_pages, page_bytes]`` - pages of the V41_FP4 layout: 256 B of e2m1 codes per token (even index in - the low nibble), then 32 e4m3 scales per token (one per 16 values), - ``scale = e4m3(clamp(amax / 6, 2**-9, 448))``. A NaN element poisons its - tile: NaN scale, zero codes.""" - num_pages, page_size, d = k.shape - assert d == 512 - x = k.float().view(num_pages, page_size, 32, 16) - amax = torch.nan_to_num(x.abs(), nan=float("inf")).amax(dim=-1) - scale = torch.clamp(amax / 6.0, 2.0**-9, 448.0).to(torch.float8_e4m3fn) - scale = torch.where(torch.isinf(amax), torch.full_like(scale, float("nan")), scale) - codes = quantize_to_e2m1_codes(x / scale.float().unsqueeze(-1)) - codes = codes.view(num_pages, page_size, 512) - packed = codes[..., 0::2] | (codes[..., 1::2] << 4) - raw = page_size * 288 - if page_bytes is None: - page_bytes = -(-raw // 256) * 256 - assert page_bytes >= raw - out = torch.zeros((num_pages, page_bytes), dtype=torch.uint8, device=k.device) - out[:, : page_size * 256] = packed.reshape(num_pages, page_size * 256) - out[:, page_size * 256 : raw] = scale.view(torch.uint8).reshape( - num_pages, page_size * 32 - ) - return out - - -def dequantize_k_cache_v41_fp4(pages: torch.Tensor, page_size: int) -> torch.Tensor: - """Inverse of :func:`quantize_k_cache_v41_fp4`: ``[num_pages, page_size, 512]`` - bf16. ``e2m1 * e4m3`` has at most 2 + 4 significant bits, so the product is - exact in bf16, as in the kernel.""" - num_pages = pages.shape[0] - pages = pages.view(torch.uint8) - data = pages[:, : page_size * 256].reshape(num_pages, page_size, 256) - scale = pages[:, page_size * 256 : page_size * 288].reshape( - num_pages, page_size, 32 - ) - codes = torch.empty( - (num_pages, page_size, 512), dtype=torch.uint8, device=pages.device - ) - codes[..., 0::2] = data & 0xF - codes[..., 1::2] = data >> 4 - values = dequantize_e2m1_codes(codes).view(num_pages, page_size, 32, 16) - out = values * scale.view(torch.float8_e4m3fn).float().unsqueeze(-1) - return out.view(num_pages, page_size, 512).to(torch.bfloat16) From f33acae343ac8e99c823dad45b9bda9f4a216f2a Mon Sep 17 00:00:00 2001 From: BBuf <1182563586@qq.com> Date: Wed, 16 Sep 2026 22:27:27 +0800 Subject: [PATCH 26/30] c2: drop the dead duplicate freqs_cis load flash_c2_decode_kernel loads the RoPE frequencies twice for the same thread and the same position: once right after the norm weight, and again inside the `tx >= kNopeThreads` branch that consumes them. `freq` is never read between the two, so the first load is dead and the second re-issues the same 8 bytes per lane. Keep the early load -- it is issued before the softmax and the RMSNorm reduction, so its latency overlaps that work -- and remove the reload inside the branch. Verified bit-exact on GB300 (sm_103) against the pre-change kernel: identical sha256 of `out`, the paged cache and the pair-state ring for V4, V41 and V41_FP4 decode plus a draft_len=4 verify batch, with the JIT cache cleared between runs. A deliberate eps perturbation was used as a negative control to confirm the harness recompiles and detects kernel changes. Co-Authored-By: Claude Opus 5 (1M context) --- python/sglang/kernels/jit/csrc/deepseek_v4/c2.cuh | 1 - 1 file changed, 1 deletion(-) diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/c2.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/c2.cuh index 6f34c6be3eb1..0de8e45d6ea0 100644 --- a/python/sglang/kernels/jit/csrc/deepseek_v4/c2.cuh +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/c2.cuh @@ -199,7 +199,6 @@ __global__ __launch_bounds__(kHeadDim / kC2VecSize) void flash_c2_decode_kernel( if (tx >= kNopeThreads) { // Match rope_tail()'s bf16 rounding before fake quantization. // Only odd positions reach here; the latent represents `pos - 1`. - freq.load(params.freqs_cis + (pos - 1) * kRopeDim, tx - kNopeThreads); #pragma unroll for (uint32_t i = 0; i < kVecSize / 2; ++i) { const auto x_real = staged[i * 2 + 0]; From 922b0d547394182da1cfbd30b40951ff8a05d3c7 Mon Sep 17 00:00:00 2001 From: BBuf <1182563586@qq.com> Date: Wed, 16 Sep 2026 22:28:53 +0800 Subject: [PATCH 27/30] mhc: stop allocating a throwaway sqrsum for the residual projection hc_mix_stats_sinkhorn_deepgemm ran the two compensated projections as tf32_hc_prenorm_gemm(x_flat, high, mix_hi, sq, _NUM_SPLITS) tf32_hc_prenorm_gemm(x_flat, low, mix_lo, unused_sq, _NUM_SPLITS) where `unused_sq` exists only to be discarded -- a (_NUM_SPLITS, m) fp32 buffer, 4 MiB at m = 65536, allocated on every prefill call. The row sum of squares is a function of x_flat alone, so the second call recomputes exactly what the first already wrote; the two calls are ordered on the same stream, so it can simply write back into `sq`. Verified on GB300 (sm_103) against the installed DeepGEMM: the sqrsum output is independent of `fn`, is fully written rather than accumulated, and the patched two-call sequence leaves `sq` bitwise identical to the original. Co-Authored-By: Claude Opus 5 (1M context) --- .../sglang/kernels/ops/layernorm/hc_mix_stats_deepgemm.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/python/sglang/kernels/ops/layernorm/hc_mix_stats_deepgemm.py b/python/sglang/kernels/ops/layernorm/hc_mix_stats_deepgemm.py index 6cb245e65469..15b9d41ff2cb 100644 --- a/python/sglang/kernels/ops/layernorm/hc_mix_stats_deepgemm.py +++ b/python/sglang/kernels/ops/layernorm/hc_mix_stats_deepgemm.py @@ -42,9 +42,12 @@ def hc_mix_stats_sinkhorn_deepgemm( mix_hi = torch.empty((_NUM_SPLITS, m, 24), dtype=torch.float32, device=dev) mix_lo = torch.empty_like(mix_hi) sq = torch.empty((_NUM_SPLITS, m), dtype=torch.float32, device=dev) - unused_sq = torch.empty_like(sq) tf32_hc_prenorm_gemm(x_flat, high, mix_hi, sq, _NUM_SPLITS) - tf32_hc_prenorm_gemm(x_flat, low, mix_lo, unused_sq, _NUM_SPLITS) + # The row sum of squares depends only on x_flat, so the second projection + # recomputes exactly the same values; let it overwrite sq instead of + # allocating a (_NUM_SPLITS, m) fp32 scratch that is thrown away + # (4 MiB at m = 65536). + tf32_hc_prenorm_gemm(x_flat, low, mix_lo, sq, _NUM_SPLITS) _hc_mix_reduce_sinkhorn_kernel[(m,)]( mix_hi, sq, From 8af4e434fe6703ef2e87d87370773bbe68ad24f0 Mon Sep 17 00:00:00 2001 From: hnyls2002 Date: Wed, 16 Sep 2026 12:38:55 -0700 Subject: [PATCH 28/30] merge c1/c2 wrappers; split small metadata into its homes; move torch_quant next to the kernels; align compress kernel names --- .../kernels/jit/csrc/deepseek_v4/c1.cuh | 10 +- .../kernels/jit/csrc/deepseek_v4/c2.cuh | 8 +- .../sgl_kernel/deepseek_v4/kv_layout.cuh | 2 +- .../sglang/kernels/ops/attention/dsv4/c2.py | 106 ------------- .../dsv4/{c1.py => low_ratio_compress.py} | 104 ++++++++++++- .../ops/attention/dsv4/metadata_kernel.py | 67 +++++++- .../ops/attention/dsv4/small_metadata.py | 144 ------------------ .../ops}/attention/dsv4/torch_quant.py | 14 +- .../attention/dsv4_attn_metadata_kernels.py | 68 +++++++++ .../test_deepseek_v4_compress_plan_bounds.py} | 2 +- 10 files changed, 249 insertions(+), 276 deletions(-) delete mode 100644 python/sglang/kernels/ops/attention/dsv4/c2.py rename python/sglang/kernels/ops/attention/dsv4/{c1.py => low_ratio_compress.py} (50%) delete mode 100644 python/sglang/kernels/ops/attention/dsv4/small_metadata.py rename python/sglang/{srt/layers => kernels/ops}/attention/dsv4/torch_quant.py (91%) rename test/registered/{kernels/ops/attention/test_deepseek_v4_compress_plan_draft_pad.py => kernel/attention/test_deepseek_v4_compress_plan_bounds.py} (98%) diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/c1.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/c1.cuh index ea76c8630448..4a14db009650 100644 --- a/python/sglang/kernels/jit/csrc/deepseek_v4/c1.cuh +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/c1.cuh @@ -25,7 +25,7 @@ namespace sglang { /// kernel's input is the GEMM output and its RoPE position is `positions`, not /// `positions - 1`. `kv_output` is the pre-RoPE latent, for the index-K /// branch's `wk` projection. -struct C1Params { +struct Compress1DecodeParams { const bf16_t* __restrict__ kv_input; // [num_tokens, kHeadDim] bf16 bf16_t* __restrict__ kv_output; // [num_tokens, kHeadDim] bf16, pre-RoPE const bf16_t* __restrict__ norm_weight; // [kHeadDim] bf16 @@ -61,8 +61,8 @@ template < typename LocT, deepseek_v4::KVLayout kLayout, bool kUsePDL> -__global__ -__launch_bounds__(kHeadDim / kC1VecSize) void flash_c1_decode_kernel(const __grid_constant__ C1Params params) { +__global__ __launch_bounds__(kHeadDim / kC1VecSize) void flash_c1_decode_kernel( + const __grid_constant__ Compress1DecodeParams params) { using namespace device; using deepseek_v4::KVLayout; using deepseek_v4::fp8::cast_to_ue8m0; @@ -237,7 +237,7 @@ __launch_bounds__(kHeadDim / kC1VecSize) void flash_c1_decode_kernel(const __gri /// \brief Host side of `flash_c1_decode_kernel`. template -struct FlashC1DecodeKernel { +struct FlashCompress1Kernel { static constexpr int32_t kPageBits = std::bit_width(kPageSize) - 1; static constexpr int64_t kPageBytes = deepseek_v4::kv_page_bytes(kPageSize); static constexpr uint32_t kBlockSize = kHeadDim / kC1VecSize; @@ -300,7 +300,7 @@ struct FlashC1DecodeKernel { const auto num_tokens = static_cast(N.unwrap()); if (num_tokens == 0) return; - const auto params = C1Params{ + const auto params = Compress1DecodeParams{ .kv_input = static_cast(kv_input.data_ptr()), .kv_output = static_cast(kv_output.data_ptr()), .norm_weight = static_cast(norm_weight.data_ptr()), diff --git a/python/sglang/kernels/jit/csrc/deepseek_v4/c2.cuh b/python/sglang/kernels/jit/csrc/deepseek_v4/c2.cuh index 0de8e45d6ea0..d7544574ff1d 100644 --- a/python/sglang/kernels/jit/csrc/deepseek_v4/c2.cuh +++ b/python/sglang/kernels/jit/csrc/deepseek_v4/c2.cuh @@ -22,7 +22,7 @@ namespace sglang { /// /// `kv_input` and `kv_state` rows are `2 * kHeadDim` floats, kv then score. /// `kv_output` is the pre-RoPE latent, for the index-K branch's `wk`. -struct C2Params { +struct Compress2DecodeParams { const float* __restrict__ kv_input; // [num_tokens, 2 * kHeadDim] fp32 /// `CompressStatePool`'s flat `KVAndScore` buffer, `[size, 2 * kHeadDim]` /// fp32 with kv in the low half and score in the high half. A request's @@ -73,7 +73,7 @@ template < typename LocT, deepseek_v4::KVLayout kLayout, bool kUsePDL> -__global__ __launch_bounds__(kHeadDim / kC2VecSize) void flash_c2_decode_kernel(const C2Params params) { +__global__ __launch_bounds__(kHeadDim / kC2VecSize) void flash_c2_decode_kernel(const Compress2DecodeParams params) { using namespace device; using deepseek_v4::KVLayout; using deepseek_v4::fp8::cast_to_ue8m0; @@ -275,7 +275,7 @@ __global__ __launch_bounds__(kHeadDim / kC2VecSize) void flash_c2_decode_kernel( } template -struct FlashC2DecodeKernel { +struct FlashCompress2Kernel { static constexpr uint32_t kBlockSize = kHeadDim / kC2VecSize; static constexpr int32_t kPageBits = std::bit_width(kPageSize) - 1; static constexpr int64_t kPageBytes = deepseek_v4::kv_page_bytes(kPageSize); @@ -374,7 +374,7 @@ struct FlashC2DecodeKernel { CHECK_HOST(!is_verify || num_tokens % draft_len == 0); CHECK_HOST(!is_verify || ring_size > draft_len) << "the pair-state ring (" << ring_size << ") must be wider than the draft length (" << draft_len << ")"; - const auto params = C2Params{ + const auto params = Compress2DecodeParams{ .kv_input = static_cast(kv_input.data_ptr()), .kv_state = static_cast(kv_state.data_ptr()), .kv_output = static_cast(kv_output.data_ptr()), diff --git a/python/sglang/kernels/jit/include/sgl_kernel/deepseek_v4/kv_layout.cuh b/python/sglang/kernels/jit/include/sgl_kernel/deepseek_v4/kv_layout.cuh index 5b201ee8b6db..6eabad45f0f7 100644 --- a/python/sglang/kernels/jit/include/sgl_kernel/deepseek_v4/kv_layout.cuh +++ b/python/sglang/kernels/jit/include/sgl_kernel/deepseek_v4/kv_layout.cuh @@ -27,7 +27,7 @@ // The reader requires the rows of a page to be contiguous and the page stride // to be a multiple of kPageAlign (its TMA row stride), which is what // kv_page_bytes pads to. The pure-torch reference of the V41 (fp8) quantizer is -// `sglang.srt.layers.attention.dsv4.torch_quant`. +// `sglang.kernels.ops.attention.dsv4.torch_quant`. namespace sglang { diff --git a/python/sglang/kernels/ops/attention/dsv4/c2.py b/python/sglang/kernels/ops/attention/dsv4/c2.py deleted file mode 100644 index 65b9e13f168e..000000000000 --- a/python/sglang/kernels/ops/attention/dsv4/c2.py +++ /dev/null @@ -1,106 +0,0 @@ -"""Fused ratio-2 decode pair-pooling, RMSNorm and optional main-KV write. - -Closed-form softmax and FMA contraction can differ from torch by fp32 ulps; -bf16 rounding boundaries can preserve those differences. The pooling tests -use a tolerance, while state updates and stores from a given latent are bitwise. -Positions and state-ring indices describe the per-request decode schedule. -""" - -from __future__ import annotations - -from typing import TYPE_CHECKING, Optional, Union - -import torch - -from sglang.kernels.jit.utils import ( - cache_once, - is_arch_support_pdl, - load_jit, - make_cpp_args, -) - -from .kv_layout import KVLayout -from .utils import make_name - -if TYPE_CHECKING: - from tvm_ffi.module import Module - - -@cache_once -def _jit_c2_module( - head_dim: int, - rope_dim: int, - page_size: int, - layout: KVLayout, -) -> Module: - args = make_cpp_args( - head_dim, - rope_dim, - page_size, - layout.cpp_name, - is_arch_support_pdl(), - ) - return load_jit( - make_name("c2"), - *args, - cuda_files=["deepseek_v4/c2.cuh"], - cuda_wrappers=[ - ("decode_fusion", f"FlashC2DecodeKernel<{args}>::run_decode_fusion"), - ], - ) - - -def c2_decode_or_verify_norm_rope_store( - kv_input: torch.Tensor, - kv_state: torch.Tensor, - norm_weight: torch.Tensor, - positions: torch.Tensor, - req: torch.Tensor, - raw_out_loc: torch.Tensor, - eps: float, - freqs_cis: torch.Tensor, - k_cache: torch.Tensor, - *, - page_size: int, - ring_size: int, - draft_len: int = 1, - layout: Union[KVLayout, str] = KVLayout.V4, - out: Optional[torch.Tensor] = None, -) -> torch.Tensor: - """Pair-pool ``kv_input`` against ``kv_state``, RMSNorm, and write the main KV slot. - - ``out`` contains the pre-RoPE latent for the index-K branch's ``wk`` projection. - The cache store uses ``raw_out_loc // 2`` as its slot. - - :param freqs_cis: ``[max_pos, rope_dim]`` fp32, real/imag interleaved -- - ``torch.view_as_real(freqs).flatten(-2)``. Indexed - in-kernel at ``positions - 1``, the position the latent - stands for, so there is no gather launch. - :param k_cache: the compressed KV pool buffer for this layer. - :param page_size: slots per page of that pool (``page_size // ratio``). - :param layout: the pool's :class:`KVLayout`. The fp8 layouts (``V4``, - ``V41``) store the fp4 fake-quantized value; ``V41_FP4`` - stores the e2m1 codes themselves, rounding once. - """ - num_tokens, fused_dim = kv_input.shape - head_dim = fused_dim // 2 - if out is None: - out = kv_input.new_empty((num_tokens, head_dim), dtype=torch.bfloat16) - - layout = KVLayout.parse(layout) - module = _jit_c2_module(head_dim, freqs_cis.shape[-1], page_size, layout) - module.decode_fusion( - kv_input, - kv_state, - out, - norm_weight, - positions, - req, - raw_out_loc, - eps, - freqs_cis, - k_cache, - ring_size, - draft_len, - ) - return out diff --git a/python/sglang/kernels/ops/attention/dsv4/c1.py b/python/sglang/kernels/ops/attention/dsv4/low_ratio_compress.py similarity index 50% rename from python/sglang/kernels/ops/attention/dsv4/c1.py rename to python/sglang/kernels/ops/attention/dsv4/low_ratio_compress.py index 2feb12318468..ebcf8e0d4d94 100644 --- a/python/sglang/kernels/ops/attention/dsv4/c1.py +++ b/python/sglang/kernels/ops/attention/dsv4/low_ratio_compress.py @@ -1,11 +1,18 @@ -"""Fused ratio-1 decode RMSNorm, RoPE, fp4 fake-quant and FlashMLA cache write. +"""Fused ratio-1 and ratio-2 decode compressors: RMSNorm, RoPE and the FlashMLA +cache write in one launch. -The input is bf16 with no pooling. RoPE uses the token's own position, and the -compressed slot equals the FULL slot; out_loc == 0 marks graph padding. -The pre-RoPE latent is also returned for the index-key projection. +Ratio 1 takes the bf16 ``wkv`` projection as is: RoPE uses the token's own +position and the compressed slot equals the FULL slot. Ratio 2 pair-pools the +token against the pending partner in the state ring first; its closed-form +softmax and FMA contraction can differ from torch by fp32 ulps, so pooling is +compared with a tolerance while stores from a given latent are bitwise. +``out_loc == 0`` marks a padded graph row on both paths, and both return the +pre-RoPE latent for the index-key projection. """ -from typing import Optional, Union +from __future__ import annotations + +from typing import TYPE_CHECKING, Optional, Union import torch @@ -19,6 +26,9 @@ from .kv_layout import KVLayout from .utils import make_name +if TYPE_CHECKING: + from tvm_ffi.module import Module + @cache_once def _jit_c1_module(head_dim: int, rope_dim: int, page_size: int, layout: KVLayout): @@ -30,11 +40,11 @@ def _jit_c1_module(head_dim: int, rope_dim: int, page_size: int, layout: KVLayou is_arch_support_pdl(), ) return load_jit( - make_name("c1"), + make_name("c1_decode"), *args, cuda_files=["deepseek_v4/c1.cuh"], cuda_wrappers=[ - ("decode_fusion", f"FlashC1DecodeKernel<{args}>::run_decode_fusion"), + ("decode_fusion", f"FlashCompress1Kernel<{args}>::run_decode_fusion"), ], ) @@ -101,3 +111,83 @@ def c1_decode_norm_rope_store( float(eps), ) return out + + +@cache_once +def _jit_c2_module( + head_dim: int, + rope_dim: int, + page_size: int, + layout: KVLayout, +) -> Module: + args = make_cpp_args( + head_dim, + rope_dim, + page_size, + layout.cpp_name, + is_arch_support_pdl(), + ) + return load_jit( + make_name("c2_decode"), + *args, + cuda_files=["deepseek_v4/c2.cuh"], + cuda_wrappers=[ + ("decode_fusion", f"FlashCompress2Kernel<{args}>::run_decode_fusion"), + ], + ) + + +def c2_decode_norm_rope_store( + kv_input: torch.Tensor, + kv_state: torch.Tensor, + norm_weight: torch.Tensor, + positions: torch.Tensor, + req: torch.Tensor, + raw_out_loc: torch.Tensor, + eps: float, + freqs_cis: torch.Tensor, + k_cache: torch.Tensor, + *, + page_size: int, + ring_size: int, + draft_len: int = 1, + layout: Union[KVLayout, str] = KVLayout.V4, + out: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """Pair-pool ``kv_input`` against ``kv_state``, RMSNorm, and write the main KV slot. + + ``out`` contains the pre-RoPE latent for the index-K branch's ``wk`` projection. + The cache store uses ``raw_out_loc // 2`` as its slot. + + :param freqs_cis: ``[max_pos, rope_dim]`` fp32, real/imag interleaved -- + ``torch.view_as_real(freqs).flatten(-2)``. Indexed + in-kernel at ``positions - 1``, the position the latent + stands for, so there is no gather launch. + :param k_cache: the compressed KV pool buffer for this layer. + :param page_size: slots per page of that pool (``page_size // ratio``). + :param layout: the pool's :class:`KVLayout`. The fp8 layouts (``V4``, + ``V41``) store the fp4 fake-quantized value; ``V41_FP4`` + stores the e2m1 codes themselves, rounding once. + """ + num_tokens, fused_dim = kv_input.shape + head_dim = fused_dim // 2 + if out is None: + out = kv_input.new_empty((num_tokens, head_dim), dtype=torch.bfloat16) + + layout = KVLayout.parse(layout) + module = _jit_c2_module(head_dim, freqs_cis.shape[-1], page_size, layout) + module.decode_fusion( + kv_input, + kv_state, + out, + norm_weight, + positions, + req, + raw_out_loc, + eps, + freqs_cis, + k_cache, + ring_size, + draft_len, + ) + return out diff --git a/python/sglang/kernels/ops/attention/dsv4/metadata_kernel.py b/python/sglang/kernels/ops/attention/dsv4/metadata_kernel.py index 5647f08cac14..8a647544c04a 100644 --- a/python/sglang/kernels/ops/attention/dsv4/metadata_kernel.py +++ b/python/sglang/kernels/ops/attention/dsv4/metadata_kernel.py @@ -1,4 +1,4 @@ -from typing import Optional, Tuple +from typing import NamedTuple, Optional, Tuple import torch import triton @@ -275,3 +275,68 @@ def init_compression_metadata( page_size, compute_page_indices, ) + + +@triton.jit +def _low_ratio_metadata( + LENS, + LOC, + OUT1, + LEN1, + SPARSE1, + PAGE1, + OUT2, + LEN2, + SPARSE2, + PAGE2, + TOPK: tl.constexpr, + PADDED: tl.constexpr, + BLOCK: tl.constexpr, +): + row = tl.program_id(0) + length = tl.load(LENS + row).to(tl.int32) + loc = tl.load(LOC + row).to(tl.int64) + len1, len2 = tl.maximum(length, 1), tl.maximum(length >> 1, 1) + tl.store(OUT1 + row, loc) + tl.store(OUT2 + row, tl.where((length & 1) == 0, loc >> 1, -1)) + tl.store(LEN1 + row, len1) + tl.store(LEN2 + row, len2) + tl.store(SPARSE1 + row, tl.minimum(len1, TOPK)) + tl.store(SPARSE2 + row, tl.minimum(len2, TOPK)) + cols = tl.arange(0, BLOCK) + tl.store(PAGE1 + row * PADDED + cols, -1, cols < PADDED) + tl.store(PAGE2 + row * PADDED + cols, -1, cols < PADDED) + + +class LowRatioMetadata(NamedTuple): + """Per-request slots and lengths of the ratio-1 and ratio-2 compressed caches.""" + + c1_out_loc: torch.Tensor + c1_seq_lens: torch.Tensor + c1_sparse_lens: torch.Tensor + c1_page_indices: torch.Tensor + c2_out_loc: torch.Tensor + c2_seq_lens: torch.Tensor + c2_sparse_lens: torch.Tensor + c2_page_indices: torch.Tensor + + +def build_low_ratio_metadata(seq_lens, out_loc, topk) -> LowRatioMetadata: + assert seq_lens.numel() == out_loc.numel() + rows = seq_lens.numel() + kw = dict(device=seq_lens.device, dtype=torch.int32) + padded = triton.cdiv(topk, 64) * 64 + outputs = [] + for _ in range(2): + outputs.extend( + [ + torch.empty(rows, device=out_loc.device, dtype=torch.int64), + torch.empty(rows, **kw), + torch.empty(rows, **kw), + torch.empty((rows, padded), **kw), + ] + ) + _low_ratio_metadata[(rows,)]( + seq_lens, out_loc, *outputs, topk, padded, triton.next_power_of_2(padded) + ) + return LowRatioMetadata(*outputs) diff --git a/python/sglang/kernels/ops/attention/dsv4/small_metadata.py b/python/sglang/kernels/ops/attention/dsv4/small_metadata.py deleted file mode 100644 index b15d958adc56..000000000000 --- a/python/sglang/kernels/ops/attention/dsv4/small_metadata.py +++ /dev/null @@ -1,144 +0,0 @@ -"""Small-batch V4.1 page and compression metadata.""" - -from typing import NamedTuple - -import torch -import triton -import triton.language as tl - -from sglang.kernels.ops.attention.dsv4_attn_metadata_kernels import ( - PageTablePositionsResult, -) - - -@triton.jit -def _small_page_table( - REQ_TO_TOKEN, - REQS, - LENS, - OUT_LENS, - POS, - PAGES, - SWA, - STRIDE: tl.constexpr, - NUM_PAGES: tl.constexpr, - PAGE_SIZE: tl.constexpr, - WINDOW: tl.constexpr, - BLOCK: tl.constexpr, -): - row, tile = tl.program_id(0), tl.program_id(1) - if tile == 0: - length = tl.load(LENS + row).to(tl.int32) - tl.store(OUT_LENS + row, length) - tl.store(POS + row, length - 1) - tl.store(SWA + row, tl.minimum(length, WINDOW)) - req = tl.load(REQS + row).to(tl.int64) - p = tile * BLOCK + tl.arange(0, BLOCK) - slot = tl.load( - REQ_TO_TOKEN + req * STRIDE + p.to(tl.int64) * PAGE_SIZE, - mask=p < NUM_PAGES, - other=0, - ).to(tl.int32) - tl.store(PAGES + row * NUM_PAGES + p, slot // PAGE_SIZE, mask=p < NUM_PAGES) - - -def page_table_positions_small( - *, - req_to_token, - req_pool_indices_repeated, - seq_lens_casual, - max_seq_len, - page_size, - swa_window, -): - assert page_size > 0 and page_size & (page_size - 1) == 0 - rows = seq_lens_casual.numel() - pages = triton.cdiv(max_seq_len, page_size) - kw = dict(device=seq_lens_casual.device, dtype=torch.int32) - lengths, positions, swa = [torch.empty(rows, **kw) for _ in range(3)] - table = torch.empty((rows, pages), **kw) - _small_page_table[(rows, triton.cdiv(pages, 256))]( - req_to_token, - req_pool_indices_repeated, - seq_lens_casual, - lengths, - positions, - table, - swa, - req_to_token.stride(0), - pages, - page_size, - swa_window, - 256, - ) - return PageTablePositionsResult( - seq_lens_casual=lengths, - positions_casual=positions, - page_table=table, - swa_topk_lengths=swa, - ) - - -@triton.jit -def _low_ratio_metadata( - LENS, - LOC, - OUT1, - LEN1, - SPARSE1, - PAGE1, - OUT2, - LEN2, - SPARSE2, - PAGE2, - TOPK: tl.constexpr, - PADDED: tl.constexpr, - BLOCK: tl.constexpr, -): - row = tl.program_id(0) - length = tl.load(LENS + row).to(tl.int32) - loc = tl.load(LOC + row).to(tl.int64) - len1, len2 = tl.maximum(length, 1), tl.maximum(length >> 1, 1) - tl.store(OUT1 + row, loc) - tl.store(OUT2 + row, tl.where((length & 1) == 0, loc >> 1, -1)) - tl.store(LEN1 + row, len1) - tl.store(LEN2 + row, len2) - tl.store(SPARSE1 + row, tl.minimum(len1, TOPK)) - tl.store(SPARSE2 + row, tl.minimum(len2, TOPK)) - cols = tl.arange(0, BLOCK) - tl.store(PAGE1 + row * PADDED + cols, -1, cols < PADDED) - tl.store(PAGE2 + row * PADDED + cols, -1, cols < PADDED) - - -class LowRatioMetadata(NamedTuple): - """Per-request slots and lengths of the ratio-1 and ratio-2 compressed caches.""" - - c1_out_loc: torch.Tensor - c1_seq_lens: torch.Tensor - c1_sparse_lens: torch.Tensor - c1_page_indices: torch.Tensor - c2_out_loc: torch.Tensor - c2_seq_lens: torch.Tensor - c2_sparse_lens: torch.Tensor - c2_page_indices: torch.Tensor - - -def low_ratio_metadata(seq_lens, out_loc, topk) -> LowRatioMetadata: - assert seq_lens.numel() == out_loc.numel() - rows = seq_lens.numel() - kw = dict(device=seq_lens.device, dtype=torch.int32) - padded = triton.cdiv(topk, 64) * 64 - outputs = [] - for _ in range(2): - outputs.extend( - [ - torch.empty(rows, device=out_loc.device, dtype=torch.int64), - torch.empty(rows, **kw), - torch.empty(rows, **kw), - torch.empty((rows, padded), **kw), - ] - ) - _low_ratio_metadata[(rows,)]( - seq_lens, out_loc, *outputs, topk, padded, triton.next_power_of_2(padded) - ) - return LowRatioMetadata(*outputs) diff --git a/python/sglang/srt/layers/attention/dsv4/torch_quant.py b/python/sglang/kernels/ops/attention/dsv4/torch_quant.py similarity index 91% rename from python/sglang/srt/layers/attention/dsv4/torch_quant.py rename to python/sglang/kernels/ops/attention/dsv4/torch_quant.py index 0d043b58bd2d..3f3712c990ba 100644 --- a/python/sglang/srt/layers/attention/dsv4/torch_quant.py +++ b/python/sglang/kernels/ops/attention/dsv4/torch_quant.py @@ -71,15 +71,15 @@ def fake_quant_compressed_kv(x: torch.Tensor) -> torch.Tensor: # --------------------------------------------------------------------------- -def cast_scale_inv_to_ue8m0(scale_inv: torch.Tensor) -> torch.Tensor: - """``2 ** ceil(log2(max(scale_inv, 1e-4)))`` as fp32, computed on the IEEE - bits so that it is exact at (and just above) powers of two.""" - scale_inv = scale_inv.float() - scale = ceil_pow2(torch.clamp_min(scale_inv, 1e-4)) +def ceil_pow2_scale(x: torch.Tensor) -> torch.Tensor: + """``2 ** ceil(log2(max(x, 1e-4)))`` as fp32, computed on the IEEE bits so + that it is exact at (and just above) powers of two.""" + x = x.float() + scale = ceil_pow2(torch.clamp_min(x, 1e-4)) # ceil_pow2 works on the bits of a finite value; a NaN or inf amax passes # through (both become the ue8m0 NaN byte, but only the NaN one turns the # whole tile's payload into NaN). - return torch.where(torch.isfinite(scale_inv), scale, scale_inv) + return torch.where(torch.isfinite(x), scale, x) def quantize_k_cache_v41( @@ -91,7 +91,7 @@ def quantize_k_cache_v41( num_pages, page_size, d = k.shape assert d == 512 x = k.float().view(num_pages, page_size, 16, 32) - scale = cast_scale_inv_to_ue8m0(x.abs().amax(dim=-1) / 448.0) + scale = ceil_pow2_scale(x.abs().amax(dim=-1) / 448.0) data = (x / scale.unsqueeze(-1)).to(torch.float8_e4m3fn).view(torch.uint8) scale_u8 = scale.to(torch.float8_e8m0fnu).view(torch.uint8) raw = page_size * 528 diff --git a/python/sglang/kernels/ops/attention/dsv4_attn_metadata_kernels.py b/python/sglang/kernels/ops/attention/dsv4_attn_metadata_kernels.py index bbb28952c9e1..d93bbee70a15 100644 --- a/python/sglang/kernels/ops/attention/dsv4_attn_metadata_kernels.py +++ b/python/sglang/kernels/ops/attention/dsv4_attn_metadata_kernels.py @@ -524,3 +524,71 @@ def build_causal_swa_page_indices_triton( BLOCK_K=BLOCK_K, ) return out + + +@triton.jit +def _small_page_table( + REQ_TO_TOKEN, + REQS, + LENS, + OUT_LENS, + POS, + PAGES, + SWA, + STRIDE: tl.constexpr, + NUM_PAGES: tl.constexpr, + PAGE_SIZE: tl.constexpr, + WINDOW: tl.constexpr, + BLOCK: tl.constexpr, +): + row, tile = tl.program_id(0), tl.program_id(1) + if tile == 0: + length = tl.load(LENS + row).to(tl.int32) + tl.store(OUT_LENS + row, length) + tl.store(POS + row, length - 1) + tl.store(SWA + row, tl.minimum(length, WINDOW)) + req = tl.load(REQS + row).to(tl.int64) + p = tile * BLOCK + tl.arange(0, BLOCK) + slot = tl.load( + REQ_TO_TOKEN + req * STRIDE + p.to(tl.int64) * PAGE_SIZE, + mask=p < NUM_PAGES, + other=0, + ).to(tl.int32) + tl.store(PAGES + row * NUM_PAGES + p, slot // PAGE_SIZE, mask=p < NUM_PAGES) + + +def build_page_table_positions_small( + *, + req_to_token, + req_pool_indices_repeated, + seq_lens_casual, + max_seq_len, + page_size, + swa_window, +): + assert page_size > 0 and page_size & (page_size - 1) == 0 + rows = seq_lens_casual.numel() + pages = triton.cdiv(max_seq_len, page_size) + kw = dict(device=seq_lens_casual.device, dtype=torch.int32) + lengths, positions, swa = [torch.empty(rows, **kw) for _ in range(3)] + table = torch.empty((rows, pages), **kw) + _small_page_table[(rows, triton.cdiv(pages, 256))]( + req_to_token, + req_pool_indices_repeated, + seq_lens_casual, + lengths, + positions, + table, + swa, + req_to_token.stride(0), + pages, + page_size, + swa_window, + 256, + ) + return PageTablePositionsResult( + seq_lens_casual=lengths, + positions_casual=positions, + page_table=table, + swa_topk_lengths=swa, + ) diff --git a/test/registered/kernels/ops/attention/test_deepseek_v4_compress_plan_draft_pad.py b/test/registered/kernel/attention/test_deepseek_v4_compress_plan_bounds.py similarity index 98% rename from test/registered/kernels/ops/attention/test_deepseek_v4_compress_plan_draft_pad.py rename to test/registered/kernel/attention/test_deepseek_v4_compress_plan_bounds.py index ce0649534a10..37bd285bffc6 100644 --- a/test/registered/kernels/ops/attention/test_deepseek_v4_compress_plan_draft_pad.py +++ b/test/registered/kernel/attention/test_deepseek_v4_compress_plan_bounds.py @@ -17,7 +17,7 @@ register_cuda_ci(est_time=20, stage="base-b-kernel-unit", runner_config="1-gpu-large") -class TestCompressWritePlanDraftPad(CustomTestCase): +class TestCompressWritePlanBounds(CustomTestCase): def test_64k_prefill_preserves_last_token(self): """65536 tokens fit uint16 indices; the last token must not wrap or vanish.""" for cr in (4, 128): From 235af929b6202e11f55b6a8e31898ace899b8c3e Mon Sep 17 00:00:00 2001 From: hnyls2002 Date: Wed, 16 Sep 2026 17:03:21 -0700 Subject: [PATCH 29/30] mhc: fold compensated hc_mix_stats variants into mhc.py; derive sizes from hc_mult; share the slice count --- .../ops/layernorm/hc_mix_stats_bf16x3.py | 100 ---------- .../ops/layernorm/hc_mix_stats_deepgemm.py | 70 ------- python/sglang/kernels/ops/layernorm/mhc.py | 188 ++++++++++++++++++ 3 files changed, 188 insertions(+), 170 deletions(-) delete mode 100644 python/sglang/kernels/ops/layernorm/hc_mix_stats_bf16x3.py delete mode 100644 python/sglang/kernels/ops/layernorm/hc_mix_stats_deepgemm.py diff --git a/python/sglang/kernels/ops/layernorm/hc_mix_stats_bf16x3.py b/python/sglang/kernels/ops/layernorm/hc_mix_stats_bf16x3.py deleted file mode 100644 index f6fe6000cd13..000000000000 --- a/python/sglang/kernels/ops/layernorm/hc_mix_stats_bf16x3.py +++ /dev/null @@ -1,100 +0,0 @@ -"""Compensated mHC prefill projection with a shared activation load. - -Keep three BF16 components of the FP32 weights and accumulate their products -separately. The 16 fixed K slices bound FP32 accumulation error, as in the -compensated DeepGEMM path, while avoiding its second activation read/reduction. -""" - -import torch -import triton -import triton.language as tl - - -def split_bf16_hc_weight(weight: torch.Tensor): - assert weight.dtype == torch.float32 and weight.is_contiguous() - high = weight.bfloat16() - residual = weight - high.float() - middle = residual.bfloat16() - low = (residual - middle.float()).bfloat16() - return high, middle, low - - -@triton.jit -def _hc_mix_stats_bf16x3(X, W_HI, W_MID, W_LO, MIX, SQ, M, BLOCK_M: tl.constexpr): - # M stays runtime-valued so variable prefill lengths reuse the same binary. - rows = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M) - cols = tl.arange(0, 32) - # 20480 input features / 16 independent slices. - start = tl.program_id(1) * 1280 - ks = start + tl.arange(0, 64) - hi = tl.zeros((BLOCK_M, 32), tl.float32) - mid = tl.zeros((BLOCK_M, 32), tl.float32) - lo = tl.zeros((BLOCK_M, 32), tl.float32) - sq = tl.zeros((BLOCK_M,), tl.float32) - for block in range(20): - k = ks + block * 64 - x = tl.load( - X + rows[:, None].to(tl.int64) * 20480 + k[None, :], - rows[:, None] < M, - 0, - ) - offsets = cols[None, :] * 20480 + k[:, None] - w_hi = tl.load(W_HI + offsets, cols[None, :] < 24, 0) - w_mid = tl.load(W_MID + offsets, cols[None, :] < 24, 0) - w_lo = tl.load(W_LO + offsets, cols[None, :] < 24, 0) - hi = tl.dot(x, w_hi, hi) - mid = tl.dot(x, w_mid, mid) - lo = tl.dot(x, w_lo, lo) - xf = x.to(tl.float32) - sq += tl.sum(xf * xf, 1) - offsets = (tl.program_id(1) * M + rows[:, None]) * 24 + cols[None, :] - tl.store(MIX + offsets, (hi + mid) + lo, (rows[:, None] < M) & (cols[None, :] < 24)) - tl.store(SQ + tl.program_id(1) * M + rows, sq, rows < M) - - -def hc_mix_stats_sinkhorn_bf16x3( - x: torch.Tensor, - weight_parts, - scale: torch.Tensor, - base: torch.Tensor, - sinkhorn_iters: int, - rms_eps: float, - hc_eps: float, -): - from sglang.kernels.ops.layernorm.mhc import _hc_mix_reduce_sinkhorn_kernel - - m = x.shape[0] - assert x.shape == (m, 20480) and x.is_contiguous() - assert x.dtype == torch.bfloat16 and 4096 <= m <= 65536 - assert len(weight_parts) == 3 - assert all( - w.shape == (24, 20480) and w.dtype == torch.bfloat16 and w.is_contiguous() - for w in weight_parts - ) - mix = torch.empty((16, m, 24), device=x.device, dtype=torch.float32) - sq = torch.empty((16, m), device=x.device, dtype=torch.float32) - pre = torch.empty((m, 4), device=x.device, dtype=torch.float32) - post = torch.empty_like(pre) - comb = torch.empty((m, 4, 4), device=x.device, dtype=torch.float32) - _hc_mix_stats_bf16x3[(triton.cdiv(m, 128), 16)]( - x, *weight_parts, mix, sq, m, 128, num_warps=4, num_stages=3 - ) - _hc_mix_reduce_sinkhorn_kernel[(m,)]( - mix, - sq, - scale, - base, - pre, - post, - comb, - m, - 1.0 / 20480, - rms_eps, - MIX=24, - HC=4, - NUM_SLICES=16, - ITERS=sinkhorn_iters, - EPS=hc_eps, - num_warps=1, - ) - return pre, post, comb diff --git a/python/sglang/kernels/ops/layernorm/hc_mix_stats_deepgemm.py b/python/sglang/kernels/ops/layernorm/hc_mix_stats_deepgemm.py deleted file mode 100644 index 15b9d41ff2cb..000000000000 --- a/python/sglang/kernels/ops/layernorm/hc_mix_stats_deepgemm.py +++ /dev/null @@ -1,70 +0,0 @@ -"""Compensated FP32 mHC projections for SM100 batches with at least 128 rows. - -The small-row and batch-invariant paths remain in mhc.py. Native TF32 discards -too much of the FP32 projection weights, so evaluate their high and residual -components separately and bound accumulation length with a fixed split count. -""" - -import torch - -_NUM_SPLITS = 16 - - -def split_tf32_hc_weight(weight: torch.Tensor): - assert weight.dtype == torch.float32 and weight.is_contiguous() - high = (weight.view(torch.int32) & -8192).view(torch.float32) - return high, weight - high - - -def hc_mix_stats_sinkhorn_deepgemm( - x_flat: torch.Tensor, - weight_parts, - hc_scale: torch.Tensor, - hc_base: torch.Tensor, - sinkhorn_iters: int, - rms_eps: float, - hc_eps: float, -): - from sglang.kernels.ops.layernorm.mhc import _hc_mix_reduce_sinkhorn_kernel - from sglang.srt.layers.deep_gemm_wrapper.entrypoint import tf32_hc_prenorm_gemm - - assert x_flat.dtype == torch.bfloat16 and x_flat.is_contiguous() - m, k = x_flat.shape - high, low = weight_parts - assert k == 20480 and high.shape == low.shape == (24, k) - dev = x_flat.device - pre = torch.empty((m, 4), dtype=torch.float32, device=dev) - post = torch.empty_like(pre) - comb = torch.empty((m, 4, 4), dtype=torch.float32, device=dev) - if m == 0: - return pre, post, comb - - mix_hi = torch.empty((_NUM_SPLITS, m, 24), dtype=torch.float32, device=dev) - mix_lo = torch.empty_like(mix_hi) - sq = torch.empty((_NUM_SPLITS, m), dtype=torch.float32, device=dev) - tf32_hc_prenorm_gemm(x_flat, high, mix_hi, sq, _NUM_SPLITS) - # The row sum of squares depends only on x_flat, so the second projection - # recomputes exactly the same values; let it overwrite sq instead of - # allocating a (_NUM_SPLITS, m) fp32 scratch that is thrown away - # (4 MiB at m = 65536). - tf32_hc_prenorm_gemm(x_flat, low, mix_lo, sq, _NUM_SPLITS) - _hc_mix_reduce_sinkhorn_kernel[(m,)]( - mix_hi, - sq, - hc_scale, - hc_base, - pre, - post, - comb, - m, - 1.0 / k, - rms_eps, - MIX=24, - HC=4, - NUM_SLICES=_NUM_SPLITS, - ITERS=sinkhorn_iters, - EPS=hc_eps, - part_mix_residual_ptr=mix_lo, - num_warps=1, - ) - return pre, post, comb diff --git a/python/sglang/kernels/ops/layernorm/mhc.py b/python/sglang/kernels/ops/layernorm/mhc.py index 2c7a89dc1025..10574d59c2aa 100644 --- a/python/sglang/kernels/ops/layernorm/mhc.py +++ b/python/sglang/kernels/ops/layernorm/mhc.py @@ -2368,6 +2368,194 @@ def hc_mix_stats_sinkhorn( return pre, post, comb +# Compensated projections: the fp32 mixing weight is split into components the +# tensor cores take exactly (three bf16 parts, or a tf32 high part and its fp32 +# residual) and accumulated over a fixed number of K slices, so the fp32 +# accumulation error stays bounded like the DeepGEMM path's. +_HC_MIX_COMPENSATED_SLICES = 16 +_HC_MIX_BF16X3_BLOCK_M = 128 + + +def split_bf16_hc_weight(weight: torch.Tensor): + assert weight.dtype == torch.float32 and weight.is_contiguous() + high = weight.bfloat16() + residual = weight - high.float() + middle = residual.bfloat16() + low = (residual - middle.float()).bfloat16() + return high, middle, low + + +def split_tf32_hc_weight(weight: torch.Tensor): + assert weight.dtype == torch.float32 and weight.is_contiguous() + high = (weight.view(torch.int32) & -8192).view(torch.float32) + return high, weight - high + + +@triton.jit +def _hc_mix_stats_bf16x3_kernel( + X, + W_HI, + W_MID, + W_LO, + MIX, + SQ, + M, + K: tl.constexpr, + K_PER_SLICE: tl.constexpr, + MIX_COLS: tl.constexpr, + MIX_PAD: tl.constexpr, + BLOCK_K: tl.constexpr, + BLOCK_M: tl.constexpr, +): + # M stays runtime-valued so variable prefill lengths reuse the same binary. + rows = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M) + cols = tl.arange(0, MIX_PAD) + start = tl.program_id(1) * K_PER_SLICE + ks = start + tl.arange(0, BLOCK_K) + hi = tl.zeros((BLOCK_M, MIX_PAD), tl.float32) + mid = tl.zeros((BLOCK_M, MIX_PAD), tl.float32) + lo = tl.zeros((BLOCK_M, MIX_PAD), tl.float32) + sq = tl.zeros((BLOCK_M,), tl.float32) + for block in range(K_PER_SLICE // BLOCK_K): + k = ks + block * BLOCK_K + x = tl.load( + X + rows[:, None].to(tl.int64) * K + k[None, :], + rows[:, None] < M, + 0, + ) + offsets = cols[None, :] * K + k[:, None] + w_hi = tl.load(W_HI + offsets, cols[None, :] < MIX_COLS, 0) + w_mid = tl.load(W_MID + offsets, cols[None, :] < MIX_COLS, 0) + w_lo = tl.load(W_LO + offsets, cols[None, :] < MIX_COLS, 0) + hi = tl.dot(x, w_hi, hi) + mid = tl.dot(x, w_mid, mid) + lo = tl.dot(x, w_lo, lo) + xf = x.to(tl.float32) + sq += tl.sum(xf * xf, 1) + offsets = (tl.program_id(1) * M + rows[:, None]) * MIX_COLS + cols[None, :] + tl.store( + MIX + offsets, (hi + mid) + lo, (rows[:, None] < M) & (cols[None, :] < MIX_COLS) + ) + tl.store(SQ + tl.program_id(1) * M + rows, sq, rows < M) + + +def hc_mix_stats_sinkhorn_bf16x3( + x: torch.Tensor, + weight_parts, + scale: torch.Tensor, + base: torch.Tensor, + sinkhorn_iters: int, + rms_eps: float, + hc_eps: float, + hc_mult: int = 4, +): + m, k = x.shape + mix = (2 + hc_mult) * hc_mult + slices = _HC_MIX_COMPENSATED_SLICES + assert x.is_contiguous() and x.dtype == torch.bfloat16 and 4096 <= m <= 65536 + assert k % (slices * _HC_MIX_BLOCK_K) == 0 + assert len(weight_parts) == 3 + assert all( + w.shape == (mix, k) and w.dtype == torch.bfloat16 and w.is_contiguous() + for w in weight_parts + ) + part_mix = torch.empty((slices, m, mix), device=x.device, dtype=torch.float32) + sq = torch.empty((slices, m), device=x.device, dtype=torch.float32) + pre = torch.empty((m, hc_mult), device=x.device, dtype=torch.float32) + post = torch.empty_like(pre) + comb = torch.empty((m, hc_mult, hc_mult), device=x.device, dtype=torch.float32) + _hc_mix_stats_bf16x3_kernel[(triton.cdiv(m, _HC_MIX_BF16X3_BLOCK_M), slices)]( + x, + *weight_parts, + part_mix, + sq, + m, + K=k, + K_PER_SLICE=k // slices, + MIX_COLS=mix, + MIX_PAD=triton.next_power_of_2(mix), + BLOCK_K=_HC_MIX_BLOCK_K, + BLOCK_M=_HC_MIX_BF16X3_BLOCK_M, + num_warps=4, + num_stages=3, + ) + _hc_mix_reduce_sinkhorn_kernel[(m,)]( + part_mix, + sq, + scale, + base, + pre, + post, + comb, + m, + 1.0 / k, + rms_eps, + MIX=mix, + HC=hc_mult, + NUM_SLICES=slices, + ITERS=sinkhorn_iters, + EPS=hc_eps, + num_warps=1, + ) + return pre, post, comb + + +def hc_mix_stats_sinkhorn_deepgemm( + x_flat: torch.Tensor, + weight_parts, + hc_scale: torch.Tensor, + hc_base: torch.Tensor, + sinkhorn_iters: int, + rms_eps: float, + hc_eps: float, + hc_mult: int = 4, +): + from sglang.srt.layers.deep_gemm_wrapper.entrypoint import tf32_hc_prenorm_gemm + + assert x_flat.dtype == torch.bfloat16 and x_flat.is_contiguous() + m, k = x_flat.shape + mix = (2 + hc_mult) * hc_mult + slices = _HC_MIX_COMPENSATED_SLICES + high, low = weight_parts + assert high.shape == low.shape == (mix, k) + dev = x_flat.device + pre = torch.empty((m, hc_mult), dtype=torch.float32, device=dev) + post = torch.empty_like(pre) + comb = torch.empty((m, hc_mult, hc_mult), dtype=torch.float32, device=dev) + if m == 0: + return pre, post, comb + + mix_hi = torch.empty((slices, m, mix), dtype=torch.float32, device=dev) + mix_lo = torch.empty_like(mix_hi) + sq = torch.empty((slices, m), dtype=torch.float32, device=dev) + tf32_hc_prenorm_gemm(x_flat, high, mix_hi, sq, slices) + # The row sum of squares depends only on x_flat, so the second projection + # recomputes exactly the same values; let it overwrite sq instead of + # allocating a (slices, m) fp32 scratch that is thrown away + # (4 MiB at m = 65536). + tf32_hc_prenorm_gemm(x_flat, low, mix_lo, sq, slices) + _hc_mix_reduce_sinkhorn_kernel[(m,)]( + mix_hi, + sq, + hc_scale, + hc_base, + pre, + post, + comb, + m, + 1.0 / k, + rms_eps, + MIX=mix, + HC=hc_mult, + NUM_SLICES=slices, + ITERS=sinkhorn_iters, + EPS=hc_eps, + part_mix_residual_ptr=mix_lo, + num_warps=1, + ) + return pre, post, comb + + def hc_combine( x_flat: torch.Tensor, pre: torch.Tensor, hc: int, out_dtype: torch.dtype ) -> torch.Tensor: From 6d606f06746bf387d2714e7393bfad7f96302787 Mon Sep 17 00:00:00 2001 From: hnyls2002 Date: Wed, 16 Sep 2026 17:07:32 -0700 Subject: [PATCH 30/30] trim comments --- python/sglang/kernels/ops/layernorm/mhc.py | 31 ++++++++-------------- 1 file changed, 11 insertions(+), 20 deletions(-) diff --git a/python/sglang/kernels/ops/layernorm/mhc.py b/python/sglang/kernels/ops/layernorm/mhc.py index 10574d59c2aa..3195265e9539 100644 --- a/python/sglang/kernels/ops/layernorm/mhc.py +++ b/python/sglang/kernels/ops/layernorm/mhc.py @@ -365,9 +365,8 @@ def hc_split_sinkhorn( ): b, s, _ = mixes.size() if b * s == 0: - # DP attention's idle forward carries no tokens. Every backend below - # derives its grid from the token count, and CUDA rejects a launch with - # a zero-sized grid, so answer the empty batch directly. + # DP attention's idle forward carries no tokens, and every backend below + # derives a CUDA grid from the token count; a zero-sized grid is rejected. return ( mixes.new_empty(b, s, hc_mult), mixes.new_empty(b, s, hc_mult), @@ -2054,9 +2053,7 @@ def _hc_mix_stats_partial_kernel( BLOCK_K: tl.constexpr, DOT_PRECISION: tl.constexpr, ): - """Mixing dot products and row sum of squares over one K slice; the slicing - and tiles are compile-time constants, so a row's fp32 operation sequence - does not depend on the batch size.""" + """Mixing dot products and row sum of squares over one K slice.""" pid_m = tl.program_id(0) pid_s = tl.program_id(1) offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M) @@ -2127,7 +2124,7 @@ def _hc_mix_stats_reduce_kernel( # K slicing, BLOCK_K and dot precision must stay independent of M; -# a row must produce the same bits alone and in any batch. +# a row must produce the same bits alone and in any batch routed to this backend. _HC_MIX_SLICE_CHOICES = (80, 64, 40, 32, 16, 8, 4, 2, 1) _HC_MIX_BLOCK_M = 32 _HC_MIX_BLOCK_K = 64 @@ -2155,7 +2152,6 @@ def _block_m_for(m: int) -> int: def _num_stages_for(m: int, k: int) -> int: # GB300 verify batches benefit from a smaller shared-memory footprint. - # This changes memory scheduling only; K tiles and reduction order stay fixed. if get_platform().is_blackwell and k == 20480 and 64 <= m <= 384: return 1 return _HC_MIX_NUM_STAGES @@ -2176,7 +2172,7 @@ def hc_mix_stats(x_flat: torch.Tensor, hc_fn: torch.Tensor, eps: float) -> torch x_flat is [M, K] in any float dtype; hc_fn is [MIX, K] fp32; returns [M, MIX] fp32. K slicing and reduction order are independent of M, so each row is bitwise - identical whether computed alone or in a batch. + identical whether computed alone or in any batch routed to this backend. """ assert x_flat.dim() == 2 and hc_fn.dim() == 2 assert x_flat.stride(1) == 1 and hc_fn.stride(1) == 1 @@ -2248,9 +2244,7 @@ def _hc_mix_reduce_sinkhorn_kernel( EPS: tl.constexpr, part_mix_residual_ptr=None, ): - """One CTA per row keeps the sinkhorn reductions two-dimensional. - Per-row arithmetic follows the slice reduction, then the Triton sinkhorn. - """ + """One CTA per row keeps the sinkhorn reductions two-dimensional.""" row = tl.program_id(0) if row >= m: return @@ -2308,8 +2302,8 @@ def hc_mix_stats_sinkhorn( ): """Fuse the reduce and sinkhorn stages of hc_mix_stats followed by hc_split_sinkhorn. - The split-K kernel fixes the reduction order and preserves batch invariance. - Sinkhorn uses the Triton port's transcendental lowering, which differs from TileLang. + The split-K kernel fixes the reduction order; sinkhorn uses the Triton port's + transcendental lowering, which differs from TileLang. """ assert x_flat.dim() == 2 and hc_fn.dim() == 2 assert x_flat.stride(1) == 1 and hc_fn.stride(1) == 1 @@ -2370,8 +2364,7 @@ def hc_mix_stats_sinkhorn( # Compensated projections: the fp32 mixing weight is split into components the # tensor cores take exactly (three bf16 parts, or a tf32 high part and its fp32 -# residual) and accumulated over a fixed number of K slices, so the fp32 -# accumulation error stays bounded like the DeepGEMM path's. +# residual), accumulated over a fixed slice count to bound the fp32 error. _HC_MIX_COMPENSATED_SLICES = 16 _HC_MIX_BF16X3_BLOCK_M = 128 @@ -2529,10 +2522,8 @@ def hc_mix_stats_sinkhorn_deepgemm( mix_lo = torch.empty_like(mix_hi) sq = torch.empty((slices, m), dtype=torch.float32, device=dev) tf32_hc_prenorm_gemm(x_flat, high, mix_hi, sq, slices) - # The row sum of squares depends only on x_flat, so the second projection - # recomputes exactly the same values; let it overwrite sq instead of - # allocating a (slices, m) fp32 scratch that is thrown away - # (4 MiB at m = 65536). + # sq depends only on x_flat, so the second projection recomputes the same + # values; overwriting it avoids a throwaway (slices, m) scratch, 4 MiB at m=65536. tf32_hc_prenorm_gemm(x_flat, low, mix_lo, sq, slices) _hc_mix_reduce_sinkhorn_kernel[(m,)]( mix_hi,