From 7c23d2255a1f3d3f850701d635003aad195bb70a Mon Sep 17 00:00:00 2001 From: Xinyuan Tong <115166877+JustinTong0323@users.noreply.github.com> Date: Mon, 22 Jun 2026 21:10:43 +0100 Subject: [PATCH] [minimax-m3] Split 1/4: sparse attention ops + JIT kernels + config foundation (#28712) --- .../kernels/fused_moe_triton/common_utils.py | 37 +- .../tuning_fused_moe_triton.py | 6 +- docker/rocm.Dockerfile | 3 +- .../csrc/minimax/fused_gemma_qknorm_rope.cuh | 200 ++++ .../csrc/minimax/fused_store_kv_index.cuh | 171 +++ .../csrc/minimax/minimax_decode_topk.cuh | 535 +++++++++ .../csrc/minimax/per_token_quant_ue8m0.cuh | 280 +++++ .../jit_kernel/csrc/moe/moe_topk_sigmoid.cuh | 544 +++++++++ .../sglang/jit_kernel/minimax_decode_topk.py | 129 ++ .../sglang/jit_kernel/minimax_m3/__init__.py | 25 + .../jit_kernel/minimax_m3/qk_norm_rope.py | 659 ++++++++++ .../sglang/jit_kernel/minimax_m3/rmsnorm.py | 148 +++ python/sglang/jit_kernel/minimax_m3/swiglu.py | 201 ++++ .../sglang/jit_kernel/minimax_qknorm_rope.py | 170 +++ .../sglang/jit_kernel/minimax_quant_ue8m0.py | 105 ++ .../jit_kernel/minimax_store_kv_index.py | 79 ++ python/sglang/jit_kernel/moe_topk_sigmoid.py | 105 ++ .../jit_kernel/tests/test_minimax_m3_mxfp8.py | 265 ++++ .../tests/test_minimax_m3_rmsnorm.py | 78 ++ .../jit_kernel/tests/test_moe_topk_sigmoid.py | 493 ++++++++ python/sglang/srt/configs/model_config.py | 71 +- python/sglang/srt/environ.py | 7 + .../minimax_sparse_ops/common/index.py | 67 ++ .../minimax_sparse_ops/common/utils.py | 262 ++++ .../decode/flash_with_topk_idx.py | 1065 +++++++++++++++++ .../minimax_sparse_ops/decode/topk_sparse.py | 419 +++++++ .../minimax_sparse_ops/minimax_sparse.py | 234 ++++ .../attention/minimax_sparse_ops/msa.py | 353 ++++++ .../naive/flash_with_topk_idx.py | 84 ++ .../minimax_sparse_ops/naive/topk_sparse.py | 88 ++ .../prefill/flash_with_topk_idx.py | 563 +++++++++ .../minimax_sparse_ops/prefill/topk_sparse.py | 356 ++++++ .../tests/test_flash_with_topk_idx.py | 480 ++++++++ .../tests/test_sparse_gqa.py | 362 ++++++ .../srt/layers/moe/fused_moe_triton/layer.py | 2 + .../sglang/srt/layers/moe/moe_runner/base.py | 4 + .../srt/layers/moe/moe_runner/triton.py | 116 +- .../minimax_m3_gfx950_mxfp8_compact_moe.json | 70 ++ ...dtype=fp8_w8a8,block_shape=[128, 128].json | 164 +++ .../E=129,N=384,device_name=NVIDIA_H200.json | 146 +++ .../moe/moe_runner/triton_utils/fused_moe.py | 32 +- .../triton_utils/mxfp8_moe_amd_gfx95.py | 433 +++++++ .../layers/quantization/mxfp8_amd_gfx95.py | 333 ++++++ .../minimax/bench_minimax_decode_topk.py | 121 ++ .../minimax/bench_minimax_qknorm_rope.py | 146 +++ .../minimax/bench_minimax_store_kv_index.py | 67 ++ .../jit/minimax/test_minimax_decode_topk.py | 144 +++ .../test_minimax_decode_topk_page_table.py | 178 +++ .../jit/minimax/test_minimax_qknorm_rope.py | 178 +++ .../minimax/test_minimax_store_kv_index.py | 78 ++ .../jit/test_minimax_m3_qk_norm_rope.py | 334 ++++++ 51 files changed, 11157 insertions(+), 33 deletions(-) create mode 100644 python/sglang/jit_kernel/csrc/minimax/fused_gemma_qknorm_rope.cuh create mode 100644 python/sglang/jit_kernel/csrc/minimax/fused_store_kv_index.cuh create mode 100644 python/sglang/jit_kernel/csrc/minimax/minimax_decode_topk.cuh create mode 100644 python/sglang/jit_kernel/csrc/minimax/per_token_quant_ue8m0.cuh create mode 100644 python/sglang/jit_kernel/csrc/moe/moe_topk_sigmoid.cuh create mode 100644 python/sglang/jit_kernel/minimax_decode_topk.py create mode 100644 python/sglang/jit_kernel/minimax_m3/__init__.py create mode 100644 python/sglang/jit_kernel/minimax_m3/qk_norm_rope.py create mode 100644 python/sglang/jit_kernel/minimax_m3/rmsnorm.py create mode 100644 python/sglang/jit_kernel/minimax_m3/swiglu.py create mode 100644 python/sglang/jit_kernel/minimax_qknorm_rope.py create mode 100644 python/sglang/jit_kernel/minimax_quant_ue8m0.py create mode 100644 python/sglang/jit_kernel/minimax_store_kv_index.py create mode 100644 python/sglang/jit_kernel/moe_topk_sigmoid.py create mode 100644 python/sglang/jit_kernel/tests/test_minimax_m3_mxfp8.py create mode 100644 python/sglang/jit_kernel/tests/test_minimax_m3_rmsnorm.py create mode 100644 python/sglang/jit_kernel/tests/test_moe_topk_sigmoid.py create mode 100644 python/sglang/srt/layers/attention/minimax_sparse_ops/common/index.py create mode 100644 python/sglang/srt/layers/attention/minimax_sparse_ops/common/utils.py create mode 100644 python/sglang/srt/layers/attention/minimax_sparse_ops/decode/flash_with_topk_idx.py create mode 100644 python/sglang/srt/layers/attention/minimax_sparse_ops/decode/topk_sparse.py create mode 100644 python/sglang/srt/layers/attention/minimax_sparse_ops/minimax_sparse.py create mode 100644 python/sglang/srt/layers/attention/minimax_sparse_ops/msa.py create mode 100644 python/sglang/srt/layers/attention/minimax_sparse_ops/naive/flash_with_topk_idx.py create mode 100644 python/sglang/srt/layers/attention/minimax_sparse_ops/naive/topk_sparse.py create mode 100644 python/sglang/srt/layers/attention/minimax_sparse_ops/prefill/flash_with_topk_idx.py create mode 100644 python/sglang/srt/layers/attention/minimax_sparse_ops/prefill/topk_sparse.py create mode 100644 python/sglang/srt/layers/attention/minimax_sparse_ops/tests/test_flash_with_topk_idx.py create mode 100644 python/sglang/srt/layers/attention/minimax_sparse_ops/tests/test_sparse_gqa.py create mode 100644 python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/minimax_m3_gfx950_mxfp8_compact_moe.json create mode 100644 python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=128,N=384,device_name=AMD_Instinct_MI300X,dtype=fp8_w8a8,block_shape=[128, 128].json create mode 100644 python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=129,N=384,device_name=NVIDIA_H200.json create mode 100644 python/sglang/srt/layers/moe/moe_runner/triton_utils/mxfp8_moe_amd_gfx95.py create mode 100644 python/sglang/srt/layers/quantization/mxfp8_amd_gfx95.py create mode 100644 test/registered/jit/benchmark/minimax/bench_minimax_decode_topk.py create mode 100644 test/registered/jit/benchmark/minimax/bench_minimax_qknorm_rope.py create mode 100644 test/registered/jit/benchmark/minimax/bench_minimax_store_kv_index.py create mode 100644 test/registered/jit/minimax/test_minimax_decode_topk.py create mode 100644 test/registered/jit/minimax/test_minimax_decode_topk_page_table.py create mode 100644 test/registered/jit/minimax/test_minimax_qknorm_rope.py create mode 100644 test/registered/jit/minimax/test_minimax_store_kv_index.py create mode 100644 test/registered/jit/test_minimax_m3_qk_norm_rope.py diff --git a/benchmark/kernels/fused_moe_triton/common_utils.py b/benchmark/kernels/fused_moe_triton/common_utils.py index 66c1f495c..cb7422d6a 100644 --- a/benchmark/kernels/fused_moe_triton/common_utils.py +++ b/benchmark/kernels/fused_moe_triton/common_utils.py @@ -1,4 +1,5 @@ import json +from types import SimpleNamespace from typing import Dict, List, TypedDict import torch @@ -7,7 +8,7 @@ from sglang.srt.layers.moe.moe_runner.triton_utils.fused_moe import get_config_d from sglang.srt.layers.moe.moe_runner.triton_utils.fused_moe_triton_config import ( get_config_file_name, ) -from sglang.srt.utils import is_hip +from sglang.srt.utils import is_gfx95_supported, is_hip from sglang.srt.utils.hf_transformers_utils import get_config @@ -59,7 +60,13 @@ def get_model_config( assert len(block_shape) == 2 # Replace config with text_config for encoder-decoder models after getting block_shape and architecture if hasattr(config, "text_config"): - config = config.get_text_config() + text_config = config.get_text_config() + # Some models (e.g. MiniMax-M3) carry text_config as a plain dict; wrap + # it so downstream attribute access works uniformly. + if isinstance(text_config, dict): + config = SimpleNamespace(**text_config) + else: + config = text_config hidden_size = config.hidden_size if architecture == "DbrxForCausalLM": @@ -148,6 +155,17 @@ def get_model_config( E = config.num_experts // ep_size topk = config.num_experts_per_tok intermediate_size = config.moe_intermediate_size + elif architecture == "MiniMaxM3SparseForConditionalGeneration": + # Serving fuses the shared expert into the routed-expert tensor by + # default (E = num_local_experts + 1), so tune for that shape unless + # fusion is explicitly disabled. + E = config.num_local_experts // ep_size + ( + 0 if disable_shared_experts_fusion else 1 + ) + topk = config.num_experts_per_tok + ( + 0 if disable_shared_experts_fusion or topk_ids_dir is None else 1 + ) + intermediate_size = config.intermediate_size else: # Default: Mixtral E = config.num_local_experts // ep_size @@ -158,12 +176,25 @@ def get_model_config( intermediate_size, tp_size, ep_size ) + # gfx942 (MI300X) remaps MXFP8 [1,32]->[128,128] at load; tune must match. + # sm100/gfx95 run native [1,32] and must not remap (mirrors fp8.py). + if ( + architecture == "MiniMaxM3SparseForConditionalGeneration" + and is_hip() + and not is_gfx95_supported() + and block_shape == [1, 32] + ): + block_shape = [128, 128] + + # text_config may not carry torch_dtype; fall back to bf16. + torch_dtype = getattr(config, "torch_dtype", None) or torch.bfloat16 + return { "num_experts": E, "topk": topk, "hidden_size": hidden_size, "shard_intermediate_size": shard_intermediate_size, - "dtype": config.torch_dtype, + "dtype": torch_dtype, "block_shape": block_shape, "architecture": architecture, } diff --git a/benchmark/kernels/fused_moe_triton/tuning_fused_moe_triton.py b/benchmark/kernels/fused_moe_triton/tuning_fused_moe_triton.py index 3efd128a5..536cc6913 100644 --- a/benchmark/kernels/fused_moe_triton/tuning_fused_moe_triton.py +++ b/benchmark/kernels/fused_moe_triton/tuning_fused_moe_triton.py @@ -248,8 +248,10 @@ class BenchmarkWorker: torch.get_device_module().manual_seed_all(0) self.seed = seed # Get the device ID to allocate tensors and kernels - # on the respective GPU. - self.device_id = int(ray.get_gpu_ids()[0]) + # on the respective GPU. Ray isolates each worker to a single visible + # GPU via CUDA_VISIBLE_DEVICES, so the local ordinal is always 0. On + # ROCm using the global ray gpu id here raises "invalid device ordinal". + self.device_id = 0 if is_hip() else int(ray.get_gpu_ids()[0]) set_global_server_args_for_scheduler(server_args) def benchmark( diff --git a/docker/rocm.Dockerfile b/docker/rocm.Dockerfile index 4e50dd006..70d77d542 100644 --- a/docker/rocm.Dockerfile +++ b/docker/rocm.Dockerfile @@ -569,7 +569,8 @@ RUN if [ "$BUILD_TRITON" = "1" ]; then \ && cd triton-custom \ && git checkout ${TRITON_COMMIT} \ && pip install -r python/requirements.txt \ - && pip install -e .; \ + && pip install -e . \ + && if [ -d python/triton_kernels ]; then pip install -e python/triton_kernels --no-deps; fi; \ fi # ----------------------- diff --git a/python/sglang/jit_kernel/csrc/minimax/fused_gemma_qknorm_rope.cuh b/python/sglang/jit_kernel/csrc/minimax/fused_gemma_qknorm_rope.cuh new file mode 100644 index 000000000..98f9ca32d --- /dev/null +++ b/python/sglang/jit_kernel/csrc/minimax/fused_gemma_qknorm_rope.cuh @@ -0,0 +1,200 @@ +#include +#include + +#include +#include +#include +#include + +#include +#include + +#include + +namespace { + +// Up to this many norm "groups" are fused into one launch. A group is a +// contiguous run of heads (within the per-token row) that share one norm +// weight and all receive RoPE. Heads not covered by any group (e.g. the V / +// index-V regions) are simply never assigned a job and left untouched. +// +// For MiniMax-M3 the groups are, in order: Q (main), K (main), index-Q, +// index-K -- which is why 4 slots are enough. This mirrors the multi-branch +// single-launch design of fused_store_kv_index.cuh (one kernel handles main +// K/V + index K/V), here applied to GemmaRMSNorm + partial RoPE. +constexpr int kMaxGroups = 4; + +struct FusedGemmaQKNormParams { + bf16_t* __restrict__ qkv; + const bf16_t* __restrict__ weight[kMaxGroups]; // per-group norm weight [head_dim] + uint32_t group_offset[kMaxGroups]; // head offset of the group in the row + uint32_t group_count[kMaxGroups]; // number of heads in the group + uint32_t num_groups; + uint32_t total_heads; // sum of group_count[0..num_groups) + const float* __restrict__ cos_sin_cache; + const void* __restrict__ positions; // dtype depends on PosT template + uint32_t num_tokens; + int64_t token_stride; + float eps; +}; + +template +struct FusedTrait { + static_assert(kHeadDim == 128 && kRopeDim == 64, "kernel specialized for HEAD_DIM=128, ROTARY_DIM=64"); + static constexpr uint32_t kWorkerSize = device::kWarpThreads; // full warp per head + SGL_DEVICE static void forward(const FusedGemmaQKNormParams& params) { + using namespace device; + const auto tx = threadIdx.x; + const auto bx = blockIdx.x; + const auto lane_id = tx % kWorkerSize; + const auto work_id = (bx * blockDim.x + tx) / kWorkerSize; + const auto total_heads = params.total_heads; + if (work_id >= params.num_tokens * total_heads) return; + const auto token_id = work_id / total_heads; + const auto grouped_head = work_id % total_heads; + + // Resolve which group this job belongs to (num_groups <= kMaxGroups, so a + // short scan is cheaper than any precomputed table and warp-uniform). + uint32_t g = 0, base = 0; + for (uint32_t i = 0; i < params.num_groups; ++i) { + const auto cnt = params.group_count[i]; + if (grouped_head < base + cnt) { + g = i; + break; + } + base += cnt; + } + const auto local_head = grouped_head - base; + const auto head_id = params.group_offset[g] + local_head; + const auto weight = params.weight[g]; + const auto input = params.qkv + token_id * params.token_stride + head_id * kHeadDim; + + // prefetch weight and rope index + const auto idx_0 = lane_id + 0; // rope first half [0,32) + const auto idx_1 = lane_id + 32; // rope second half [32, 64) + const auto idx_2 = lane_id + 64; // pass + const auto idx_3 = lane_id + 96; // pass + const auto w0_bf16 = weight[idx_0]; + const auto w1_bf16 = weight[idx_1]; + const auto w2_bf16 = weight[idx_2]; + const auto w3_bf16 = weight[idx_3]; + const auto rope_idx = static_cast(params.positions)[token_id]; + PDLWaitPrimary(); + + // load input and compute RMS, fp32 accumulation for stability + const auto i0_bf16 = input[idx_0]; + const auto i1_bf16 = input[idx_1]; + const auto i2_bf16 = input[idx_2]; + const auto i3_bf16 = input[idx_3]; + const auto [i0, i1] = cast(bf16x2_t{i0_bf16, i1_bf16}); + const auto [i2, i3] = cast(bf16x2_t{i2_bf16, i3_bf16}); + const auto ss = warp::reduce_sum(i0 * i0 + i1 * i1 + i2 * i2 + i3 * i3); + const auto inv_rms = rsqrtf(ss / static_cast(kHeadDim) + params.eps); + + // apply norm + const auto [w0, w1] = cast(bf16x2_t{w0_bf16, w1_bf16}); + const auto [w2, w3] = cast(bf16x2_t{w2_bf16, w3_bf16}); + const auto n0 = i0 * inv_rms * (1.0f + w0); + const auto n1 = i1 * inv_rms * (1.0f + w1); + const auto n2 = i2 * inv_rms * (1.0f + w2); + const auto n3 = i3 * inv_rms * (1.0f + w3); + + const auto cs = params.cos_sin_cache + rope_idx * kRopeDim; + const auto cos = cs[lane_id]; + const auto sin = cs[lane_id + kRopeDim / 2]; + + // apply rope to the first kRopeDim dims, and write back + device::PDLTriggerSecondary(); + const auto o0 = n0 * cos - n1 * sin; + const auto o1 = n1 * cos + n0 * sin; + const auto [o0_bf16, o1_bf16] = cast(fp32x2_t{o0, o1}); + const auto [o2_bf16, o3_bf16] = cast(fp32x2_t{n2, n3}); + input[idx_0] = o0_bf16; + input[idx_1] = o1_bf16; + input[idx_2] = o2_bf16; + input[idx_3] = o3_bf16; + } +}; + +template +__global__ void fused_gemma_qknorm_rope_kernel(const __grid_constant__ FusedGemmaQKNormParams params) { + return Trait::forward(params); +} + +// Multi-group fused GemmaRMSNorm + partial NeoX RoPE, in place over `qkv`. +// +// Up to kMaxGroups norm groups are passed as (weight, head offset, head count) +// triples. `w0..w3` are the per-group norm weights ([head_dim] bf16 each); the +// `offN` / `cntN` scalars give each group's head offset (within the per-token +// row) and head count. These offsets/counts are host-known constants (passed as +// scalars, never device tensors) so the launch stays CUDA-graph capturable. +// Weight slots beyond `num_groups` may be dummies (e.g. == w0); the kernel +// never reads them because `num_groups` bounds the group scan. The main Q/K and +// index-Q/index-K heads are all normed and rotated in one launch; the V / +// index-V heads, lying outside every group, are left untouched. +template +void fused_gemma_qknorm_rope( + tvm::ffi::TensorView qkv, + tvm::ffi::TensorView w0, + tvm::ffi::TensorView w1, + tvm::ffi::TensorView w2, + tvm::ffi::TensorView w3, + tvm::ffi::TensorView cos_sin_cache, + tvm::ffi::TensorView positions, + int64_t off0, + int64_t cnt0, + int64_t off1, + int64_t cnt1, + int64_t off2, + int64_t cnt2, + int64_t off3, + int64_t cnt3, + int64_t num_groups, + double eps) { + using namespace host; + auto N = SymbolicSize{"num_tokens"}; + auto device = SymbolicDevice{}; + constexpr auto D = HEAD_DIM; + constexpr auto R = ROTARY_DIM; + device.set_options(); + TensorMatcher({N, -1}).with_dtype().with_device(device).verify(qkv); + TensorMatcher({D}).with_dtype().with_device(device).verify(w0).verify(w1).verify(w2).verify(w3); + TensorMatcher({-1, R}).with_dtype().with_device(device).verify(cos_sin_cache); + TensorMatcher({N}).with_dtype().with_device(device).verify(positions); + + RuntimeCheck(num_groups >= 1 && num_groups <= kMaxGroups); + + const tvm::ffi::TensorView weights[kMaxGroups] = {w0, w1, w2, w3}; + const int64_t offsets[kMaxGroups] = {off0, off1, off2, off3}; + const int64_t counts[kMaxGroups] = {cnt0, cnt1, cnt2, cnt3}; + + auto params = FusedGemmaQKNormParams{}; + params.qkv = static_cast(qkv.data_ptr()); + params.cos_sin_cache = static_cast(cos_sin_cache.data_ptr()); + params.positions = positions.data_ptr(); + params.num_tokens = static_cast(N.unwrap()); + params.num_groups = static_cast(num_groups); + params.token_stride = static_cast(qkv.stride(0)); + params.eps = static_cast(eps); + + uint32_t total_heads = 0; + for (int64_t i = 0; i < kMaxGroups; ++i) { + const auto cnt = (i < num_groups) ? counts[i] : 0; + params.weight[i] = static_cast(weights[i].data_ptr()); + params.group_offset[i] = static_cast(i < num_groups ? offsets[i] : 0); + params.group_count[i] = static_cast(cnt); + total_heads += static_cast(cnt); + } + params.total_heads = total_heads; + + using Trait = FusedTrait; + const auto needed_threads = static_cast(params.num_tokens) * total_heads * device::kWarpThreads; + RuntimeCheck(needed_threads < std::numeric_limits::max()); + if (needed_threads == 0) return; + const uint32_t block_size = 256u; + const uint32_t num_blocks = div_ceil(static_cast(needed_threads), block_size); + LaunchKernel(num_blocks, block_size, device.unwrap()) // + .enable_pdl(kUsePDL)(fused_gemma_qknorm_rope_kernel, params); +} + +} // namespace diff --git a/python/sglang/jit_kernel/csrc/minimax/fused_store_kv_index.cuh b/python/sglang/jit_kernel/csrc/minimax/fused_store_kv_index.cuh new file mode 100644 index 000000000..0aac50472 --- /dev/null +++ b/python/sglang/jit_kernel/csrc/minimax/fused_store_kv_index.cuh @@ -0,0 +1,171 @@ +#include // For TensorMatcher, SymbolicSize/DType/Device +#include // For RuntimeCheck, div_ceil, pointer::offset + +#include // For LaunchKernel, SGL_DEVICE, PDL helpers +#include // For AlignedVector + +#include +#include + +#include + +namespace { + +struct StoreKVIndexParams { + const void* __restrict__ input_ptrs[4]; + void* __restrict__ cache_ptrs[4]; + const void* __restrict__ indices; + int64_t input_stride_bytes[4]; // k stride, v stride, idx_k stride, idx_v stride + int64_t cache_stride_bytes[4]; // k stride, v stride, idx_k stride, idx_v stride + uint32_t num_k_heads; + uint32_t num_total_heads; // 2 * num_k_heads + 1 + has_index_v + uint32_t total_jobs; // batch_size * heads_per_token +}; + +template +struct StoreTrait { + static_assert(kHeadBytes % 16 == 0, "head bytes must be a multiple of 16 (128-bit vector)"); + using vec_t = device::AlignedVector; // 16 bytes / thread + static constexpr uint32_t kWorkerSize = kHeadBytes / 16; // threads per head + + template + SGL_DEVICE static void forward(const StoreKVIndexParams& params) { + using namespace device; + const uint32_t tid = blockIdx.x * blockDim.x + threadIdx.x; + const uint32_t work_id = tid / kWorkerSize; // which (token, head) job + const uint32_t lane = tid % kWorkerSize; // which 16-byte chunk of the head + if (work_id >= params.total_jobs) return; + + const uint32_t num_total_heads = params.num_total_heads; + const uint32_t token_id = work_id / num_total_heads; + const uint32_t slot_id = work_id % num_total_heads; + const uint32_t num_k_heads = params.num_k_heads; + + const auto loc = static_cast(params.indices)[token_id]; + uint32_t head_id, which; + if (slot_id < num_k_heads) { + head_id = slot_id; + which = 0; // K + } else if (slot_id < 2 * num_k_heads) { + head_id = slot_id - num_k_heads; + which = 1; // V + } else if (slot_id == 2 * num_k_heads) { + head_id = 0; + which = 2; // idx K + } else { + head_id = 0; + which = 3; // idx V + } + const auto cache_ptr = static_cast(params.cache_ptrs[which]); + const auto input_ptr = static_cast(params.input_ptrs[which]); + const auto cache_stride = params.cache_stride_bytes[which]; + const auto input_stride = params.input_stride_bytes[which]; + const auto src = pointer::offset(input_ptr, token_id * input_stride, head_id * kHeadBytes); + const auto dst = pointer::offset(cache_ptr, loc * cache_stride, head_id * kHeadBytes); + PDLWaitPrimary(); + vec_t chunk; + chunk.load(src, lane); + chunk.store(dst, lane); + PDLTriggerSecondary(); + } +}; + +template +__global__ void store_kv_index_kernel(const __grid_constant__ StoreKVIndexParams params) { + Trait::template forward(params); +} + +// idx_v / idx_v_cache may be dummies (== idx_k / idx_k_cache) when the layer +// has no index value; the caller signals "no V" via heads_per_token == 2*hkv+1 +// so the index-V branch is never taken. +template +void store_kv_index( + tvm::ffi::TensorView k, + tvm::ffi::TensorView v, + tvm::ffi::TensorView k_cache, + tvm::ffi::TensorView v_cache, + tvm::ffi::TensorView idx_k, + tvm::ffi::TensorView idx_k_cache, + tvm::ffi::TensorView idx_v, + tvm::ffi::TensorView idx_v_cache, + tvm::ffi::TensorView indices, + int64_t num_kv_heads, + int64_t heads_per_token) { + using namespace host; + auto B = SymbolicSize{"batch"}; + auto Mrow = SymbolicSize{"main_row"}; // num_kv_heads * head_dim + auto Drow = SymbolicSize{"idx_row"}; // head_dim + auto dtype = SymbolicDType{}; + auto indice_dtype = SymbolicDType{}; + auto device = SymbolicDevice{}; + device.set_options(); + + // All cache/source tensors share the same store dtype (fast-path precondition). + TensorMatcher({B, Mrow}) // + .with_strides({-1, 1}) + .with_dtype(dtype) + .with_device(device) + .verify(k) + .verify(v); + TensorMatcher({-1, Mrow}) // + .with_strides({-1, 1}) + .with_dtype(dtype) + .with_device(device) + .verify(k_cache) + .verify(v_cache); + TensorMatcher({B, Drow}) // + .with_strides({-1, 1}) + .with_dtype(dtype) + .with_device(device) + .verify(idx_k) + .verify(idx_v); + TensorMatcher({-1, Drow}) // + .with_strides({-1, 1}) + .with_dtype(dtype) + .with_device(device) + .verify(idx_k_cache) + .verify(idx_v_cache); + TensorMatcher({B}) // + .with_dtype(indice_dtype) + .with_device(device) + .verify(indices); + + const int64_t dsize = dtype_bytes(dtype.unwrap()); + RuntimeCheck(kHeadBytes == Drow.unwrap() * dsize); + RuntimeCheck(Mrow.unwrap() == num_kv_heads * Drow.unwrap()); + + const auto params = StoreKVIndexParams{ + .input_ptrs = {k.data_ptr(), v.data_ptr(), idx_k.data_ptr(), idx_v.data_ptr()}, + .cache_ptrs = {k_cache.data_ptr(), v_cache.data_ptr(), idx_k_cache.data_ptr(), idx_v_cache.data_ptr()}, + .indices = indices.data_ptr(), + .input_stride_bytes = + { + k.stride(0) * dsize, + v.stride(0) * dsize, + idx_k.stride(0) * dsize, + idx_v.stride(0) * dsize, + }, + .cache_stride_bytes = + { + k_cache.stride(0) * dsize, + v_cache.stride(0) * dsize, + idx_k_cache.stride(0) * dsize, + idx_v_cache.stride(0) * dsize, + }, + .num_k_heads = static_cast(num_kv_heads), + .num_total_heads = static_cast(heads_per_token), + .total_jobs = static_cast(B.unwrap() * heads_per_token), + }; + if (params.total_jobs == 0) return; + + using Trait = StoreTrait; + constexpr uint32_t kBlockSize = 256u; + const uint32_t num_blocks = div_ceil(params.total_jobs * Trait::kWorkerSize, kBlockSize); + const auto kernel = indice_dtype.is_type() // + ? store_kv_index_kernel + : store_kv_index_kernel; + LaunchKernel(num_blocks, kBlockSize, device.unwrap()) // + .enable_pdl(kUsePDL)(kernel, params); +} + +} // namespace diff --git a/python/sglang/jit_kernel/csrc/minimax/minimax_decode_topk.cuh b/python/sglang/jit_kernel/csrc/minimax/minimax_decode_topk.cuh new file mode 100644 index 000000000..ec2277edc --- /dev/null +++ b/python/sglang/jit_kernel/csrc/minimax/minimax_decode_topk.cuh @@ -0,0 +1,535 @@ +#include +#include + +#include +#include + +#include +#include + +#include +#include +#if defined(__HIP_PLATFORM_AMD__) +static constexpr unsigned long long kWarpSyncMask = 0xFFFFFFFFFFFFFFFFull; +#else +#include +static constexpr unsigned int kWarpSyncMask = 0xFFFFFFFFu; +#endif + +namespace { + +// Block top-k selection over a per-(head, batch) row of block scores, run by one +// CTA of TopKTrait::kCTASize threads. Picks the `topk` highest-scoring block ids +// (k_eff = min(topk, num_blocks)). Three size regimes, chosen by num_blocks: +// * <= kSmallThreshold : O(n^2) rank-by-compare (no radix). +// * <= kCTASize : 4-pass 8-bit radix, one element per thread in a reg. +// * <= kMaxNumBlocks : 4-pass 8-bit radix, kIters elements per thread cached +// in registers (row read from global exactly once); +// liveness is a uint32_t bitmask, selection is an +// in-loop scatter -- nothing is cached in shared memory. +// The trivial case num_blocks <= topk (every block selected) is handled by the +// kernels below, outside the Trait. +struct TopKTrait { + static constexpr uint32_t kMaxTopK = 32; + static constexpr uint32_t kCTASize = 512; + static constexpr uint32_t kNumWarps = kCTASize / device::kWarpThreads; + static constexpr uint32_t kMaxNumBlocks = 4096; // block topk + static constexpr uint32_t kSmallThreshold = 8 * kNumWarps; + static constexpr uint32_t kRadixBits = 8; + static constexpr uint32_t kRadixSize = 1 << kRadixBits; + static constexpr float kNegInf = -std::numeric_limits::infinity(); + + struct Smem { + uint32_t warp_sum[kNumWarps]; + alignas(128) uint32_t counter; + alignas(128) uint32_t counter_final; + alignas(128) uint32_t threshold_bin; + uint32_t equal_count; + uint32_t above_count; + uint32_t histogram[2][kRadixSize]; // 8 bit radix + float small_scores[kSmallThreshold]; // small (O(n^2)) path only + }; + + SGL_DEVICE static void forward( + const float* __restrict__ scores, + const uint32_t num_blocks, + int32_t* __restrict__ topk_out, + const uint32_t topk, + Smem* smem) { + using namespace device; + const auto tx = threadIdx.x; + __builtin_assume(tx < kCTASize); + const auto warp_id = tx / kWarpThreads; + const auto lane_id = tx % kWarpThreads; + + constexpr auto is_greater = [](float x, float y, int32_t delta) { + return (x > y) || ((x == y) && delta < 0); // lower block id wins + }; + constexpr auto warp_inclusive_sum = [](uint32_t lane_id, uint32_t val) { +#pragma unroll + for (uint32_t offset = 1; offset < 32; offset *= 2) { + uint32_t n = __shfl_up_sync(kWarpSyncMask, val, offset, 32); + if (lane_id >= offset) val += n; + } + return val; + }; + constexpr auto clip_nan = [](float x) { return x != x ? kNegInf : x; }; + constexpr auto score_to_key = [](float x) { + uint32_t b = __float_as_uint(x); + return (b & 0x80000000u) ? ~b : (b | 0x80000000u); + }; + // Find the radix bin holding the topk_remain-th largest of `total_active` + // elements currently counted in `histogram`. Writes threshold_bin (the bin), + // above_count (elements strictly above it), equal_count (elements in it). + const auto find_threshold = [&](uint32_t* histogram, uint32_t total_active, uint32_t topk_remain) { + using namespace device; + uint32_t hist_val = 0; + uint32_t warp_inc = 0; + if (tx < kRadixSize) { + hist_val = histogram[tx]; + warp_inc = warp_inclusive_sum(lane_id, hist_val); + if (lane_id == kWarpThreads - 1) smem->warp_sum[warp_id] = warp_inc; + } + __syncthreads(); + if (tx < kRadixSize) { + const auto inter = warp::reduce_sum(lane_id < warp_id ? smem->warp_sum[lane_id] : 0); + const auto prefix = inter + warp_inc; // count in bins [0, tx] + const auto above = total_active - prefix; // count in bins ABOVE tx + if (above < topk_remain && above + hist_val >= topk_remain) { + smem->threshold_bin = tx; + smem->above_count = above; + smem->equal_count = hist_val; + } + } + __syncthreads(); + }; + + if (num_blocks <= kSmallThreshold) { + // O(n^2) compare: each block's rank = #blocks that outrank it; the ones + // with rank < topk are selected (rank is its position in topk_out). + static_assert(kSmallThreshold <= kCTASize); + if (tx < num_blocks) smem->small_scores[tx] = clip_nan(scores[tx]); + __syncthreads(); + constexpr uint32_t kNumCandidates = kSmallThreshold / kNumWarps; + constexpr uint32_t kNumTargets = kSmallThreshold / kWarpThreads; + float candidates[kNumCandidates]; + float target[kNumTargets]; +#pragma unroll + for (uint32_t i = 0; i < kNumTargets; ++i) { + const auto idx = lane_id + i * kWarpThreads; + target[i] = (idx < num_blocks) ? smem->small_scores[idx] : kNegInf; + } +#pragma unroll + for (uint32_t i = 0; i < kNumCandidates; ++i) { + const auto idx = warp_id + i * kNumWarps; + candidates[i] = (idx < num_blocks) ? smem->small_scores[idx] : kNegInf; + } + +#pragma unroll + for (uint32_t i = 0; i < kNumCandidates; ++i) { + const int32_t idx = warp_id + i * kNumWarps; + if (idx >= static_cast(num_blocks)) break; + uint32_t rank = 0; +#pragma unroll + for (uint32_t j = 0; j < kNumTargets; ++j) { + const int32_t delta = lane_id + j * kWarpThreads - idx; + // partial rank = how many of this lane's targets outrank the candidate + rank += is_greater(target[j], candidates[i], delta); + } + // full rank = sum of the per-lane partial ranks across the warp + rank = warp::reduce_sum(rank); + if (rank < topk) topk_out[rank] = idx; + } + } else if (num_blocks <= kCTASize) { + // 4-pass 8-bit radix select, one element per thread held in a register. + bool active = tx < num_blocks; + const auto value = active ? clip_nan(scores[tx]) : kNegInf; + const auto key = score_to_key(value); + uint32_t topk_remain = topk; + uint32_t write_pos = topk; // sentinel: not selected + if (tx < kRadixSize) smem->histogram[0][tx] = 0; + if (tx == kRadixSize) smem->counter = smem->counter_final = 0; + __syncthreads(); + uint32_t total_active = num_blocks; + +#pragma unroll + for (int round = 0; round < 4; round++) { + const uint32_t shift = 24 - round * 8; + const uint32_t bin = (key >> shift) & 0xFFu; + const auto hist_idx = round % 2; + const auto histogram = smem->histogram[hist_idx]; + + if (active) atomicAdd(&histogram[bin], 1); + if (round < 3 && tx < kRadixSize) smem->histogram[hist_idx ^ 1][tx] = 0; + __syncthreads(); + + find_threshold(histogram, total_active, topk_remain); + + const auto threshold_bin = smem->threshold_bin; + const auto above_count = smem->above_count; + const auto equal_count = smem->equal_count; + + if (round < 3) total_active = equal_count; + topk_remain -= above_count; + + // scatter: above -> selected now; equal at the last pass -> keep the rest + if (active) { + if (bin > threshold_bin) { + write_pos = atomicAdd(&smem->counter, 1); + active = false; + } else if (bin < threshold_bin) { + active = false; + } else if (round == 3) { + write_pos = topk - topk_remain + atomicAdd(&smem->counter_final, 1); + } + // bin == threshold && round < 3: stay active for the next pass + } + + if (round == 3 || topk_remain == 0) break; + } + + if (write_pos < topk) topk_out[write_pos] = tx; + } else { + // num_blocks in (kCTASize, kMaxNumBlocks]: each thread caches its (up to + // kIters) slice of the row in registers -- read from global exactly ONCE -- + // then runs the same 4-pass radix select as the single-element path looped + // over those slots. Liveness is a uint32_t bitmask (bit i = slot i still in + // the running set), so there is no per-element flag array; selection is an + // in-loop scatter, so there is no per-element position array. Nothing is + // cached in shared memory beyond the histogram. + constexpr uint32_t kIters = kMaxNumBlocks / kCTASize; + static_assert(kIters <= 32, "active liveness is packed into a uint32_t"); + uint32_t key[kIters]; + uint32_t active = 0; +#pragma unroll + for (uint32_t i = 0; i < kIters; ++i) { + const uint32_t idx = i * kCTASize + tx; + if (idx < num_blocks) { + key[i] = score_to_key(clip_nan(scores[idx])); + active |= 1u << i; + } + } + if (tx < kRadixSize) smem->histogram[0][tx] = 0; + if (tx == kRadixSize) smem->counter = smem->counter_final = 0; + __syncthreads(); + + uint32_t topk_remain = topk; + uint32_t total_active = num_blocks; + +#pragma unroll + for (int round = 0; round < 4; ++round) { + const uint32_t shift = 24 - round * 8; + const auto hb = round & 1; + +#pragma unroll + for (uint32_t i = 0; i < kIters; ++i) + if (active & (1u << i)) atomicAdd(&smem->histogram[hb][(key[i] >> shift) & 0xFFu], 1); + if (round < 3 && tx < kRadixSize) smem->histogram[hb ^ 1][tx] = 0; + __syncthreads(); + + find_threshold(smem->histogram[hb], total_active, topk_remain); + const auto threshold_bin = smem->threshold_bin; + const auto above_count = smem->above_count; + const auto equal_count = smem->equal_count; + + if (round < 3) total_active = equal_count; + topk_remain -= above_count; + +#pragma unroll + for (uint32_t i = 0; i < kIters; ++i) { + if (active & (1u << i)) { + const uint32_t bin = (key[i] >> shift) & 0xFFu; + if (bin > threshold_bin) { + topk_out[atomicAdd(&smem->counter, 1)] = i * kCTASize + tx; + active &= ~(1u << i); + } else if (bin < threshold_bin) { + active &= ~(1u << i); + } else if (round == 3) { + const auto pos = topk - topk_remain + atomicAdd(&smem->counter_final, 1); + if (pos < topk) topk_out[pos] = i * kCTASize + tx; + } + // bin == threshold && round < 3: slot stays live for the next pass + } + } + + if (round == 3 || topk_remain == 0) break; + } + } + } +}; + +// ------------------------------------------------------------------------- +// Kernels: one CTA (kCTASize threads) per (head, batch) row. The trivial case +// num_blocks <= topk (every block selected) is special-judged here, outside the +// Trait; otherwise the Trait selects the top-k block ids. +// ------------------------------------------------------------------------- + +// Block-id output: topk_idx[h, b, 0:k_eff) = selected block ids (front-packed, +// unordered), [k_eff:topk) = -1. +template +__global__ void minimax_decode_topk_block_kernel( + const float* __restrict__ score, + const SeqLenT* __restrict__ seq_lens, + int32_t* __restrict__ topk_idx, + int batch, + int num_heads, + int max_seqblock, + int block_size, + int topk) { + const int b = blockIdx.x; // grid.x = batch + const int h = blockIdx.y; // grid.y = num_heads + const int tx = threadIdx.x; + + // seq_lens is from an earlier kernel; prefetch it (and the cheap setup) before + // waiting on the score producer so the prologue overlaps its tail (PDL). + const int64_t seq_len = static_cast(seq_lens[b]); + const int num_blocks_raw = static_cast((seq_len + block_size - 1) / block_size); + // Never scan past the materialized score columns. + const int num_blocks = num_blocks_raw < max_seqblock ? num_blocks_raw : max_seqblock; + int32_t* __restrict__ out = topk_idx + (static_cast(h) * batch + b) * topk; + device::PDLWaitPrimary(); + + if (num_blocks <= topk) { // trivial: identity, -1 padded + for (int i = tx; i < topk; i += TopKTrait::kCTASize) + out[i] = (i < num_blocks) ? i : -1; + return; + } + + const float* __restrict__ row = score + (static_cast(h) * batch + b) * max_seqblock; + __shared__ TopKTrait::Smem smem; + TopKTrait::forward(row, static_cast(num_blocks), out, static_cast(topk), &smem); +} + +// Page-table output: for each (batch b, kv-head h) pseudo-request emit the +// trtllm/fa3 page table -- selected blocks sorted ascending (so the final partial +// block's pages land last), each expanded to its ppb = block_size/page_size pages +// via req_to_token -- plus the effective KV length seq_lens_out. +// +// DP attention (num_kv_heads > 1): each kv head selects its OWN blocks, so the +// per-request page table can't be shared across heads. We flatten (b, h) into +// num_heads*batch pseudo-requests laid out batch-major (row = b*num_heads + h, +// matching q.view(bs, nkv, gqa, d).reshape(bs*nkv, gqa, d)). seq_lens / slot_ids / +// req_to_token are per-batch (head-independent: a token's cache slot is the same +// for every head). The page index is head-encoded (head-minor) as +// base_page*num_heads + h, which is exactly the page index into an HND cache +// [num_pages, nkv, page_size, D] reshaped to [num_pages*nkv, 1, page_size, D] (a +// free view when the cache is contiguous HND). num_heads == 1 (h == 0) reproduces +// the single-kv-head TP>=4 behavior (page index == base_page). +template +__global__ void minimax_decode_topk_page_table_kernel( + const float* __restrict__ score, + const SeqLenT* __restrict__ seq_lens, + const int32_t* __restrict__ req_to_token, + const int64_t* __restrict__ slot_ids, + int32_t* __restrict__ page_table, + int32_t* __restrict__ seq_lens_out, + int batch, + int num_heads, + int max_seqblock, + int block_size, + int topk, + int page_size, + int r2t_stride, + int max_kv_len, + int max_sparse_pages) { + const int b = blockIdx.x; // grid.x = batch + const int h = blockIdx.y; // grid.y = num_heads (kv head) + const int tx = threadIdx.x; + + // Prefetch seq_lens / slot_ids (from earlier kernels) and the cheap setup + // before waiting on the score producer, so the prologue overlaps its tail (PDL). + const int64_t seq_len = static_cast(seq_lens[b]); + const int num_blocks_raw = static_cast((seq_len + block_size - 1) / block_size); + const int num_blocks = num_blocks_raw < max_seqblock ? num_blocks_raw : max_seqblock; + const int ppb = block_size / page_size; + const int64_t out_row = static_cast(b) * num_heads + h; // flattened pseudo-request + int32_t* __restrict__ pt_row = page_table + out_row * max_sparse_pages; + const int64_t r2t_base = static_cast(slot_ids[b]) * r2t_stride; + device::PDLWaitPrimary(); + + if (num_blocks <= topk) { // trivial: every block selected, all tokens valid + if (tx == 0) seq_lens_out[out_row] = static_cast(seq_len); + // block id == ascending slot, so the partial final block's pages land last + const int total = num_blocks * ppb; + for (int e = tx; e < total; e += TopKTrait::kCTASize) { + const int slot = e / ppb; + const int pp = e % ppb; + int tok = slot * block_size + pp * page_size; + if (tok >= max_kv_len) tok = max_kv_len - 1; + pt_row[e] = req_to_token[r2t_base + tok] / page_size * num_heads + h; + } + return; + } + + const int k_eff = topk; // num_blocks > topk + const float* __restrict__ row = score + (static_cast(h) * batch + b) * max_seqblock; // head-major score + __shared__ TopKTrait::Smem smem; + __shared__ int32_t s_topk[TopKTrait::kMaxTopK]; + TopKTrait::forward(row, static_cast(num_blocks), s_topk, static_cast(topk), &smem); + __syncthreads(); // s_topk fully written before the transform reads it + + // Sort the selected block ids ascending (k_eff <= kMaxTopK is tiny) so the + // partial final block lands last, accumulating the effective KV length in the + // same pass: each selected block contributes min(block_size, seq_len - c*block) + // valid tokens (only the final block can be partial). + __shared__ int32_t s_sorted[TopKTrait::kMaxTopK]; + __shared__ int s_eff_kv; + if (tx == 0) s_eff_kv = 0; + __syncthreads(); + for (int slot = tx; slot < k_eff; slot += TopKTrait::kCTASize) { + const int32_t v = s_topk[slot]; + int rank = 0; + for (int j = 0; j < k_eff; ++j) + rank += (s_topk[j] < v); + s_sorted[rank] = v; + const int rem = static_cast(seq_len - static_cast(v) * block_size); + atomicAdd(&s_eff_kv, rem < block_size ? rem : block_size); + } + __syncthreads(); + if (tx == 0) seq_lens_out[out_row] = s_eff_kv; + + // Parallel page emit: one thread per output page. + const int total = k_eff * ppb; + for (int e = tx; e < total; e += TopKTrait::kCTASize) { + const int slot = e / ppb; + const int pp = e % ppb; + int tok = s_sorted[slot] * block_size + pp * page_size; + if (tok >= max_kv_len) tok = max_kv_len - 1; + pt_row[e] = req_to_token[r2t_base + tok] / page_size * num_heads + h; + } +} + +// ------------------------------------------------------------------------- +// Launchers +// ------------------------------------------------------------------------- +template +void minimax_decode_topk( + tvm::ffi::TensorView score, // [H, B, S] fp32 + tvm::ffi::TensorView seq_lens, // [B] int32/int64 + tvm::ffi::TensorView topk_idx, // [H, B, T] int32 + int64_t block_size, + int64_t topk) { + using namespace host; + + SymbolicSize H = {"num_heads"}; + SymbolicSize B = {"batch"}; + SymbolicSize S = {"max_seqblock"}; + SymbolicSize T = {"topk"}; + SymbolicDevice device_; + device_.set_options(); + + TensorMatcher({H, B, S}).with_dtype().with_device(device_).verify(score); + TensorMatcher({B}).with_dtype().with_device(device_).verify(seq_lens); + TensorMatcher({H, B, T}).with_dtype().with_device(device_).verify(topk_idx); + + const int num_heads = static_cast(H.unwrap()); + const int batch = static_cast(B.unwrap()); + const int max_seqblock = static_cast(S.unwrap()); + const int topk_i = static_cast(T.unwrap()); + const DLDevice device = device_.unwrap(); + + RuntimeCheck( + static_cast(topk_i) == topk, + "minimax_decode_topk: topk arg (", + topk, + ") must match topk_idx last dim (", + topk_i, + ")"); + RuntimeCheck(block_size > 0, "block_size must be > 0, got ", block_size); + if (batch == 0 || num_heads == 0) return; + + const dim3 grid(static_cast(batch), static_cast(num_heads)); + LaunchKernel(grid, TopKTrait::kCTASize, device, 0) + .enable_pdl(kUsePDL)( + minimax_decode_topk_block_kernel, + static_cast(score.data_ptr()), + static_cast(seq_lens.data_ptr()), + static_cast(topk_idx.data_ptr()), + batch, + num_heads, + max_seqblock, + static_cast(block_size), + topk_i); +} + +// Page-table variant: emit the per-(batch, kv-head) paged page table consumed by +// the dense backend (trtllm_mha / fa3) plus the effective KV length, instead of +// block ids. For DP attention (num_kv_heads > 1) each kv head selects its own +// blocks, so (b, h) pseudo-requests are flattened batch-major into the output +// (B*num_heads rows); num_heads == 1 is the TP>=4 single-kv-head case. The page +// index is head-encoded (head-minor) as base_page*num_heads + h -- the index into +// an HND cache [num_pages, nkv, ps, D] reshaped to [num_pages*nkv, 1, ps, D]. +// page_table and seq_lens_out are allocated by the caller. +template +void minimax_decode_topk_page_table( + tvm::ffi::TensorView score, // [H, B, S] fp32 (H = num_kv_heads) + tvm::ffi::TensorView seq_lens, // [B] int32/int64 + tvm::ffi::TensorView req_to_token, // [max_reqs, max_kv_len] int32 + tvm::ffi::TensorView slot_ids, // [B] int64 (req_pool_indices) + tvm::ffi::TensorView page_table, // [B*H, max_sparse_pages] int32 (out) + tvm::ffi::TensorView seq_lens_out, // [B*H] int32 (effective KV length, out) + int64_t block_size, + int64_t topk, + int64_t page_size) { + using namespace host; + + SymbolicSize H = {"num_heads"}; + SymbolicSize B = {"batch"}; + SymbolicSize S = {"max_seqblock"}; + SymbolicSize R = {"max_reqs"}; + SymbolicSize KV = {"max_kv_len"}; + SymbolicSize BH = {"batch_heads"}; + SymbolicSize P = {"max_sparse_pages"}; + SymbolicDevice device_; + device_.set_options(); + + TensorMatcher({H, B, S}).with_dtype().with_device(device_).verify(score); + TensorMatcher({B}).with_dtype().with_device(device_).verify(seq_lens); + TensorMatcher({R, KV}).with_dtype().with_device(device_).verify(req_to_token); + TensorMatcher({B}).with_dtype().with_device(device_).verify(slot_ids); + TensorMatcher({BH, P}).with_dtype().with_device(device_).verify(page_table); + TensorMatcher({BH}).with_dtype().with_device(device_).verify(seq_lens_out); + + const int num_heads = static_cast(H.unwrap()); + const int batch = static_cast(B.unwrap()); + const int max_seqblock = static_cast(S.unwrap()); + const int max_kv_len = static_cast(KV.unwrap()); + const int max_sparse_pages = static_cast(P.unwrap()); + const int r2t_stride = static_cast(req_to_token.stride(0)); + const DLDevice device = device_.unwrap(); + + RuntimeCheck( + BH.unwrap() == static_cast(batch) * num_heads, + "page_table rows (", + BH.unwrap(), + ") must equal batch*num_heads (", + static_cast(batch) * num_heads, + ")"); + RuntimeCheck( + block_size > 0 && page_size > 0 && block_size % page_size == 0, + "block_size must be a positive multiple of page_size"); + RuntimeCheck(topk <= static_cast(TopKTrait::kMaxTopK), "topk exceeds kMaxTopK for page-table mode"); + if (batch == 0 || num_heads == 0) return; + + const dim3 grid(static_cast(batch), static_cast(num_heads)); + LaunchKernel(grid, TopKTrait::kCTASize, device, 0) + .enable_pdl(kUsePDL)( + minimax_decode_topk_page_table_kernel, + static_cast(score.data_ptr()), + static_cast(seq_lens.data_ptr()), + static_cast(req_to_token.data_ptr()), + static_cast(slot_ids.data_ptr()), + static_cast(page_table.data_ptr()), + static_cast(seq_lens_out.data_ptr()), + batch, + num_heads, + max_seqblock, + static_cast(block_size), + static_cast(topk), + static_cast(page_size), + r2t_stride, + max_kv_len, + max_sparse_pages); +} + +} // namespace diff --git a/python/sglang/jit_kernel/csrc/minimax/per_token_quant_ue8m0.cuh b/python/sglang/jit_kernel/csrc/minimax/per_token_quant_ue8m0.cuh new file mode 100644 index 000000000..b82f59995 --- /dev/null +++ b/python/sglang/jit_kernel/csrc/minimax/per_token_quant_ue8m0.cuh @@ -0,0 +1,280 @@ +#include +#include + +#include +#include +#include +#include +#include + +#include + +#include +#include + +#include + +namespace { + +using deepseek_v4::fp8::cast_to_ue8m0; +using deepseek_v4::fp8::pack_fp8; + +// Per-token group quant to FP8-e4m3 with a fused UE8M0 scale. Each group of +// kGroupSize columns gets one UE8M0 exponent byte written contiguously in +// row-major order into ``x_sf`` (int32 [num_tokens, num_groups/4], 4 group +// bytes per int32). This is the deep_gemm "transform_sf" pack done inline in +// the quant, reusing the dsv4 ``cast_to_ue8m0``/``pack_fp8`` primitives -- it +// is byte-identical to ``per_token_group_quant_fp8(scale_ue8m0=True)`` followed +// by ``transform_sf_into_required_layout`` (both round via ceil(log2(absmax/ +// FP8_MAX))), but emits no separate transpose/pack kernel. +struct PerTokenQuantUe8m0Params { + const bf16_t* __restrict__ x; // [num_tokens, hidden] + fp8_e4m3_t* __restrict__ x_q; // [num_tokens, hidden] + int32_t* __restrict__ x_sf; // [num_tokens, num_groups/4]; written as bytes + uint32_t num_tokens; + uint32_t hidden; + uint32_t num_groups; // hidden / kGroupSize +}; + +template +__global__ __launch_bounds__(1024, 2) void // + per_token_quant_ue8m0_kernel(const PerTokenQuantUe8m0Params __grid_constant__ params) { + using namespace device; + constexpr uint32_t kVecElems = 8; // 8 bf16 = 16B load per thread + static_assert(kGroupSize % kVecElems == 0, "group_size must be a multiple of 8"); + constexpr uint32_t kThreadsPerGroup = kGroupSize / kVecElems; + using InputVec = AlignedVector; + using OutputVec = AlignedVector; + + const uint32_t token_id = blockIdx.x; + const uint32_t tid = threadIdx.x; + PDLWaitPrimary(); + + const auto token_in = params.x + static_cast(token_id) * params.hidden; + const auto token_out = params.x_q + static_cast(token_id) * params.hidden; + + InputVec in_vec; + in_vec.load(token_in, tid); + float local_max = 0.0f; + float vals[kVecElems]; +#pragma unroll + for (uint32_t i = 0; i < kVecElems / 2; ++i) { + const auto [v0, v1] = cast(in_vec[i]); + vals[2 * i + 0] = v0; + vals[2 * i + 1] = v1; + local_max = fmaxf(local_max, fmaxf(fabsf(v0), fabsf(v1))); + } + // Absmax across the kThreadsPerGroup threads that cover one group. + local_max = warp::reduce_max(local_max); + const float absmax = fmaxf(local_max, 1e-10f); + const float raw_scale = absmax / math::FP8_E4M3_MAX; + const uint32_t ue8m0_exp = cast_to_ue8m0(raw_scale); + const float inv_scale = __uint_as_float((127u + 127u - ue8m0_exp) << 23); + + OutputVec out_vec; +#pragma unroll + for (uint32_t i = 0; i < kVecElems / 2; ++i) { + out_vec[i] = pack_fp8(vals[2 * i + 0] * inv_scale, vals[2 * i + 1] * inv_scale); + } + out_vec.store(token_out, tid); + + const uint32_t group_id = tid / kThreadsPerGroup; + const uint32_t within_group_id = tid % kThreadsPerGroup; + if (within_group_id == 0 && group_id < params.num_groups) { + const uint32_t byte_off = token_id * params.num_groups + group_id; + reinterpret_cast(params.x_sf)[byte_off] = static_cast(ue8m0_exp); + } + PDLTriggerSecondary(); +} + +// Fused quant + scatter: like per_token_quant_ue8m0_kernel, but instead of +// writing the fp8/scale for the single source token, it scatters them straight +// into the permuted grouped-GEMM input -- replicating each token to its ``topk`` +// destination rows -- so the separate fill_gateup_input_triton_kernel launch (and +// the intermediate x_q/x_sf buffers) are eliminated. The fp8 value + UE8M0 scale +// are computed exactly once per token (identical to the non-fused kernel); only +// the stores differ. +struct PerTokenQuantUe8m0ScatterParams { + const bf16_t* __restrict__ x; // [num_tokens, hidden] + fp8_e4m3_t* __restrict__ gateup_input; // [E, m_max, hidden] + int32_t* __restrict__ gateup_input_scale; // [E, num_groups/4, m_max] int32 (MN-major), written as bytes + const int32_t* __restrict__ src2dst; // [num_tokens, topk] -> dst row = expert*m_max + slot + const int32_t* __restrict__ topk_ids; // [num_tokens, topk]; <0 = skip + uint32_t num_tokens; + uint32_t hidden; + uint32_t num_groups; // hidden / kGroupSize + uint32_t topk; + uint32_t m_max; +}; + +template +__global__ __launch_bounds__(1024, 2) void // + per_token_quant_ue8m0_scatter_kernel(const PerTokenQuantUe8m0ScatterParams __grid_constant__ params) { + using namespace device; + constexpr uint32_t kVecElems = 8; // 8 bf16 = 16B load per thread + static_assert(kGroupSize % kVecElems == 0, "group_size must be a multiple of 8"); + static_assert(kTopK <= kWarpThreads, "kTopK must fit in a warp for the lane-parallel scale write"); + constexpr uint32_t kThreadsPerGroup = kGroupSize / kVecElems; + using InputVec = AlignedVector; + using OutputVec = AlignedVector; + + const uint32_t token_id = blockIdx.x; + const uint32_t tid = threadIdx.x; + PDLWaitPrimary(); + + const auto token_in = params.x + static_cast(token_id) * params.hidden; + + InputVec in_vec; + in_vec.load(token_in, tid); + float local_max = 0.0f; + float vals[kVecElems]; +#pragma unroll + for (uint32_t i = 0; i < kVecElems / 2; ++i) { + const auto [v0, v1] = cast(in_vec[i]); + vals[2 * i + 0] = v0; + vals[2 * i + 1] = v1; + local_max = fmaxf(local_max, fmaxf(fabsf(v0), fabsf(v1))); + } + // Butterfly reduce: every thread of the group ends up with the group absmax + // (so all of them can write scale bytes below, no broadcast needed). + local_max = warp::reduce_max(local_max); + const float absmax = fmaxf(local_max, 1e-10f); + const float raw_scale = absmax / math::FP8_E4M3_MAX; + const uint32_t ue8m0_exp = cast_to_ue8m0(raw_scale); + const float inv_scale = __uint_as_float((127u + 127u - ue8m0_exp) << 23); + + OutputVec out_vec; +#pragma unroll + for (uint32_t i = 0; i < kVecElems / 2; ++i) { + out_vec[i] = pack_fp8(vals[2 * i + 0] * inv_scale, vals[2 * i + 1] * inv_scale); + } + + // Read this token's kTopK destinations once (fully unrolled). + const auto* src2dst_row = params.src2dst + static_cast(token_id) * kTopK; + const auto* topk_ids_row = params.topk_ids + static_cast(token_id) * kTopK; + int32_t dst_rows[kTopK]; +#pragma unroll + for (uint32_t i = 0; i < kTopK; ++i) { + dst_rows[i] = (topk_ids_row[i] >= 0) ? src2dst_row[i] : -1; + } + + const uint32_t group_id = tid / kThreadsPerGroup; + const uint32_t within_group = tid % kThreadsPerGroup; + const uint32_t c = group_id / 4u; // packed int32 index along the group axis + const uint32_t b = group_id % 4u; // byte within that int32 + const uint64_t scale_g4 = params.num_groups / 4u; + auto* scale_bytes = reinterpret_cast(params.gateup_input_scale); + + // 1) Output fp8: every thread replicates its 16B chunk to all kTopK dst rows. +#pragma unroll + for (uint32_t i = 0; i < kTopK; ++i) { + const int32_t dst = dst_rows[i]; + if (dst < 0) continue; + out_vec.store(params.gateup_input + static_cast(dst) * params.hidden, tid); + } + + // 2) Scale bytes: distribute the kTopK experts across the group's threads + // (each already holds the group exponent) so they write in parallel instead + // of one leader looping. lane `within_group` handles experts {within_group, + // within_group + kThreadsPerGroup, ...}; for kTopK <= kThreadsPerGroup that is + // exactly one expert per lane (no loop). +#pragma unroll + for (uint32_t i = within_group; i < kTopK; i += kThreadsPerGroup) { + const int32_t dst = dst_rows[i]; + if (dst < 0) continue; + const uint32_t expert = static_cast(dst) / params.m_max; + const uint32_t m = static_cast(dst) % params.m_max; + // int32 element [expert, c, m] of [E, G/4, m_max], byte b within it. + const uint64_t int32_index = + static_cast(expert) * scale_g4 * params.m_max + static_cast(c) * params.m_max + m; + scale_bytes[int32_index * 4u + b] = static_cast(ue8m0_exp); + } + + PDLTriggerSecondary(); +} + +template +void per_token_quant_ue8m0_scatter( + tvm::ffi::TensorView x, + tvm::ffi::TensorView gateup_input, + tvm::ffi::TensorView gateup_input_scale, + tvm::ffi::TensorView src2dst, + tvm::ffi::TensorView topk_ids, + int64_t topk, + int64_t m_max) { + using namespace host; + auto device = SymbolicDevice{}; + auto M = SymbolicSize{"num_tokens"}; + auto H = SymbolicSize{"hidden"}; + auto E = SymbolicSize{"num_experts"}; + auto MM = SymbolicSize{"m_max"}; + auto G4 = SymbolicSize{"num_groups_div_4"}; + device.set_options(); + TensorMatcher({M, H}).with_dtype().with_device(device).verify(x); + TensorMatcher({E, MM, H}).with_dtype().with_device(device).verify(gateup_input); + TensorMatcher({E, G4, MM}).with_dtype().with_device(device).verify(gateup_input_scale); + + const uint32_t num_tokens = static_cast(M.unwrap()); + const uint32_t hidden = static_cast(H.unwrap()); + RuntimeCheck(hidden % kGroupSize == 0, "hidden ", hidden, " not divisible by group_size ", kGroupSize); + const uint32_t num_groups = hidden / static_cast(kGroupSize); + RuntimeCheck(num_groups % 4 == 0, "num_groups must be a multiple of 4 for int32 packing"); + RuntimeCheck(static_cast(G4.unwrap()) * 4 == num_groups, "scale G/4 mismatch"); + RuntimeCheck(static_cast(MM.unwrap()) == static_cast(m_max), "m_max mismatch"); + const uint32_t threads = hidden / 8; // kVecElems + RuntimeCheck(threads <= 1024, "hidden/8 must be <= 1024, got ", threads); + RuntimeCheck(topk == kTopK, "topk does not match compiled template"); + + const auto params = PerTokenQuantUe8m0ScatterParams{ + .x = static_cast(x.data_ptr()), + .gateup_input = static_cast(gateup_input.data_ptr()), + .gateup_input_scale = static_cast(gateup_input_scale.data_ptr()), + .src2dst = static_cast(src2dst.data_ptr()), + .topk_ids = static_cast(topk_ids.data_ptr()), + .num_tokens = num_tokens, + .hidden = hidden, + .num_groups = num_groups, + .topk = static_cast(kTopK), + .m_max = static_cast(m_max), + }; + if (num_tokens == 0) return; + constexpr auto kernel = per_token_quant_ue8m0_scatter_kernel; + LaunchKernel(num_tokens, threads, device.unwrap()) // + .enable_pdl(kUsePDL)(kernel, params); +} + +template +void per_token_quant_ue8m0(tvm::ffi::TensorView x, tvm::ffi::TensorView x_q, tvm::ffi::TensorView x_sf) { + using namespace host; + auto device = SymbolicDevice{}; + auto M = SymbolicSize{"num_tokens"}; + auto H = SymbolicSize{"hidden"}; + auto G4 = SymbolicSize{"num_groups_div_4"}; + device.set_options(); + TensorMatcher({M, H}).with_dtype().with_device(device).verify(x); + TensorMatcher({M, H}).with_dtype().with_device(device).verify(x_q); + TensorMatcher({M, G4}).with_dtype().with_device(device).verify(x_sf); + + const uint32_t num_tokens = static_cast(M.unwrap()); + const uint32_t hidden = static_cast(H.unwrap()); + RuntimeCheck(hidden % kGroupSize == 0, "hidden ", hidden, " not divisible by group_size ", kGroupSize); + const uint32_t num_groups = hidden / static_cast(kGroupSize); + RuntimeCheck(static_cast(G4.unwrap()) * 4 == num_groups); + const uint32_t threads = hidden / 8; // kVecElems + RuntimeCheck(threads <= 1024, "hidden/8 must be <= 1024, got ", threads); + + const auto params = PerTokenQuantUe8m0Params{ + .x = static_cast(x.data_ptr()), + .x_q = static_cast(x_q.data_ptr()), + .x_sf = static_cast(x_sf.data_ptr()), + .num_tokens = num_tokens, + .hidden = hidden, + .num_groups = num_groups, + }; + if (num_tokens == 0) return; + constexpr auto kernel = per_token_quant_ue8m0_kernel; + LaunchKernel(num_tokens, threads, device.unwrap()) // + .enable_pdl(kUsePDL)(kernel, params); +} + +} // namespace diff --git a/python/sglang/jit_kernel/csrc/moe/moe_topk_sigmoid.cuh b/python/sglang/jit_kernel/csrc/moe/moe_topk_sigmoid.cuh new file mode 100644 index 000000000..d447278e8 --- /dev/null +++ b/python/sglang/jit_kernel/csrc/moe/moe_topk_sigmoid.cuh @@ -0,0 +1,544 @@ + +// Adapt from https://github.com/vllm-project/vllm/blob/v0.7.3/csrc/moe/topk_softmax_kernels.cu +// which is originally adapted from +// https://github.com/NVIDIA/TensorRT-LLM/blob/v0.7.1/cpp/tensorrt_llm/kernels/mixtureOfExperts/moe_kernels.cu +#pragma once + +#include + +#include + +#include +#include +#include +#include + +// CUDA 12.9+ deprecated cub::Max/Min in favour of cuda::maximum/minimum +#if CUDA_VERSION >= 12090 +#include +using MaxReduceOp = cuda::maximum<>; +using MinReduceOp = cuda::minimum<>; +#else +using MaxReduceOp = cub::Max; +using MinReduceOp = cub::Min; +#endif + +#include +#include +#include + +using tvm::ffi::TensorView; + +#ifndef MOE_TOPK_SIGMOID_WARP_SIZE +#define MOE_TOPK_SIGMOID_WARP_SIZE 32 +#endif + +namespace { + +static constexpr int WARP_SIZE = MOE_TOPK_SIGMOID_WARP_SIZE; + +#define MAX(a, b) ((a) > (b) ? (a) : (b)) +#define MIN(a, b) ((a) < (b) ? (a) : (b)) + +// --------------------------------------------------------------------------- +// Aligned array — avoids CUTLASS dependency; identical semantics. +// --------------------------------------------------------------------------- +template (sizeof(T) * N)> +class alignas(Alignment) AlignedArray { + T data[N]; +}; + +// --------------------------------------------------------------------------- +// Type conversion helper +// --------------------------------------------------------------------------- +template +__device__ float convert_to_float(T x) { + if constexpr (std::is_same_v) { + return __half2float(x); + } else if constexpr (std::is_same_v) { + return __bfloat162float(x); + } else { + return static_cast(x); + } +} + +// --------------------------------------------------------------------------- +// moeSigmoid — fallback sigmoid kernel (used for non-power-of-2 experts) +// --------------------------------------------------------------------------- +template +__launch_bounds__(TPB) __global__ + void moeSigmoid(const T* input, const bool* finished, float* output, const int num_cols) { + const int thread_row_offset = blockIdx.x * num_cols; + + if ((finished != nullptr) && finished[blockIdx.x]) { + return; + } + + for (int ii = threadIdx.x; ii < num_cols; ii += TPB) { + const int idx = thread_row_offset + ii; + float val = convert_to_float(input[idx]); + val = 1.0f / (1.0f + expf(-val)); + output[idx] = val; + } +} + +// --------------------------------------------------------------------------- +// moeTopK — fallback top-k kernel (used for non-power-of-2 experts) +// --------------------------------------------------------------------------- +template +__launch_bounds__(TPB) __global__ void moeTopK( + const float* inputs_after_sigmoid, + const bool* finished, + float* output, + int* indices, + const int num_experts, + const int k, + const int start_expert, + const int end_expert, + const bool renormalize, + const float* correction_bias, + double routed_scaling_factor, + int num_fused_shared_experts) { + using cub_kvp = cub::KeyValuePair; + using BlockReduce = cub::BlockReduce; + __shared__ typename BlockReduce::TempStorage tmpStorage; + + cub_kvp thread_kvp; + cub::ArgMax arg_max; + + const int block_row = blockIdx.x; + const int topk = k + num_fused_shared_experts; + + const bool row_is_active = finished ? !finished[block_row] : true; + const int thread_read_offset = blockIdx.x * num_experts; + float row_sum_for_renormalize = 0; + + for (int k_idx = 0; k_idx < k; ++k_idx) { + thread_kvp.key = 0; + thread_kvp.value = -1.f; + + cub_kvp inp_kvp; + for (int expert = threadIdx.x; expert < num_experts; expert += TPB) { + const int idx = thread_read_offset + expert; + inp_kvp.key = expert; + inp_kvp.value = inputs_after_sigmoid[idx]; + if (correction_bias != nullptr) { + inp_kvp.value += correction_bias[expert]; + } + + for (int prior_k = 0; prior_k < k_idx; ++prior_k) { + const int prior_winning_expert = indices[topk * block_row + prior_k]; + if (prior_winning_expert == expert) { + inp_kvp = thread_kvp; + } + } + + thread_kvp = arg_max(inp_kvp, thread_kvp); + } + + const cub_kvp result_kvp = BlockReduce(tmpStorage).Reduce(thread_kvp, arg_max); + if (threadIdx.x == 0) { + const int expert = result_kvp.key; + const bool node_uses_expert = expert >= start_expert && expert < end_expert; + const bool should_process_row = row_is_active && node_uses_expert; + + const int idx = topk * block_row + k_idx; + float val; + if (correction_bias != nullptr) { + val = inputs_after_sigmoid[thread_read_offset + expert]; + } else { + val = result_kvp.value; + } + output[idx] = val; + indices[idx] = should_process_row ? (expert - start_expert) : num_experts; + assert(indices[idx] >= 0); + row_sum_for_renormalize += val; + } + __syncthreads(); + } + + if (num_fused_shared_experts > 0 && threadIdx.x == 0) { + const int last_idx = topk * block_row + k; + if (renormalize) { + output[last_idx] = 1.0f; + } else { + output[last_idx] = row_sum_for_renormalize / routed_scaling_factor; + } + indices[last_idx] = num_experts; + } + if (renormalize && threadIdx.x == 0) { + float row_sum_for_renormalize_inv = routed_scaling_factor / (row_sum_for_renormalize + 1e-20f); + for (int k_idx = 0; k_idx < k; ++k_idx) { + const int idx = topk * block_row + k_idx; + output[idx] = output[idx] * row_sum_for_renormalize_inv; + } + } +} + +// --------------------------------------------------------------------------- +// topkGatingSigmoid — optimised kernel for power-of-2 expert counts +// --------------------------------------------------------------------------- +template +__launch_bounds__(WARPS_PER_CTA* WARP_SIZE) __global__ void topkGatingSigmoid( + const T* input, + const bool* finished, + float* output, + const int num_rows, + int* indices, + const int k, + const int start_expert, + const int end_expert, + const bool renormalize, + const float* correction_bias, + double routed_scaling_factor, + int num_fused_shared_experts) { + static_assert(VPT == (VPT & -VPT), "VPT must be power of 2"); + static_assert(NUM_EXPERTS == (NUM_EXPERTS & -NUM_EXPERTS), "NUM_EXPERTS must be power of 2"); + static_assert(BYTES_PER_LDG == (BYTES_PER_LDG & -BYTES_PER_LDG), "BYTES_PER_LDG must be power of 2"); + static_assert(BYTES_PER_LDG <= 16, "BYTES_PER_LDG must be leq 16"); + + static constexpr int ELTS_PER_LDG = BYTES_PER_LDG / sizeof(T); + static constexpr int ELTS_PER_ROW = NUM_EXPERTS; + static constexpr int THREADS_PER_ROW = ELTS_PER_ROW / VPT; + static constexpr int LDG_PER_THREAD = VPT / ELTS_PER_LDG; + + static_assert(VPT % ELTS_PER_LDG == 0, ""); + static_assert(WARP_SIZE % THREADS_PER_ROW == 0, ""); + static_assert(THREADS_PER_ROW == (THREADS_PER_ROW & -THREADS_PER_ROW), ""); + static_assert(THREADS_PER_ROW <= WARP_SIZE, ""); + + static constexpr int ELTS_PER_WARP = WARP_SIZE * VPT; + static constexpr int ROWS_PER_WARP = ELTS_PER_WARP / ELTS_PER_ROW; + static constexpr int ROWS_PER_CTA = WARPS_PER_CTA * ROWS_PER_WARP; + + static_assert(ELTS_PER_WARP % ELTS_PER_ROW == 0, ""); + + const int cta_base_row = blockIdx.x * ROWS_PER_CTA; + const int warp_base_row = cta_base_row + threadIdx.y * ROWS_PER_WARP; + const int thread_row_in_warp = threadIdx.x / THREADS_PER_ROW; + const int thread_row = warp_base_row + thread_row_in_warp; + const int topk = k + num_fused_shared_experts; + + if (thread_row >= num_rows) { + return; + } + const bool row_is_active = finished ? !finished[thread_row] : true; + + const T* thread_row_ptr = input + thread_row * ELTS_PER_ROW; + const int thread_group_idx = threadIdx.x % THREADS_PER_ROW; + const int first_elt_read_by_thread = thread_group_idx * ELTS_PER_LDG; + const T* thread_read_ptr = thread_row_ptr + first_elt_read_by_thread; + + using AccessType = AlignedArray; + + T row_chunk_temp[VPT]; + AccessType* row_chunk_vec_ptr = reinterpret_cast(&row_chunk_temp); + const AccessType* vec_thread_read_ptr = reinterpret_cast(thread_read_ptr); +#pragma unroll + for (int ii = 0; ii < LDG_PER_THREAD; ++ii) { + row_chunk_vec_ptr[ii] = vec_thread_read_ptr[ii * THREADS_PER_ROW]; + } + + float row_chunk[VPT]; +#pragma unroll + for (int ii = 0; ii < VPT; ++ii) { + float val = convert_to_float(row_chunk_temp[ii]); + val = 1.0f / (1.0f + expf(-val)); + if (correction_bias != nullptr) { + const int group_id = ii / ELTS_PER_LDG; + const int local_id = ii % ELTS_PER_LDG; + const int expert_idx = first_elt_read_by_thread + group_id * THREADS_PER_ROW * ELTS_PER_LDG + local_id; + val = val + correction_bias[expert_idx]; + } + row_chunk[ii] = val; + } + + int start_col = first_elt_read_by_thread; + static constexpr int COLS_PER_GROUP_LDG = ELTS_PER_LDG * THREADS_PER_ROW; + + float row_sum_for_renormalize = 0; + + for (int k_idx = 0; k_idx < k; ++k_idx) { + float max_val = row_chunk[0]; + int expert = start_col; +#pragma unroll + for (int ldg = 0, col = start_col; ldg < LDG_PER_THREAD; ++ldg, col += COLS_PER_GROUP_LDG) { +#pragma unroll + for (int ii = 0; ii < ELTS_PER_LDG; ++ii) { + float val = row_chunk[ldg * ELTS_PER_LDG + ii]; + if (val > max_val) { + max_val = val; + expert = col + ii; + } + } + } + +#pragma unroll + for (int mask = THREADS_PER_ROW / 2; mask > 0; mask /= 2) { + float other_max = __shfl_xor_sync(0xffffffff, max_val, mask, THREADS_PER_ROW); + int other_expert = __shfl_xor_sync(0xffffffff, expert, mask, THREADS_PER_ROW); + if (other_max > max_val || (other_max == max_val && other_expert < expert)) { + max_val = other_max; + expert = other_expert; + } + } + + if (thread_group_idx == 0) { + const bool node_uses_expert = expert >= start_expert && expert < end_expert; + const bool should_process_row = row_is_active && node_uses_expert; + + const int idx = topk * thread_row + k_idx; + float out_val; + if (correction_bias != nullptr) { + out_val = convert_to_float(thread_row_ptr[expert]); + out_val = 1.0f / (1.0f + expf(-out_val)); + } else { + out_val = max_val; + } + output[idx] = out_val; + indices[idx] = should_process_row ? (expert - start_expert) : NUM_EXPERTS; + row_sum_for_renormalize += out_val; + } + + if (k_idx + 1 < k) { + const int ldg_group_for_expert = expert / COLS_PER_GROUP_LDG; + const int thread_to_clear_in_group = (expert / ELTS_PER_LDG) % THREADS_PER_ROW; + if (thread_group_idx == thread_to_clear_in_group) { + const int offset_for_expert = expert % ELTS_PER_LDG; + row_chunk[ldg_group_for_expert * ELTS_PER_LDG + offset_for_expert] = -10000.f; + } + } + } + + if (num_fused_shared_experts > 0 && thread_group_idx == 0) { + const int last_idx = topk * thread_row + k; + if (renormalize) { + output[last_idx] = 1.0f; + } else { + output[last_idx] = row_sum_for_renormalize / routed_scaling_factor; + } + indices[last_idx] = NUM_EXPERTS; + } + if (renormalize && thread_group_idx == 0) { + float row_sum_for_renormalize_inv = routed_scaling_factor / (row_sum_for_renormalize + 1e-20f); +#pragma unroll + for (int k_idx = 0; k_idx < k; ++k_idx) { + const int idx = topk * thread_row + k_idx; + output[idx] = output[idx] * row_sum_for_renormalize_inv; + } + } +} + +// --------------------------------------------------------------------------- +// Compile-time constants helper +// --------------------------------------------------------------------------- +namespace detail { +template +struct TopkConstants { + static constexpr int ELTS_PER_LDG = BYTES_PER_LDG / sizeof(T); + static_assert(EXPERTS / (ELTS_PER_LDG * WARP_SIZE) == 0 || EXPERTS % (ELTS_PER_LDG * WARP_SIZE) == 0, ""); + static constexpr int VECs_PER_THREAD = MAX(1, EXPERTS / (ELTS_PER_LDG * WARP_SIZE)); + static constexpr int VPT = VECs_PER_THREAD * ELTS_PER_LDG; + static constexpr int THREADS_PER_ROW = EXPERTS / VPT; + static constexpr int ROWS_PER_WARP = WARP_SIZE / THREADS_PER_ROW; +}; +} // namespace detail + +// --------------------------------------------------------------------------- +// Per-expert-count launcher helper +// --------------------------------------------------------------------------- +template +void topkGatingSigmoidLauncherHelper( + const T* input, + const bool* finished, + float* output, + int* indices, + const int num_rows, + const int k, + const int start_expert, + const int end_expert, + const bool renormalize, + const float* correction_bias, + double routed_scaling_factor, + int num_fused_shared_experts, + cudaStream_t stream) { + static constexpr std::size_t MAX_BYTES_PER_LDG = 16; + static constexpr int BYTES_PER_LDG = MIN(MAX_BYTES_PER_LDG, sizeof(T) * EXPERTS); + using Constants = detail::TopkConstants; + static constexpr int VPT = Constants::VPT; + static constexpr int ROWS_PER_WARP = Constants::ROWS_PER_WARP; + const int num_warps = (num_rows + ROWS_PER_WARP - 1) / ROWS_PER_WARP; + const int num_blocks = (num_warps + WARPS_PER_TB - 1) / WARPS_PER_TB; + + dim3 block_dim(WARP_SIZE, WARPS_PER_TB); + topkGatingSigmoid<<>>( + input, + finished, + output, + num_rows, + indices, + k, + start_expert, + end_expert, + renormalize, + correction_bias, + routed_scaling_factor, + num_fused_shared_experts); +} + +// --------------------------------------------------------------------------- +// Dispatch macro — used inside topkGatingSigmoidKernelLauncher +// --------------------------------------------------------------------------- +#define LAUNCH_SIGMOID(TYPE, NUM_EXPERTS, WARPS_PER_TB) \ + topkGatingSigmoidLauncherHelper( \ + gating_output, \ + nullptr, \ + topk_weights, \ + topk_indices, \ + num_tokens, \ + topk - num_fused_shared_experts, \ + 0, \ + num_experts, \ + renormalize, \ + correction_bias, \ + routed_scaling_factor, \ + num_fused_shared_experts, \ + stream) + +// --------------------------------------------------------------------------- +// Main launcher: dispatches on num_experts +// --------------------------------------------------------------------------- +template +void topkGatingSigmoidKernelLauncher( + const T* gating_output, + float* topk_weights, + int* topk_indices, + float* sigmoid_workspace, + const int num_tokens, + const int num_experts, + const int topk, + const bool renormalize, + const float* correction_bias, + double routed_scaling_factor, + int num_fused_shared_experts, + cudaStream_t stream) { + static constexpr int WARPS_PER_TB = 4; + switch (num_experts) { + case 1: + LAUNCH_SIGMOID(T, 1, WARPS_PER_TB); + break; + case 2: + LAUNCH_SIGMOID(T, 2, WARPS_PER_TB); + break; + case 4: + LAUNCH_SIGMOID(T, 4, WARPS_PER_TB); + break; + case 8: + LAUNCH_SIGMOID(T, 8, WARPS_PER_TB); + break; + case 16: + LAUNCH_SIGMOID(T, 16, WARPS_PER_TB); + break; + case 32: + LAUNCH_SIGMOID(T, 32, WARPS_PER_TB); + break; + case 64: + LAUNCH_SIGMOID(T, 64, WARPS_PER_TB); + break; + case 128: + LAUNCH_SIGMOID(T, 128, WARPS_PER_TB); + break; + case 256: + LAUNCH_SIGMOID(T, 256, WARPS_PER_TB); + break; + default: { + // Fallback: non-power-of-2 or >256 experts + using namespace host; + RuntimeCheck( + sigmoid_workspace != nullptr, "sigmoid_workspace must be provided for num_experts that are not a power of 2"); + static constexpr int TPB = 256; + moeSigmoid<<>>(gating_output, nullptr, sigmoid_workspace, num_experts); + moeTopK<<>>( + sigmoid_workspace, + nullptr, + topk_weights, + topk_indices, + num_experts, + topk - num_fused_shared_experts, + 0, + num_experts, + renormalize, + correction_bias, + routed_scaling_factor, + num_fused_shared_experts); + } + } +} + +#undef LAUNCH_SIGMOID + +} // namespace + +// --------------------------------------------------------------------------- +// Host launcher (tvm-ffi interface) +// --------------------------------------------------------------------------- +template +void topk_sigmoid( + TensorView gating_output, + TensorView topk_weights, + TensorView topk_ids, + TensorView workspace, + bool renormalize, + tvm::ffi::Optional correction_bias, + double routed_scaling_factor, + int num_fused_shared_experts) { + using namespace host; + + // --- Input validation --- + RuntimeCheck(gating_output.dim() == 2, "gating_output must be 2-D"); + RuntimeCheck(topk_weights.dim() == 2, "topk_weights must be 2-D"); + RuntimeCheck(topk_ids.dim() == 2, "topk_ids must be 2-D"); + + const int64_t num_tokens = gating_output.shape()[0]; + const int64_t num_experts = gating_output.shape()[1]; + const int64_t topk = topk_weights.shape()[1]; + + RuntimeCheck( + topk_weights.shape()[0] == num_tokens && topk_ids.shape()[0] == num_tokens, + "topk_weights and topk_ids must have num_tokens rows"); + RuntimeCheck(topk_ids.shape()[1] == topk, "topk_ids second dim must match topk_weights"); + RuntimeCheck(topk <= num_experts, "topk must be <= num_experts"); + RuntimeCheck(num_fused_shared_experts <= 1, "num_fused_shared_experts must be <= 1"); + + // correction_bias validation + if (correction_bias.has_value()) { + const auto& bias = correction_bias.value(); + RuntimeCheck(bias.dim() == 1, "correction_bias must be 1-D"); + RuntimeCheck(bias.shape()[0] == num_experts, "correction_bias size must equal num_experts"); + RuntimeCheck( + bias.dtype().code == DLDataTypeCode::kDLFloat && bias.dtype().bits == 32, "correction_bias must be float32"); + } + + const T* gating_ptr = static_cast(gating_output.data_ptr()); + float* weights_ptr = static_cast(topk_weights.data_ptr()); + int* indices_ptr = static_cast(topk_ids.data_ptr()); + float* workspace_ptr = static_cast(workspace.data_ptr()); + const float* bias_ptr = + correction_bias.has_value() ? static_cast(correction_bias.value().data_ptr()) : nullptr; + + cudaStream_t stream = LaunchKernel::resolve_device(gating_output.device()); + + topkGatingSigmoidKernelLauncher( + gating_ptr, + weights_ptr, + indices_ptr, + workspace_ptr, + static_cast(num_tokens), + static_cast(num_experts), + static_cast(topk), + renormalize, + bias_ptr, + routed_scaling_factor, + num_fused_shared_experts, + stream); +} diff --git a/python/sglang/jit_kernel/minimax_decode_topk.py b/python/sglang/jit_kernel/minimax_decode_topk.py new file mode 100644 index 000000000..26940d66d --- /dev/null +++ b/python/sglang/jit_kernel/minimax_decode_topk.py @@ -0,0 +1,129 @@ +"""Block top-k over per-row block scores for the MiniMax-M3 sparse decode indexer. + +Drop-in replacement for the 2-stage split-K Triton topk +(``_topk_index_partial_kernel`` + ``_topk_index_merge_kernel``): given the +decode score tensor ``[num_heads, batch, max_seqblock]`` it produces +``topk_idx`` ``[num_heads, batch, topk]`` (0-indexed block ids, front-packed, +``-1`` padded), matching the consumer ``_gqa_share_sparse_decode_kernel``. + +``minimax_decode_topk_page_table`` additionally fuses the page-table transform +for the dense paged backend (trtllm_mha / fa3) and returns the page table plus +the per-query effective KV length. +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Tuple + +import torch + +from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args + +if TYPE_CHECKING: + from tvm_ffi.module import Module + + +@cache_once +def _jit_module(seq_dtype: torch.dtype) -> Module: + args = make_cpp_args(seq_dtype, True) # SeqLenT, kUsePDL + return load_jit( + "minimax_decode_topk", + *args, + cuda_files=["minimax/minimax_decode_topk.cuh"], + cuda_wrappers=[ + ("minimax_decode_topk", f"minimax_decode_topk<{args}>"), + ( + "minimax_decode_topk_page_table", + f"minimax_decode_topk_page_table<{args}>", + ), + ], + ) + + +def minimax_decode_topk( + score: torch.Tensor, # [num_heads, batch, max_seqblock] fp32 + seq_lens: torch.Tensor, # [batch] int32/int64 + block_size: int, + topk: int, + out: torch.Tensor | None = None, # [num_heads, batch, topk] int32 +) -> torch.Tensor: + assert score.is_cuda and score.dtype == torch.float32 and score.dim() == 3 + assert seq_lens.is_cuda and seq_lens.dim() == 1 + assert seq_lens.dtype in (torch.int32, torch.int64) + num_heads, batch, max_seqblock = score.shape + assert seq_lens.shape[0] == batch + + if not score.is_contiguous(): + score = score.contiguous() + if not seq_lens.is_contiguous(): + seq_lens = seq_lens.contiguous() + + if out is None: + out = torch.empty( + (num_heads, batch, topk), dtype=torch.int32, device=score.device + ) + else: + assert out.shape == (num_heads, batch, topk) + assert out.dtype == torch.int32 and out.is_cuda + assert out.is_contiguous() + + module = _jit_module(seq_lens.dtype) + module.minimax_decode_topk(score, seq_lens, out, int(block_size), int(topk)) + return out + + +def minimax_decode_topk_page_table( + score: torch.Tensor, # [num_kv_heads, batch, max_seqblock] fp32 + seq_lens: torch.Tensor, # [batch] int32/int64 + req_to_token: torch.Tensor, # [max_reqs, max_kv_len] int32 + slot_ids: torch.Tensor, # [batch] int64 (req_pool_indices) + block_size: int, + topk: int, + page_size: int, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Fused top-k + page-table transform: select the top-k blocks and emit the + per-(batch, kv-head) paged page table consumed by the dense backend + (trtllm_mha / fa3), instead of block ids, plus the per-pseudo-request effective + KV length (cache_seqlens, from the actual selection). Both are allocated here + and returned. + + For DP attention (num_kv_heads > 1) each kv head selects its own blocks, so + (batch, head) pseudo-requests are flattened batch-major into the outputs + ``[batch*num_kv_heads, topk*block_size/page_size]`` / ``[batch*num_kv_heads]`` + (matching ``q.view(bs, nkv, gqa, d).reshape(bs*nkv, gqa, d)``). The page index + is head-encoded (head-minor) as ``base_page*num_kv_heads + head`` -- the index + into an HND cache ``[num_pages, nkv, ps, D]`` reshaped to + ``[num_pages*nkv, 1, ps, D]``; num_kv_heads==1 reproduces the single-kv-head + TP>=4 behavior (page index == base_page).""" + assert score.is_cuda and score.dtype == torch.float32 and score.dim() == 3 + num_heads, batch, max_seqblock = score.shape + assert block_size % page_size == 0 + assert req_to_token.dtype == torch.int32 and slot_ids.dtype == torch.int64 + if not score.is_contiguous(): + score = score.contiguous() + if not seq_lens.is_contiguous(): + seq_lens = seq_lens.contiguous() + if not slot_ids.is_contiguous(): + slot_ids = slot_ids.contiguous() + + max_sparse_pages = topk * (block_size // page_size) + page_table = torch.empty( + (batch * num_heads, max_sparse_pages), dtype=torch.int32, device=score.device + ) + real_seq_lens = torch.empty( + (batch * num_heads,), dtype=torch.int32, device=score.device + ) + + module = _jit_module(seq_lens.dtype) + module.minimax_decode_topk_page_table( + score, + seq_lens, + req_to_token, + slot_ids, + page_table, + real_seq_lens, + int(block_size), + int(topk), + int(page_size), + ) + return page_table, real_seq_lens diff --git a/python/sglang/jit_kernel/minimax_m3/__init__.py b/python/sglang/jit_kernel/minimax_m3/__init__.py new file mode 100644 index 000000000..bbc930f95 --- /dev/null +++ b/python/sglang/jit_kernel/minimax_m3/__init__.py @@ -0,0 +1,25 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Fused Triton kernels for MiniMax-M3 on AMD ROCm (gfx94x / gfx95x). + +Model-scoped JIT kernels (mirrors ``jit_kernel/dsv4``), split by op type: + * ``rmsnorm`` -- fused fp32 Gemma RMSNorm (plain + fused-add-residual) + * ``swiglu`` -- fused fp32 SwiGLU-OAI (split layout) +""" + +from sglang.jit_kernel.minimax_m3.rmsnorm import ( + _num_warps, + gemma_fused_add_rmsnorm, + gemma_rmsnorm, +) +from sglang.jit_kernel.minimax_m3.swiglu import ( + swiglu_oai_mxfp8_quant, + swiglu_oai_split, +) + +__all__ = [ + "gemma_rmsnorm", + "gemma_fused_add_rmsnorm", + "swiglu_oai_split", + "swiglu_oai_mxfp8_quant", + "_num_warps", +] diff --git a/python/sglang/jit_kernel/minimax_m3/qk_norm_rope.py b/python/sglang/jit_kernel/minimax_m3/qk_norm_rope.py new file mode 100644 index 000000000..9136521a9 --- /dev/null +++ b/python/sglang/jit_kernel/minimax_m3/qk_norm_rope.py @@ -0,0 +1,659 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Fused MiniMax-M3 per-head Gemma Q/K RMSNorm + partial RoPE for ROCm.""" + +from typing import Tuple + +import torch +import triton +import triton.language as tl + + +@triton.jit +def _qk_gemma_rmsnorm_rope_kernel( + q_ptr, + k_ptr, + q_out_ptr, + k_out_ptr, + q_weight_ptr, + k_weight_ptr, + positions_ptr, + cos_sin_cache_ptr, + q_stride_m, + q_stride_d, + k_stride_m, + k_stride_d, + q_heads: tl.constexpr, + k_heads: tl.constexpr, + head_dim: tl.constexpr, + rotary_dim: tl.constexpr, + eps: tl.constexpr, + is_neox_style: tl.constexpr, + BLOCK_HD: tl.constexpr, +): + token_id = tl.program_id(0) + head_program = tl.program_id(1) + cols = tl.arange(0, BLOCK_HD) + mask = cols < head_dim + half_rotary: tl.constexpr = rotary_dim // 2 + + is_q = head_program < q_heads + head_id = tl.where(is_q, head_program, head_program - q_heads) + in_ptr = tl.where(is_q, q_ptr, k_ptr) + out_ptr = tl.where(is_q, q_out_ptr, k_out_ptr) + weight_ptr = tl.where(is_q, q_weight_ptr, k_weight_ptr) + stride_m = tl.where(is_q, q_stride_m, k_stride_m) + stride_d = tl.where(is_q, q_stride_d, k_stride_d) + n_heads = tl.where(is_q, q_heads, k_heads) + + base_in = in_ptr + token_id * stride_m + head_id * head_dim * stride_d + x = tl.load(base_in + cols * stride_d, mask=mask, other=0.0).to(tl.float32) + w = tl.load(weight_ptr + cols, mask=mask, other=0.0).to(tl.float32) + var = tl.sum(x * x, axis=0) / head_dim + rstd = tl.rsqrt(var + eps) + normed = x * rstd * (1.0 + w) + # Match the unfused path: GemmaRMSNorm writes bf16/fp16, then RoPE reads + # that rounded value in the following kernel. + normed = normed.to(q_out_ptr.dtype.element_ty).to(tl.float32) + + rotary_mask = cols < rotary_dim + if is_neox_style: + partner_cols = tl.where( + cols < half_rotary, cols + half_rotary, cols - half_rotary + ) + cos_cols = tl.where(cols < half_rotary, cols, cols - half_rotary) + sign = tl.where(cols < half_rotary, -1.0, 1.0) + else: + partner_cols = tl.where((cols % 2) == 0, cols + 1, cols - 1) + cos_cols = cols // 2 + sign = tl.where((cols % 2) == 0, -1.0, 1.0) + + partner_mask = partner_cols < head_dim + x_partner = tl.load( + base_in + partner_cols * stride_d, + mask=partner_mask, + other=0.0, + ).to(tl.float32) + w_partner = tl.load( + weight_ptr + partner_cols, + mask=partner_mask, + other=0.0, + ).to(tl.float32) + partner_normed = x_partner * rstd * (1.0 + w_partner) + partner_normed = partner_normed.to(q_out_ptr.dtype.element_ty).to(tl.float32) + + pos = tl.load(positions_ptr + token_id).to(tl.int64) + cos_sin_base = cos_sin_cache_ptr + pos * rotary_dim + cos = tl.load(cos_sin_base + cos_cols, mask=rotary_mask, other=1.0).to(tl.float32) + sin = tl.load( + cos_sin_base + half_rotary + cos_cols, + mask=rotary_mask, + other=0.0, + ).to(tl.float32) + rotated = normed * cos + sign * partner_normed * sin + out = tl.where(rotary_mask, rotated, normed) + + base_out = out_ptr + token_id * n_heads * head_dim + head_id * head_dim + tl.store(base_out + cols, out.to(out_ptr.dtype.element_ty), mask=mask) + + +def qk_gemma_rmsnorm_rope( + q: torch.Tensor, + k: torch.Tensor, + q_weight: torch.Tensor, + k_weight: torch.Tensor, + positions: torch.Tensor, + cos_sin_cache: torch.Tensor, + eps: float, + head_dim: int, + rotary_dim: int, + is_neox_style: bool, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Return normalized+rotated Q/K tensors with the same shapes as ``q``/``k``.""" + assert q.dim() == 2 and k.dim() == 2 + assert positions.dim() == 1 + assert q.shape[0] == k.shape[0] == positions.shape[0] + assert q.shape[1] % head_dim == 0 + assert k.shape[1] % head_dim == 0 + assert rotary_dim <= head_dim and rotary_dim % 2 == 0 + + q_heads = q.shape[1] // head_dim + k_heads = k.shape[1] // head_dim + q_out = torch.empty(q.shape, dtype=q.dtype, device=q.device) + k_out = torch.empty(k.shape, dtype=k.dtype, device=k.device) + block_hd = triton.next_power_of_2(head_dim) + + _qk_gemma_rmsnorm_rope_kernel[(q.shape[0], q_heads + k_heads)]( + q, + k, + q_out, + k_out, + q_weight, + k_weight, + positions, + cos_sin_cache, + q.stride(0), + q.stride(1), + k.stride(0), + k.stride(1), + q_heads, + k_heads, + head_dim, + rotary_dim, + eps, + is_neox_style, + BLOCK_HD=block_hd, + num_warps=4, + ) + return q_out, k_out + + +@triton.jit +def _sparse_qk_index_gemma_rmsnorm_rope_kernel( + q_ptr, + k_ptr, + idx_q_ptr, + idx_k_ptr, + q_out_ptr, + k_out_ptr, + idx_q_out_ptr, + idx_k_out_ptr, + q_weight_ptr, + k_weight_ptr, + idx_q_weight_ptr, + idx_k_weight_ptr, + positions_ptr, + cos_sin_cache_ptr, + q_stride_m, + q_stride_d, + k_stride_m, + k_stride_d, + idx_q_stride_m, + idx_q_stride_d, + idx_k_stride_m, + idx_k_stride_d, + q_heads: tl.constexpr, + k_heads: tl.constexpr, + idx_q_heads: tl.constexpr, + head_dim: tl.constexpr, + rotary_dim: tl.constexpr, + eps: tl.constexpr, + is_neox_style: tl.constexpr, + BLOCK_HD: tl.constexpr, +): + token_id = tl.program_id(0) + head_program = tl.program_id(1) + cols = tl.arange(0, BLOCK_HD) + mask = cols < head_dim + half_rotary: tl.constexpr = rotary_dim // 2 + + main_heads: tl.constexpr = q_heads + k_heads + idx_k_program: tl.constexpr = q_heads + k_heads + idx_q_heads + + is_q = head_program < q_heads + is_k = (head_program >= q_heads) & (head_program < main_heads) + is_idx_q = (head_program >= main_heads) & (head_program < idx_k_program) + + head_id = tl.where( + is_q, + head_program, + tl.where( + is_k, + head_program - q_heads, + tl.where(is_idx_q, head_program - main_heads, 0), + ), + ) + + in_ptr = tl.where( + is_q, + q_ptr, + tl.where(is_k, k_ptr, tl.where(is_idx_q, idx_q_ptr, idx_k_ptr)), + ) + out_ptr = tl.where( + is_q, + q_out_ptr, + tl.where(is_k, k_out_ptr, tl.where(is_idx_q, idx_q_out_ptr, idx_k_out_ptr)), + ) + weight_ptr = tl.where( + is_q, + q_weight_ptr, + tl.where( + is_k, + k_weight_ptr, + tl.where(is_idx_q, idx_q_weight_ptr, idx_k_weight_ptr), + ), + ) + stride_m = tl.where( + is_q, + q_stride_m, + tl.where( + is_k, + k_stride_m, + tl.where(is_idx_q, idx_q_stride_m, idx_k_stride_m), + ), + ) + stride_d = tl.where( + is_q, + q_stride_d, + tl.where( + is_k, + k_stride_d, + tl.where(is_idx_q, idx_q_stride_d, idx_k_stride_d), + ), + ) + out_heads = tl.where( + is_q, + q_heads, + tl.where(is_k, k_heads, tl.where(is_idx_q, idx_q_heads, 1)), + ) + + base_in = in_ptr + token_id * stride_m + head_id * head_dim * stride_d + x = tl.load(base_in + cols * stride_d, mask=mask, other=0.0).to(tl.float32) + w = tl.load(weight_ptr + cols, mask=mask, other=0.0).to(tl.float32) + var = tl.sum(x * x, axis=0) / head_dim + rstd = tl.rsqrt(var + eps) + normed = x * rstd * (1.0 + w) + normed = normed.to(q_out_ptr.dtype.element_ty).to(tl.float32) + + rotary_mask = cols < rotary_dim + if is_neox_style: + partner_cols = tl.where( + cols < half_rotary, cols + half_rotary, cols - half_rotary + ) + cos_cols = tl.where(cols < half_rotary, cols, cols - half_rotary) + sign = tl.where(cols < half_rotary, -1.0, 1.0) + else: + partner_cols = tl.where((cols % 2) == 0, cols + 1, cols - 1) + cos_cols = cols // 2 + sign = tl.where((cols % 2) == 0, -1.0, 1.0) + + partner_mask = partner_cols < head_dim + x_partner = tl.load( + base_in + partner_cols * stride_d, + mask=partner_mask, + other=0.0, + ).to(tl.float32) + w_partner = tl.load( + weight_ptr + partner_cols, + mask=partner_mask, + other=0.0, + ).to(tl.float32) + partner_normed = x_partner * rstd * (1.0 + w_partner) + partner_normed = partner_normed.to(q_out_ptr.dtype.element_ty).to(tl.float32) + + pos = tl.load(positions_ptr + token_id).to(tl.int64) + cos_sin_base = cos_sin_cache_ptr + pos * rotary_dim + cos = tl.load(cos_sin_base + cos_cols, mask=rotary_mask, other=1.0).to(tl.float32) + sin = tl.load( + cos_sin_base + half_rotary + cos_cols, + mask=rotary_mask, + other=0.0, + ).to(tl.float32) + rotated = normed * cos + sign * partner_normed * sin + out = tl.where(rotary_mask, rotated, normed) + + base_out = out_ptr + token_id * out_heads * head_dim + head_id * head_dim + tl.store(base_out + cols, out.to(q_out_ptr.dtype.element_ty), mask=mask) + + +def sparse_qk_index_gemma_rmsnorm_rope( + q: torch.Tensor, + k: torch.Tensor, + idx_q: torch.Tensor, + idx_k: torch.Tensor, + q_weight: torch.Tensor, + k_weight: torch.Tensor, + idx_q_weight: torch.Tensor, + idx_k_weight: torch.Tensor, + positions: torch.Tensor, + cos_sin_cache: torch.Tensor, + eps: float, + head_dim: int, + rotary_dim: int, + is_neox_style: bool, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """Fuse main and sparse-index Gemma Q/K RMSNorm + RoPE into one launch.""" + assert q.dim() == k.dim() == idx_q.dim() == idx_k.dim() == 2 + assert positions.dim() == 1 + assert q.shape[0] == k.shape[0] == idx_q.shape[0] == idx_k.shape[0] + assert q.shape[0] == positions.shape[0] + assert q.shape[1] % head_dim == 0 + assert k.shape[1] % head_dim == 0 + assert idx_q.shape[1] % head_dim == 0 + assert idx_k.shape[1] == head_dim + assert rotary_dim <= head_dim and rotary_dim % 2 == 0 + + q_heads = q.shape[1] // head_dim + k_heads = k.shape[1] // head_dim + idx_q_heads = idx_q.shape[1] // head_dim + q_out = torch.empty(q.shape, dtype=q.dtype, device=q.device) + k_out = torch.empty(k.shape, dtype=k.dtype, device=k.device) + idx_q_out = torch.empty(idx_q.shape, dtype=idx_q.dtype, device=idx_q.device) + idx_k_out = torch.empty(idx_k.shape, dtype=idx_k.dtype, device=idx_k.device) + block_hd = triton.next_power_of_2(head_dim) + + _sparse_qk_index_gemma_rmsnorm_rope_kernel[ + (q.shape[0], q_heads + k_heads + idx_q_heads + 1) + ]( + q, + k, + idx_q, + idx_k, + q_out, + k_out, + idx_q_out, + idx_k_out, + q_weight, + k_weight, + idx_q_weight, + idx_k_weight, + positions, + cos_sin_cache, + q.stride(0), + q.stride(1), + k.stride(0), + k.stride(1), + idx_q.stride(0), + idx_q.stride(1), + idx_k.stride(0), + idx_k.stride(1), + q_heads, + k_heads, + idx_q_heads, + head_dim, + rotary_dim, + eps, + is_neox_style, + BLOCK_HD=block_hd, + num_warps=4, + ) + return q_out, k_out, idx_q_out, idx_k_out + + +@triton.jit +def _sparse_qk_index_gemma_rmsnorm_rope_cache_kernel( + q_ptr, + k_ptr, + v_ptr, + idx_q_ptr, + idx_k_ptr, + q_out_ptr, + k_out_ptr, + idx_q_out_ptr, + idx_k_out_ptr, + k_cache_ptr, + v_cache_ptr, + idx_k_cache_ptr, + loc_ptr, + q_weight_ptr, + k_weight_ptr, + idx_q_weight_ptr, + idx_k_weight_ptr, + positions_ptr, + cos_sin_cache_ptr, + q_stride_m, + q_stride_d, + k_stride_m, + k_stride_d, + v_stride_m, + v_stride_d, + idx_q_stride_m, + idx_q_stride_d, + idx_k_stride_m, + idx_k_stride_d, + k_cache_stride_s, + k_cache_stride_h, + k_cache_stride_d, + v_cache_stride_s, + v_cache_stride_h, + v_cache_stride_d, + idx_k_cache_stride_s, + idx_k_cache_stride_h, + idx_k_cache_stride_d, + q_heads: tl.constexpr, + k_heads: tl.constexpr, + idx_q_heads: tl.constexpr, + head_dim: tl.constexpr, + rotary_dim: tl.constexpr, + eps: tl.constexpr, + is_neox_style: tl.constexpr, + BLOCK_HD: tl.constexpr, +): + token_id = tl.program_id(0) + head_program = tl.program_id(1) + cols = tl.arange(0, BLOCK_HD) + mask = cols < head_dim + half_rotary: tl.constexpr = rotary_dim // 2 + + main_heads: tl.constexpr = q_heads + k_heads + idx_k_program: tl.constexpr = q_heads + k_heads + idx_q_heads + + is_q = head_program < q_heads + is_k = (head_program >= q_heads) & (head_program < main_heads) + is_idx_q = (head_program >= main_heads) & (head_program < idx_k_program) + + head_id = tl.where( + is_q, + head_program, + tl.where( + is_k, + head_program - q_heads, + tl.where(is_idx_q, head_program - main_heads, 0), + ), + ) + + in_ptr = tl.where( + is_q, + q_ptr, + tl.where(is_k, k_ptr, tl.where(is_idx_q, idx_q_ptr, idx_k_ptr)), + ) + out_ptr = tl.where( + is_q, + q_out_ptr, + tl.where(is_k, k_out_ptr, tl.where(is_idx_q, idx_q_out_ptr, idx_k_out_ptr)), + ) + weight_ptr = tl.where( + is_q, + q_weight_ptr, + tl.where( + is_k, + k_weight_ptr, + tl.where(is_idx_q, idx_q_weight_ptr, idx_k_weight_ptr), + ), + ) + stride_m = tl.where( + is_q, + q_stride_m, + tl.where( + is_k, + k_stride_m, + tl.where(is_idx_q, idx_q_stride_m, idx_k_stride_m), + ), + ) + stride_d = tl.where( + is_q, + q_stride_d, + tl.where( + is_k, + k_stride_d, + tl.where(is_idx_q, idx_q_stride_d, idx_k_stride_d), + ), + ) + out_heads = tl.where( + is_q, + q_heads, + tl.where(is_k, k_heads, tl.where(is_idx_q, idx_q_heads, 1)), + ) + + base_in = in_ptr + token_id * stride_m + head_id * head_dim * stride_d + x = tl.load(base_in + cols * stride_d, mask=mask, other=0.0).to(tl.float32) + w = tl.load(weight_ptr + cols, mask=mask, other=0.0).to(tl.float32) + var = tl.sum(x * x, axis=0) / head_dim + rstd = tl.rsqrt(var + eps) + normed = x * rstd * (1.0 + w) + normed = normed.to(q_out_ptr.dtype.element_ty).to(tl.float32) + + rotary_mask = cols < rotary_dim + if is_neox_style: + partner_cols = tl.where( + cols < half_rotary, cols + half_rotary, cols - half_rotary + ) + cos_cols = tl.where(cols < half_rotary, cols, cols - half_rotary) + sign = tl.where(cols < half_rotary, -1.0, 1.0) + else: + partner_cols = tl.where((cols % 2) == 0, cols + 1, cols - 1) + cos_cols = cols // 2 + sign = tl.where((cols % 2) == 0, -1.0, 1.0) + + partner_mask = partner_cols < head_dim + x_partner = tl.load( + base_in + partner_cols * stride_d, + mask=partner_mask, + other=0.0, + ).to(tl.float32) + w_partner = tl.load( + weight_ptr + partner_cols, + mask=partner_mask, + other=0.0, + ).to(tl.float32) + partner_normed = x_partner * rstd * (1.0 + w_partner) + partner_normed = partner_normed.to(q_out_ptr.dtype.element_ty).to(tl.float32) + + pos = tl.load(positions_ptr + token_id).to(tl.int64) + cos_sin_base = cos_sin_cache_ptr + pos * rotary_dim + cos = tl.load(cos_sin_base + cos_cols, mask=rotary_mask, other=1.0).to(tl.float32) + sin = tl.load( + cos_sin_base + half_rotary + cos_cols, + mask=rotary_mask, + other=0.0, + ).to(tl.float32) + rotated = normed * cos + sign * partner_normed * sin + out = tl.where(rotary_mask, rotated, normed) + out_typed = out.to(q_out_ptr.dtype.element_ty) + + base_out = out_ptr + token_id * out_heads * head_dim + head_id * head_dim + tl.store(base_out + cols, out_typed, mask=mask) + + loc = tl.load(loc_ptr + token_id) + cache_k_base = ( + k_cache_ptr + + loc * k_cache_stride_s + + head_id * k_cache_stride_h + + cols * k_cache_stride_d + ) + tl.store(cache_k_base, out_typed, mask=mask & is_k) + + v_base = v_ptr + token_id * v_stride_m + head_id * head_dim * v_stride_d + v_val = tl.load(v_base + cols * v_stride_d, mask=mask & is_k, other=0.0) + cache_v_base = ( + v_cache_ptr + + loc * v_cache_stride_s + + head_id * v_cache_stride_h + + cols * v_cache_stride_d + ) + tl.store(cache_v_base, v_val, mask=mask & is_k) + + is_idx_k = head_program == idx_k_program + idx_cache_base = ( + idx_k_cache_ptr + loc * idx_k_cache_stride_s + cols * idx_k_cache_stride_d + ) + tl.store(idx_cache_base, out_typed, mask=mask & is_idx_k) + + +def sparse_qk_index_gemma_rmsnorm_rope_cache( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + idx_q: torch.Tensor, + idx_k: torch.Tensor, + k_cache: torch.Tensor, + v_cache: torch.Tensor, + idx_k_cache: torch.Tensor, + out_cache_loc: torch.Tensor, + q_weight: torch.Tensor, + k_weight: torch.Tensor, + idx_q_weight: torch.Tensor, + idx_k_weight: torch.Tensor, + positions: torch.Tensor, + cos_sin_cache: torch.Tensor, + eps: float, + head_dim: int, + rotary_dim: int, + is_neox_style: bool, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """Fuse sparse Q/K/index norm+RoPE with main KV and index-K cache stores.""" + assert q.dim() == k.dim() == v.dim() == idx_q.dim() == idx_k.dim() == 2 + assert k_cache.dim() == v_cache.dim() == idx_k_cache.dim() == 3 + assert out_cache_loc.dim() == positions.dim() == 1 + assert q.shape[0] == k.shape[0] == v.shape[0] == idx_q.shape[0] == idx_k.shape[0] + assert q.shape[0] == positions.shape[0] == out_cache_loc.shape[0] + assert q.shape[1] % head_dim == 0 + assert k.shape[1] % head_dim == 0 + assert v.shape[1] == k.shape[1] + assert idx_q.shape[1] % head_dim == 0 + assert idx_k.shape[1] == head_dim + assert rotary_dim <= head_dim and rotary_dim % 2 == 0 + + q_heads = q.shape[1] // head_dim + k_heads = k.shape[1] // head_dim + idx_q_heads = idx_q.shape[1] // head_dim + assert k_cache.shape[1] == v_cache.shape[1] == k_heads + assert idx_k_cache.shape[1] == 1 + + q_out = torch.empty(q.shape, dtype=q.dtype, device=q.device) + k_out = torch.empty(k.shape, dtype=k.dtype, device=k.device) + idx_q_out = torch.empty(idx_q.shape, dtype=idx_q.dtype, device=idx_q.device) + idx_k_out = torch.empty(idx_k.shape, dtype=idx_k.dtype, device=idx_k.device) + block_hd = triton.next_power_of_2(head_dim) + + _sparse_qk_index_gemma_rmsnorm_rope_cache_kernel[ + (q.shape[0], q_heads + k_heads + idx_q_heads + 1) + ]( + q, + k, + v, + idx_q, + idx_k, + q_out, + k_out, + idx_q_out, + idx_k_out, + k_cache, + v_cache, + idx_k_cache, + out_cache_loc, + q_weight, + k_weight, + idx_q_weight, + idx_k_weight, + positions, + cos_sin_cache, + q.stride(0), + q.stride(1), + k.stride(0), + k.stride(1), + v.stride(0), + v.stride(1), + idx_q.stride(0), + idx_q.stride(1), + idx_k.stride(0), + idx_k.stride(1), + k_cache.stride(0), + k_cache.stride(1), + k_cache.stride(2), + v_cache.stride(0), + v_cache.stride(1), + v_cache.stride(2), + idx_k_cache.stride(0), + idx_k_cache.stride(1), + idx_k_cache.stride(2), + q_heads, + k_heads, + idx_q_heads, + head_dim, + rotary_dim, + eps, + is_neox_style, + BLOCK_HD=block_hd, + num_warps=4, + ) + return q_out, k_out, idx_q_out, idx_k_out diff --git a/python/sglang/jit_kernel/minimax_m3/rmsnorm.py b/python/sglang/jit_kernel/minimax_m3/rmsnorm.py new file mode 100644 index 000000000..790c3141a --- /dev/null +++ b/python/sglang/jit_kernel/minimax_m3/rmsnorm.py @@ -0,0 +1,148 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Fused Gemma RMSNorm Triton kernels for MiniMax-M3 on AMD ROCm. + +Gemma RMSNorm = ``normalize(x) * (1 + weight)``, computed in a single fp32 pass. +On ROCm with AITER, ``GemmaRMSNorm.forward_hip`` otherwise falls back to a +~8-op PyTorch sequence: ``sgl_kernel``'s Gemma kernels are CUDA-only, and +AITER's ``rmsnorm2d_fwd`` requires weight.dtype == activation.dtype (fp32 +weight + bf16 activation silently corrupts on gfx950). These kernels read +strided inputs, so they serve both the full-hidden norms and the per-head +q/k/index norms (non-contiguous ``qkv.split`` views). +""" + +import torch +import triton +import triton.language as tl + + +@triton.jit +def _gemma_rmsnorm_kernel( + x_ptr, + w_ptr, + out_ptr, + n_cols, + stride_row, + stride_col, + eps, + BLOCK_N: tl.constexpr, +): + row = tl.program_id(0) + cols = tl.arange(0, BLOCK_N) + mask = cols < n_cols + x = tl.load(x_ptr + row * stride_row + cols * stride_col, mask=mask, other=0.0).to( + tl.float32 + ) + var = tl.sum(x * x, axis=0) / n_cols + rstd = 1.0 / tl.sqrt(var + eps) + w = tl.load(w_ptr + cols, mask=mask, other=0.0).to(tl.float32) + out = x * rstd * (1.0 + w) + tl.store( + out_ptr + row * n_cols + cols, + out.to(out_ptr.dtype.element_ty), + mask=mask, + ) + + +@triton.jit +def _gemma_fused_add_rmsnorm_kernel( + x_ptr, + res_ptr, + w_ptr, + out_ptr, + res_out_ptr, + n_cols, + stride_xrow, + stride_xcol, + stride_rrow, + stride_rcol, + eps, + BLOCK_N: tl.constexpr, +): + row = tl.program_id(0) + cols = tl.arange(0, BLOCK_N) + mask = cols < n_cols + x = tl.load( + x_ptr + row * stride_xrow + cols * stride_xcol, mask=mask, other=0.0 + ).to(tl.float32) + r = tl.load( + res_ptr + row * stride_rrow + cols * stride_rcol, mask=mask, other=0.0 + ).to(tl.float32) + s = x + r + # residual_out is the pre-norm sum (consumed by the next layer's add). + tl.store( + res_out_ptr + row * n_cols + cols, + s.to(res_out_ptr.dtype.element_ty), + mask=mask, + ) + var = tl.sum(s * s, axis=0) / n_cols + rstd = 1.0 / tl.sqrt(var + eps) + w = tl.load(w_ptr + cols, mask=mask, other=0.0).to(tl.float32) + out = s * rstd * (1.0 + w) + tl.store( + out_ptr + row * n_cols + cols, + out.to(out_ptr.dtype.element_ty), + mask=mask, + ) + + +def _num_warps(block_n: int) -> int: + if block_n >= 4096: + return 16 + if block_n >= 1024: + return 8 + return 4 + + +def gemma_rmsnorm(x: torch.Tensor, weight: torch.Tensor, eps: float) -> torch.Tensor: + """Gemma RMSNorm = normalize(x) * (1 + weight), fp32 math, single pass.""" + orig_shape = x.shape + n = orig_shape[-1] + x2 = x.reshape(-1, n) + m = x2.shape[0] + out = torch.empty((m, n), dtype=x.dtype, device=x.device) + block_n = triton.next_power_of_2(n) + _gemma_rmsnorm_kernel[(m,)]( + x2, + weight, + out, + n, + x2.stride(0), + x2.stride(1), + eps, + BLOCK_N=block_n, + num_warps=_num_warps(block_n), + ) + return out.reshape(orig_shape) + + +def gemma_fused_add_rmsnorm( + x: torch.Tensor, + residual: torch.Tensor, + weight: torch.Tensor, + eps: float, +): + """Fused (x + residual) then Gemma RMSNorm; returns (normed, pre-norm sum).""" + orig_shape = x.shape + n = orig_shape[-1] + x2 = x.reshape(-1, n) + r2 = residual.reshape(-1, n) + m = x2.shape[0] + out = torch.empty((m, n), dtype=x.dtype, device=x.device) + res_out = torch.empty((m, n), dtype=x.dtype, device=x.device) + block_n = triton.next_power_of_2(n) + _gemma_fused_add_rmsnorm_kernel[(m,)]( + x2, + r2, + weight, + out, + res_out, + n, + x2.stride(0), + x2.stride(1), + r2.stride(0), + r2.stride(1), + eps, + BLOCK_N=block_n, + num_warps=_num_warps(block_n), + ) + return out.reshape(orig_shape), res_out.reshape(orig_shape) diff --git a/python/sglang/jit_kernel/minimax_m3/swiglu.py b/python/sglang/jit_kernel/minimax_m3/swiglu.py new file mode 100644 index 000000000..4001388b1 --- /dev/null +++ b/python/sglang/jit_kernel/minimax_m3/swiglu.py @@ -0,0 +1,201 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Fused SwiGLU-OAI (split layout) Triton kernel for MiniMax-M3 on AMD ROCm. + +SwiGLU-OAI on a ``[*, 2I]`` split-layout tensor (gate = first half, up = second +half): ``gate * sigmoid(alpha * gate) * (up + beta)`` with optional clamp, +computed in fp32. Used by the dense MLP / shared experts ``swigluoai`` +activation on ROCm in place of the ``@torch.compile`` bf16 elementwise variant. +""" + +from typing import Optional + +import torch +import triton +import triton.language as tl + + +@triton.jit +def _swiglu_oai_kernel( + g_ptr, + out_ptr, + n_inter, + stride_gm, + stride_gn, + stride_om, + stride_on, + alpha, + beta, + limit, + HAS_LIMIT: tl.constexpr, + BLOCK_I: tl.constexpr, +): + row = tl.program_id(0) + pid_i = tl.program_id(1) + cols = pid_i * BLOCK_I + tl.arange(0, BLOCK_I) + mask = cols < n_inter + gate = tl.load(g_ptr + row * stride_gm + cols * stride_gn, mask=mask, other=0.0).to( + tl.float32 + ) + up = tl.load( + g_ptr + row * stride_gm + (n_inter + cols) * stride_gn, + mask=mask, + other=0.0, + ).to(tl.float32) + if HAS_LIMIT: + gate = tl.minimum(gate, limit) + up = tl.minimum(tl.maximum(up, -limit), limit) + out = gate * tl.sigmoid(alpha * gate) * (up + beta) + tl.store( + out_ptr + row * stride_om + cols * stride_on, + out.to(out_ptr.dtype.element_ty), + mask=mask, + ) + + +def swiglu_oai_split( + gate_up: torch.Tensor, + alpha: float, + beta: float, + limit: Optional[float], + out_dtype: Optional[torch.dtype] = None, +) -> torch.Tensor: + """SwiGLU-OAI on a split-layout ``[*, 2I]`` tensor -> ``[*, I]`` (fp32 math).""" + orig_shape = gate_up.shape + two_i = orig_shape[-1] + n_inter = two_i // 2 + x2 = gate_up.reshape(-1, two_i) + m = x2.shape[0] + dt = out_dtype if out_dtype is not None else gate_up.dtype + out = torch.empty((m, n_inter), dtype=dt, device=gate_up.device) + # Adaptive tile (tuned on gfx950). A 512-wide tile only helps + # once the (TP-sharded) per-rank slice is large enough to be bandwidth-bound + # (~1.25-1.35x faster than 256 at TP=1 prefill for the dense MLP I=12288). + # For small sharded slices (high TP) / decode the kernel is launch-bound, so + # fall back to 256. num_warps is pinned to 4 (8 underfills this tile). + block_i = 512 if n_inter >= 2048 else 256 + grid = (m, triton.cdiv(n_inter, block_i)) + _swiglu_oai_kernel[grid]( + x2, + out, + n_inter, + x2.stride(0), + x2.stride(1), + out.stride(0), + out.stride(1), + float(alpha), + float(beta), + 0.0 if limit is None else float(limit), + HAS_LIMIT=limit is not None, + BLOCK_I=block_i, + num_warps=4, + ) + return out.reshape(*orig_shape[:-1], n_inter) + + +@triton.jit +def _swiglu_oai_mxfp8_quant_kernel( + g_ptr, + q_ptr, + scale_ptr, + n_inter, + stride_gm, + stride_gn, + stride_qm, + stride_qn, + stride_sm, + stride_sn, + alpha, + beta, + limit, + HAS_LIMIT: tl.constexpr, + BLOCK_I: tl.constexpr, +): + row = tl.program_id(0) + pid_i = tl.program_id(1) + cols = pid_i * BLOCK_I + tl.arange(0, BLOCK_I) + mask = cols < n_inter + + gate = tl.load(g_ptr + row * stride_gm + cols * stride_gn, mask=mask, other=0.0) + up = tl.load( + g_ptr + row * stride_gm + (n_inter + cols) * stride_gn, + mask=mask, + other=0.0, + ) + gate = gate.to(tl.float32) + up = up.to(tl.float32) + if HAS_LIMIT: + gate = tl.minimum(gate, limit) + up = tl.minimum(tl.maximum(up, -limit), limit) + + # Keep the activation in fp32 all the way to the E8M0 scale selection (no + # bf16 round-trip to HBM). Matches the vLLM/ame fused swiglu+quant kernel: + # marginally more accurate than the unfused bf16 two-kernel chain. + activated = gate * tl.sigmoid(alpha * gate) * (up + beta) + + groups: tl.constexpr = BLOCK_I // 32 + activated_2d = tl.reshape(activated, (groups, 32)) + valid_groups = pid_i * groups + tl.arange(0, groups) < (n_inter // 32) + + amax = tl.maximum(tl.max(tl.abs(activated_2d), axis=1), 1e-30) + # Round the E8M0 exponent up (ceil(log2(amax / e4m3_max))) so the block amax + # stays inside the e4m3 range and the full dynamic range is used. + scale_biased = tl.ceil(tl.log2(amax / 448.0)) + 127.0 + scale_biased = tl.minimum(tl.maximum(scale_biased, 0.0), 254.0) + descale = tl.reshape(tl.exp2(scale_biased - 127.0), (groups, 1)) + + q_2d = tl.clamp(activated_2d / descale, -448.0, 448.0) + q = tl.reshape(q_2d, (BLOCK_I,)).to(q_ptr.dtype.element_ty) + + tl.store(q_ptr + row * stride_qm + cols * stride_qn, q, mask=mask) + tl.store( + scale_ptr + + row * stride_sm + + (pid_i * groups + tl.arange(0, groups)) * stride_sn, + scale_biased.to(tl.uint8), + mask=valid_groups, + ) + + +def swiglu_oai_mxfp8_quant( + gate_up: torch.Tensor, + alpha: float, + beta: float, + limit: Optional[float], +) -> tuple[torch.Tensor, torch.Tensor]: + """SwiGLU-OAI on split layout, then MiniMax MXFP8 quant, in one launch. + + The activation stays in fp32 through the E8M0 scale selection (no bf16 + round-trip), matching the vLLM/ame fused swiglu+quant kernel. + """ + orig_shape = gate_up.shape + two_i = orig_shape[-1] + n_inter = two_i // 2 + assert n_inter % 32 == 0, "MiniMax MXFP8 quant requires I divisible by 32." + + x2 = gate_up.reshape(-1, two_i) + m = x2.shape[0] + q = torch.empty((m, n_inter), dtype=torch.float8_e4m3fn, device=gate_up.device) + scales = torch.empty((m, n_inter // 32), dtype=torch.uint8, device=gate_up.device) + block_i = 512 if n_inter >= 2048 else 256 + grid = (m, triton.cdiv(n_inter, block_i)) + _swiglu_oai_mxfp8_quant_kernel[grid]( + x2, + q, + scales, + n_inter, + x2.stride(0), + x2.stride(1), + q.stride(0), + q.stride(1), + scales.stride(0), + scales.stride(1), + float(alpha), + float(beta), + 0.0 if limit is None else float(limit), + HAS_LIMIT=limit is not None, + BLOCK_I=block_i, + num_warps=4, + ) + return q.reshape(*orig_shape[:-1], n_inter), scales.reshape( + *orig_shape[:-1], n_inter // 32 + ) diff --git a/python/sglang/jit_kernel/minimax_qknorm_rope.py b/python/sglang/jit_kernel/minimax_qknorm_rope.py new file mode 100644 index 000000000..7815c398f --- /dev/null +++ b/python/sglang/jit_kernel/minimax_qknorm_rope.py @@ -0,0 +1,170 @@ +"""Fused per-head Gemma-RMSNorm + partial NeoX RoPE for MiniMax-M3 attention. + +In-place over a fused QKV tensor: normalizes + rotates one or more groups of +heads (each group = a contiguous head run sharing one norm weight, all getting +RoPE), leaving every other head (V, index-V) untouched. Consumes the model's +own ``cos_sin_cache`` (fp32) so the rotation matches sglang's RotaryEmbedding +exactly. + +Two entry points: + +* :func:`minimax_qknorm_rope` -- the original main-attention call + ``(q | k | v ...)``: Q heads then K heads, both normed + roped. +* :func:`minimax_qknorm_rope_grouped` -- a multi-group launch (up to 4 groups), + used to fold the main Q/K *and* the sparse-index Q/K of one fused + qkv+index-qkv GEMM output into a single kernel launch (mirroring the + multi-branch single-launch design of ``fused_store_kv_index.cuh``). +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, List, Sequence, Tuple + +import torch + +from sglang.jit_kernel.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 + +_MAX_GROUPS = 4 + + +@cache_once +def _jit_module(pos_dtype, head_dim, rope_dim) -> Module: + args = make_cpp_args(pos_dtype, head_dim, rope_dim, is_arch_support_pdl()) + return load_jit( + "fused_gemma_qknorm_rope", + *args, + cuda_files=["minimax/fused_gemma_qknorm_rope.cuh"], + cuda_wrappers=[("fused_gemma_qknorm_rope", f"fused_gemma_qknorm_rope<{args}>")], + ) + + +@register_custom_op(mutates_args=["qkv"]) +def _fused_gemma_qknorm_rope( + qkv: torch.Tensor, + w0: torch.Tensor, + w1: torch.Tensor, + w2: torch.Tensor, + w3: torch.Tensor, + cos_sin_cache: torch.Tensor, + positions: torch.Tensor, + off0: int, + cnt0: int, + off1: int, + cnt1: int, + off2: int, + cnt2: int, + off3: int, + cnt3: int, + num_groups: int, + eps: float, +) -> None: + # Wrap the tvm-ffi kernel as a custom op so torch.compile / piecewise CUDA + # graph can trace past the otherwise-opaque FFI call. The launch is + # graph-capturable (host-side constant offsets/counts), so it stays inside + # the captured region. + module = _jit_module(positions.dtype, 128, 64) + module.fused_gemma_qknorm_rope( + qkv, + w0, + w1, + w2, + w3, + cos_sin_cache, + positions, + off0, + cnt0, + off1, + cnt1, + off2, + cnt2, + off3, + cnt3, + num_groups, + eps, + ) + + +def minimax_qknorm_rope_grouped( + qkv: torch.Tensor, + groups: Sequence[Tuple[torch.Tensor, int, int]], + cos_sin_cache: torch.Tensor, + positions: torch.Tensor, + eps: float, +) -> torch.Tensor: + """Fused GemmaRMSNorm + partial NeoX RoPE over ``groups``, in place on ``qkv``. + + ``qkv`` is ``[T, total_heads * head_dim]`` (head_dim == 128). Each group is + ``(weight, head_offset, head_count)``: ``head_count`` consecutive heads + starting at ``head_offset`` (in head units) are normed with ``weight`` + (a ``[head_dim]`` bf16 tensor, the *raw* Gemma weight -- the kernel applies + ``1 + weight``) and rotated. Heads outside every group are untouched. + + Up to 4 groups are supported in one launch. The offsets/counts are + host-side constants, so the launch is CUDA-graph capturable. + """ + groups = [(w, off, cnt) for (w, off, cnt) in groups if cnt > 0] + num_groups = len(groups) + assert ( + 1 <= num_groups <= _MAX_GROUPS + ), f"need 1..{_MAX_GROUPS} groups, got {num_groups}" + + weights: List[torch.Tensor] = [g[0] for g in groups] + offsets: List[int] = [int(g[1]) for g in groups] + counts: List[int] = [int(g[2]) for g in groups] + # Pad weight slots up to 4 with a dummy (group 0's weight); the kernel never + # reads padded slots because num_groups bounds the in-kernel group scan. + while len(weights) < _MAX_GROUPS: + weights.append(weights[0]) + offsets.append(0) + counts.append(0) + + _fused_gemma_qknorm_rope( + qkv, + weights[0], + weights[1], + weights[2], + weights[3], + cos_sin_cache, + positions, + offsets[0], + counts[0], + offsets[1], + counts[1], + offsets[2], + counts[2], + offsets[3], + counts[3], + num_groups, + eps, + ) + return qkv + + +def minimax_qknorm_rope( + qkv: torch.Tensor, + q_weight: torch.Tensor, + k_weight: torch.Tensor, + cos_sin_cache: torch.Tensor, + positions: torch.Tensor, + nq: int, + nk: int, + nv: int, # deprecated / ignored: V heads are simply left untouched + eps: float, +) -> torch.Tensor: + """Main-attention layout ``[q (nq) | k (nk) | v ...]``: norm + rope Q then K.""" + return minimax_qknorm_rope_grouped( + qkv, + [(q_weight, 0, nq), (k_weight, nq, nk)], + cos_sin_cache, + positions, + eps, + ) diff --git a/python/sglang/jit_kernel/minimax_quant_ue8m0.py b/python/sglang/jit_kernel/minimax_quant_ue8m0.py new file mode 100644 index 000000000..240e9cf75 --- /dev/null +++ b/python/sglang/jit_kernel/minimax_quant_ue8m0.py @@ -0,0 +1,105 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, Tuple + +import torch + +from sglang.jit_kernel.utils import ( + cache_once, + is_arch_support_pdl, + load_jit, + make_cpp_args, +) + +if TYPE_CHECKING: + from tvm_ffi.module import Module + + +@cache_once +def _jit_module(group_size: int) -> Module: + args = make_cpp_args(group_size, is_arch_support_pdl()) + return load_jit( + "minimax_per_token_quant_ue8m0", + *args, + cuda_files=["minimax/per_token_quant_ue8m0.cuh"], + cuda_wrappers=[ + ("per_token_quant_ue8m0", f"per_token_quant_ue8m0<{args}>"), + ], + ) + + +@cache_once +def _jit_scatter_module(group_size: int, topk: int) -> Module: + # topk is a template arg so the dst-row load/store loops fully unroll. + args = make_cpp_args(group_size, topk, is_arch_support_pdl()) + return load_jit( + "minimax_per_token_quant_ue8m0_scatter", + *args, + cuda_files=["minimax/per_token_quant_ue8m0.cuh"], + cuda_wrappers=[ + ( + "per_token_quant_ue8m0_scatter", + f"per_token_quant_ue8m0_scatter<{args}>", + ), + ], + ) + + +def per_token_quant_fp8_ue8m0( + x: torch.Tensor, group_size: int = 128 +) -> Tuple[torch.Tensor, torch.Tensor]: + """Per-token group quant to FP8-e4m3 with a fused UE8M0 (int32-packed) scale. + + Returns ``(x_q, x_sf)`` where ``x_q`` is fp8_e4m3 ``[num_tokens, hidden]`` and + ``x_sf`` is the int32-packed UE8M0 scale ``[num_tokens, hidden//group_size//4]`` + (row-major). Byte-identical to ``per_token_group_quant_fp8(scale_ue8m0=True)`` + followed by ``transform_sf_into_required_layout`` (both ceil-round the scale), + but does it in a single kernel -- no separate transpose/pack launch. + """ + assert x.is_cuda and x.dtype == torch.bfloat16 and x.dim() == 2 + assert x.is_contiguous() + num_tokens, hidden = x.shape + assert hidden % group_size == 0 + num_groups = hidden // group_size + assert num_groups % 4 == 0, "num_groups must be a multiple of 4 for int32 packing" + + x_q = torch.empty_like(x, dtype=torch.float8_e4m3fn) + x_sf = torch.empty( + (num_tokens, num_groups // 4), dtype=torch.int32, device=x.device + ) + _jit_module(group_size).per_token_quant_ue8m0(x, x_q, x_sf) + return x_q, x_sf + + +def per_token_quant_fp8_ue8m0_scatter( + x: torch.Tensor, + gateup_input: torch.Tensor, + gateup_input_scale: torch.Tensor, + src2dst: torch.Tensor, + topk_ids: torch.Tensor, + topk: int, + m_max: int, + group_size: int = 128, +) -> None: + """Fused per-token FP8/UE8M0 quant **and** scatter into the permuted grouped-GEMM + input -- a single kernel replacing ``per_token_quant_fp8_ue8m0`` + + ``fill_gateup_input_triton_kernel``. + + For each source token it computes the fp8 row + int32-packed UE8M0 scale once, + then writes them to each of the token's ``topk`` destination rows: + ``gateup_input`` fp8 ``[E, m_max, hidden]`` (row ``src2dst[token, i]``) + ``gateup_input_scale`` int32 ``[E, hidden//group//4, m_max]`` (MN-major; byte-scattered) + Slots with ``topk_ids[token, i] < 0`` are skipped. Byte-identical to the + two-kernel path on every written row. + """ + assert x.is_cuda and x.dtype == torch.bfloat16 and x.dim() == 2 + assert x.is_contiguous() + assert gateup_input.dtype == torch.float8_e4m3fn and gateup_input.dim() == 3 + assert gateup_input_scale.dtype == torch.int32 and gateup_input_scale.dim() == 3 + num_tokens, hidden = x.shape + assert hidden % group_size == 0 + num_groups = hidden // group_size + assert num_groups % 4 == 0, "num_groups must be a multiple of 4 for int32 packing" + _jit_scatter_module(group_size, int(topk)).per_token_quant_ue8m0_scatter( + x, gateup_input, gateup_input_scale, src2dst, topk_ids, int(topk), int(m_max) + ) diff --git a/python/sglang/jit_kernel/minimax_store_kv_index.py b/python/sglang/jit_kernel/minimax_store_kv_index.py new file mode 100644 index 000000000..2c5fd4408 --- /dev/null +++ b/python/sglang/jit_kernel/minimax_store_kv_index.py @@ -0,0 +1,79 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, Optional + +import torch + +from sglang.jit_kernel.utils import ( + cache_once, + is_arch_support_pdl, + load_jit, + make_cpp_args, +) + +if TYPE_CHECKING: + from tvm_ffi.module import Module + + +@cache_once +def _jit_module(head_bytes: int) -> Module: + # Build marker is (head_bytes, kUsePDL); the index dtype (int32/int64) is a + # runtime dispatch inside the C++ launcher. + args = make_cpp_args(head_bytes, is_arch_support_pdl()) + return load_jit( + "minimax_store_kv_index", + *args, + cuda_files=["minimax/fused_store_kv_index.cuh"], + cuda_wrappers=[("store_kv_index", f"store_kv_index<{args}>")], + ) + + +def store_kv_index( + k: torch.Tensor, + v: torch.Tensor, + k_cache: torch.Tensor, + v_cache: torch.Tensor, + idx_k: torch.Tensor, + idx_k_cache: torch.Tensor, + idx_v: Optional[torch.Tensor], + idx_v_cache: Optional[torch.Tensor], + indices: torch.Tensor, + *, + num_kv_heads: int, + head_bytes: int, +) -> None: + """Fused store of the MiniMax-M3 sparse caches in one launch. + + Writes the main ``k``/``v`` (``num_kv_heads`` heads each), the single index + ``idx_k`` head, and optionally the single ``idx_v`` head into their caches + at the per-token rows given by ``indices`` (out_cache_loc). In-place on the + four cache tensors. + + All tensors must share the same (store) dtype and a head_dim whose byte size + equals ``head_bytes`` (a multiple of 16). ``k``/``idx_k`` are 2D rows + ``[T, num_kv_heads*head_dim]`` / ``[T, head_dim]``; caches are the matching + ``[num_pages, ...]`` buffers. When ``idx_v`` is None there is no index value + head (the layer is a pure block selector). + """ + has_v = idx_v is not None + if not has_v: + # Pass idx_k as a dummy for the unused index-V slot; heads_per_token is + # set so the kernel never reaches the index-V branch. + idx_v = idx_k + idx_v_cache = idx_k_cache + heads_per_token = 2 * num_kv_heads + 1 + (1 if has_v else 0) + + module = _jit_module(head_bytes) + module.store_kv_index( + k, + v, + k_cache, + v_cache, + idx_k, + idx_k_cache, + idx_v, + idx_v_cache, + indices, + num_kv_heads, + heads_per_token, + ) diff --git a/python/sglang/jit_kernel/moe_topk_sigmoid.py b/python/sglang/jit_kernel/moe_topk_sigmoid.py new file mode 100644 index 000000000..7fbdd22a6 --- /dev/null +++ b/python/sglang/jit_kernel/moe_topk_sigmoid.py @@ -0,0 +1,105 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, Optional + +import torch + +from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args +from sglang.srt.utils.custom_op import register_custom_op + +if TYPE_CHECKING: + from tvm_ffi.module import Module + + +@cache_once +def _jit_moe_topk_sigmoid_module(dtype: torch.dtype) -> Module: + args = make_cpp_args(dtype) + return load_jit( + "moe_topk_sigmoid", + *args, + cuda_files=["moe/moe_topk_sigmoid.cuh"], + cuda_wrappers=[("topk_sigmoid", f"topk_sigmoid<{args}>")], + extra_cuda_cflags=["--use_fast_math"], + ) + + +@register_custom_op( + op_name="moe_topk_sigmoid_out", + mutates_args=["topk_weights", "topk_ids"], +) +def moe_topk_sigmoid_out( + gating_output: torch.Tensor, + topk_weights: torch.Tensor, + topk_ids: torch.Tensor, + workspace: torch.Tensor, + renormalize: bool, + correction_bias: Optional[torch.Tensor], + routed_scaling_factor: float, + num_fused_shared_experts: int, +) -> None: + """ + Fused sigmoid top-k MoE gate (destination-passing style). + + Args: + gating_output: [num_tokens, num_experts], fp32/fp16/bf16 + topk_weights: [num_tokens, topk], float32, pre-allocated output + topk_ids: [num_tokens, topk], int32, pre-allocated output + workspace: [num_tokens * num_experts] float32 scratch (may be size 1 + when num_experts is a supported power-of-2 ≤ 256) + renormalize: whether to renormalize weights to sum to 1 per row + correction_bias: [num_experts] float32 per-expert bias, or None + routed_scaling_factor: [num_tokens, num_experts] float32, or None + """ + module = _jit_moe_topk_sigmoid_module(gating_output.dtype) + module.topk_sigmoid( + gating_output, + topk_weights, + topk_ids, + workspace, + renormalize, + correction_bias, + routed_scaling_factor, + num_fused_shared_experts, + ) + + +def topk_sigmoid( + topk_weights: torch.Tensor, + topk_ids: torch.Tensor, + gating_output: torch.Tensor, + renormalize: bool = False, + correction_bias: Optional[torch.Tensor] = None, + routed_scaling_factor: float = 1.0, + num_fused_shared_experts: int = 0, +) -> None: + """ + Fused sigmoid top-k MoE gate with the same call signature as + ``sgl_kernel.topk_sigmoid`` (destination-passing, in-place). + + Args: + topk_weights: [num_tokens, topk] float32, written in-place + topk_ids: [num_tokens, topk] int32, written in-place + gating_output: [num_tokens, num_experts] fp32/fp16/bf16 + renormalize: whether to renormalize weights to sum to 1 per row + correction_bias: [num_experts] float32 per-expert bias, or None + """ + num_tokens = gating_output.shape[0] + num_experts = gating_output.shape[1] + + is_pow2 = num_experts != 0 and (num_experts & (num_experts - 1)) == 0 + needs_workspace = not is_pow2 or num_experts > 256 + workspace_size = num_tokens * num_experts if needs_workspace else 1 + workspace = torch.empty( + workspace_size, dtype=torch.float32, device=gating_output.device + ) + + moe_topk_sigmoid_out( + gating_output, + topk_weights, + topk_ids, + workspace, + renormalize, + correction_bias, + routed_scaling_factor, + num_fused_shared_experts, + ) diff --git a/python/sglang/jit_kernel/tests/test_minimax_m3_mxfp8.py b/python/sglang/jit_kernel/tests/test_minimax_m3_mxfp8.py new file mode 100644 index 000000000..52eedf623 --- /dev/null +++ b/python/sglang/jit_kernel/tests/test_minimax_m3_mxfp8.py @@ -0,0 +1,265 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Reference-vs-fused unit tests for the MiniMax-M3 ROCm native MXFP8 ops. + +Each fused kernel has a slow PyTorch / dequant-to-bf16 reference; these assert +the two agree within tolerance: + + * Fused MXFP8 activation quant (Triton) -> torch reference + * Native MXFP8 linear (tl.dot_scaled) -> dequant-to-bf16 @ matmul + * Native MXFP8 MoE (dot_scaled grouped GEMM) -> dequant-to-bf16 MoE math + +ROCm-only. The pure quant test runs on any ROCm arch; the native MXFP8 +``dot_scaled`` linear/MoE tests are gated to CDNA4 gfx95x (the hardware +microscaling matrix cores) -- gfx942 has no native ``dot_scaled`` MX path. + +Run: pytest python/sglang/jit_kernel/tests/test_minimax_m3_mxfp8.py -v +""" + +import pytest +import torch + +from sglang.srt.utils import is_hip + +if not is_hip(): + pytest.skip( + "MiniMax-M3 native MXFP8 ops are the ROCm path.", allow_module_level=True + ) +if not torch.cuda.is_available(): + pytest.skip("Requires a GPU.", allow_module_level=True) + +from sglang.srt.layers.quantization.mxfp8_amd_gfx95 import ( # noqa: E402 + _mxfp8_dot_scaled_linear, + _mxfp8_e4m3_quantize_torch, + _mxfp8_e4m3_quantize_triton, + dequant_mxfp8_to_bf16, +) + +DEVICE = "cuda" + + +def _gcn_arch() -> str: + try: + return torch.cuda.get_device_properties(0).gcnArchName + except Exception: # pragma: no cover - no device / non-AMD + return "" + + +requires_gfx950 = pytest.mark.skipif( + "gfx95" not in _gcn_arch(), + reason="native MXFP8 dot_scaled is a CDNA4 (gfx95x) feature; " + "gfx942 has no native dot_scaled MX path.", +) + + +def _relerr(a: torch.Tensor, b: torch.Tensor) -> float: + a = a.float() + b = b.float() + return ((a - b).norm() / (b.norm() + 1e-8)).item() + + +# --------------------------------------------------------------------------- # +# Fused MXFP8 activation quant (Triton vs torch reference) +# --------------------------------------------------------------------------- # +@pytest.mark.parametrize("shape", [(64, 4096), (1, 6144), (333, 2048)]) +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +@torch.inference_mode() +def test_mxfp8_quant_triton_matches_torch(shape, dtype): + torch.manual_seed(0) + x = torch.randn(*shape, device=DEVICE, dtype=dtype) + xq_t, s_t = _mxfp8_e4m3_quantize_torch(x) + xq_k, s_k = _mxfp8_e4m3_quantize_triton(x) + assert s_k.shape == s_t.shape == (shape[0], shape[1] // 32) + # E8M0 block exponents share the round-up ceil(log2(amax/e4m3_max))+127 + # algorithm; allow at most a 1-step difference at exact powers of two. + assert (s_k.int() - s_t.int()).abs().max().item() <= 1 + # Dequantized values agree to fp8 granularity. + deq_t = dequant_mxfp8_to_bf16(xq_t, s_t) + deq_k = dequant_mxfp8_to_bf16(xq_k, s_k) + assert _relerr(deq_k, deq_t) < 1e-2 + + +@pytest.mark.parametrize("m,inter", [(8, 512), (65, 2048)]) +@torch.inference_mode() +def test_minimax_swiglu_mxfp8_quant_matches_unfused_fp32(m, inter): + # The fused swiglu+quant kernel keeps the activation in fp32 through the + # E8M0 scale selection (no bf16 round-trip; matches the vLLM/ame kernel), so + # the reference is the unfused fp32 swiglu followed by MXFP8 quant. Not + # bit-identical because the reference quant runs in torch vs the fused triton + # path, but numerically equivalent (tight relerr, scales agree within 1 ulp). + from sglang.jit_kernel.minimax_m3 import ( + swiglu_oai_mxfp8_quant, + swiglu_oai_split, + ) + from sglang.srt.layers.quantization.mxfp8_amd_gfx95 import mxfp8_e4m3_quantize + + torch.manual_seed(0) + alpha, beta, limit = 1.702, 1.0, 7.0 + gate_up = torch.randn(m, 2 * inter, device=DEVICE, dtype=torch.bfloat16) * 0.5 + + act = swiglu_oai_split( + gate_up, alpha=alpha, beta=beta, limit=limit, out_dtype=torch.float32 + ) + q_ref, s_ref = mxfp8_e4m3_quantize(act) + q, s = swiglu_oai_mxfp8_quant(gate_up, alpha=alpha, beta=beta, limit=limit) + + assert q.shape == q_ref.shape + assert s.shape == s_ref.shape + # E8M0 block scales agree within one exponent step (last-bit amax differences). + assert (s.int() - s_ref.int()).abs().max().item() <= 1 + assert ( + _relerr(dequant_mxfp8_to_bf16(q, s), dequant_mxfp8_to_bf16(q_ref, s_ref)) < 1e-2 + ) + + +# --------------------------------------------------------------------------- # +# Native MXFP8 linear (dot_scaled) vs dequant-to-bf16 matmul +# --------------------------------------------------------------------------- # +@requires_gfx950 +@pytest.mark.parametrize("m,n,k", [(64, 256, 128), (37, 512, 256), (1, 6144, 4096)]) +@torch.inference_mode() +def test_mxfp8_native_linear(m, n, k): + torch.manual_seed(0) + w_bf16 = torch.randn(n, k, device=DEVICE, dtype=torch.bfloat16) * 0.1 + w_fp8, w_scale = _mxfp8_e4m3_quantize_torch(w_bf16) + x = torch.randn(m, k, device=DEVICE, dtype=torch.bfloat16) * 0.5 + + got = _mxfp8_dot_scaled_linear(x, w_fp8, w_scale) + # Reference consumes the SAME quantized weights (isolates activation-quant + # noise) -> dequant to bf16, plain matmul. + w_deq = dequant_mxfp8_to_bf16(w_fp8, w_scale) + ref = torch.nn.functional.linear(x, w_deq).to(x.dtype) + assert got.shape == (m, n) + assert _relerr(got, ref) < 5e-2 + + +# --------------------------------------------------------------------------- # +# Native MXFP8 MoE (dot_scaled grouped GEMM) vs dequant-to-bf16 MoE math +# --------------------------------------------------------------------------- # +def _ref_moe(x, w13, w2, topk_weights, topk_ids, alpha, beta, limit): + T, H = x.shape + inter = w2.shape[-1] + top_k = topk_ids.shape[1] + out = torch.zeros(T, H, device=x.device, dtype=torch.float32) + for t in range(T): + for j in range(top_k): + e = int(topk_ids[t, j].item()) + if e < 0 or e >= w13.shape[0]: + continue + g1 = x[t].float() @ w13[e].float().T # [2I] + gate = g1[:inter] + up = g1[inter:] + if limit is not None: + gate = gate.clamp(max=limit) + up = up.clamp(min=-limit, max=limit) + act = gate * torch.sigmoid(alpha * gate) * (up + beta) + g2 = act @ w2[e].float().T # [H] + out[t] += topk_weights[t, j].float() * g2 + return out.to(x.dtype) + + +@requires_gfx950 +@pytest.mark.parametrize( + "T,H,inter,E,top_k", [(8, 256, 512, 8, 2), (1, 512, 256, 16, 4)] +) +@torch.inference_mode() +def test_mxfp8_native_moe(T, H, inter, E, top_k): + from sglang.srt.layers.moe.moe_runner.triton_utils.mxfp8_moe_amd_gfx95 import ( + fused_moe_mxfp8_native, + ) + + torch.manual_seed(0) + alpha, beta, limit = 1.702, 1.0, 7.0 + w13_bf16 = torch.randn(E, 2 * inter, H, device=DEVICE, dtype=torch.bfloat16) * 0.1 + w2_bf16 = torch.randn(E, H, inter, device=DEVICE, dtype=torch.bfloat16) * 0.1 + w13_fp8, w13_scale = _mxfp8_e4m3_quantize_torch(w13_bf16) + w2_fp8, w2_scale = _mxfp8_e4m3_quantize_torch(w2_bf16) + + x = torch.randn(T, H, device=DEVICE, dtype=torch.bfloat16) * 0.5 + logits = torch.randn(T, E, device=DEVICE, dtype=torch.float32) + topk_weights, topk_ids = logits.softmax(dim=-1).topk(top_k, dim=-1) + topk_weights = topk_weights.to(torch.float32) + topk_ids = topk_ids.to(torch.int32) + + got = fused_moe_mxfp8_native( + x, + w13_fp8, + w13_scale, + w2_fp8, + w2_scale, + topk_weights, + topk_ids, + alpha=alpha, + beta=beta, + limit=limit, + ) + # Reference consumes the dequantized weights (same bits the kernel reads). + w13_deq = dequant_mxfp8_to_bf16(w13_fp8, w13_scale) + w2_deq = dequant_mxfp8_to_bf16(w2_fp8, w2_scale) + ref = _ref_moe(x, w13_deq, w2_deq, topk_weights, topk_ids, alpha, beta, limit) + assert got.shape == (T, H) + assert _relerr(got, ref) < 5e-2 + + +@requires_gfx950 +@torch.inference_mode() +def test_mxfp8_native_moe_ep_expert_map_filters_non_local_routes(): + from sglang.srt.layers.moe.moe_runner.triton_utils.mxfp8_moe_amd_gfx95 import ( + fused_moe_mxfp8_native, + ) + + torch.manual_seed(0) + T, H, inter = 4, 256, 512 + local_E = 3 + alpha, beta, limit = 1.702, 1.0, 7.0 + + w13_bf16 = ( + torch.randn(local_E, 2 * inter, H, device=DEVICE, dtype=torch.bfloat16) * 0.1 + ) + w2_bf16 = torch.randn(local_E, H, inter, device=DEVICE, dtype=torch.bfloat16) * 0.1 + w13_fp8, w13_scale = _mxfp8_e4m3_quantize_torch(w13_bf16) + w2_fp8, w2_scale = _mxfp8_e4m3_quantize_torch(w2_bf16) + + x = torch.randn(T, H, device=DEVICE, dtype=torch.bfloat16) * 0.5 + topk_ids_global = torch.tensor( + [[0, 1, 4], [2, 3, 5], [4, 0, 3], [5, 1, 2]], + device=DEVICE, + dtype=torch.int32, + ) + topk_weights = torch.tensor( + [ + [0.50, 0.25, 0.25], + [0.40, 0.30, 0.30], + [0.70, 0.20, 0.10], + [0.60, 0.30, 0.10], + ], + device=DEVICE, + dtype=torch.float32, + ) + expert_map = torch.tensor([0, -1, 1, -1, 2, -1], device=DEVICE, dtype=torch.int32) + + got = fused_moe_mxfp8_native( + x, + w13_fp8, + w13_scale, + w2_fp8, + w2_scale, + topk_weights, + topk_ids_global, + alpha=alpha, + beta=beta, + limit=limit, + expert_map=expert_map, + ) + + topk_ids_local = expert_map[topk_ids_global.long()] + w13_deq = dequant_mxfp8_to_bf16(w13_fp8, w13_scale) + w2_deq = dequant_mxfp8_to_bf16(w2_fp8, w2_scale) + ref = _ref_moe(x, w13_deq, w2_deq, topk_weights, topk_ids_local, alpha, beta, limit) + assert got.shape == (T, H) + assert _relerr(got, ref) < 5e-2 + + +if __name__ == "__main__": + import sys + + sys.exit(pytest.main([__file__, "-v"])) diff --git a/python/sglang/jit_kernel/tests/test_minimax_m3_rmsnorm.py b/python/sglang/jit_kernel/tests/test_minimax_m3_rmsnorm.py new file mode 100644 index 000000000..6897039c8 --- /dev/null +++ b/python/sglang/jit_kernel/tests/test_minimax_m3_rmsnorm.py @@ -0,0 +1,78 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Reference tests for MiniMax-M3 ROCm Gemma RMSNorm Triton kernels.""" + +import pytest +import torch + +from sglang.srt.utils import is_hip + +if not is_hip(): + pytest.skip( + "MiniMax-M3 Gemma RMSNorm Triton kernels are ROCm-only.", + allow_module_level=True, + ) +if not torch.cuda.is_available(): + pytest.skip("Requires a GPU.", allow_module_level=True) + +from sglang.jit_kernel.minimax_m3.rmsnorm import ( # noqa: E402 + gemma_fused_add_rmsnorm, + gemma_rmsnorm, +) + +DEVICE = "cuda" +EPS = 1e-6 + + +def _gemma_rmsnorm_ref(x: torch.Tensor, weight: torch.Tensor) -> torch.Tensor: + orig_dtype = x.dtype + x_f = x.float() + variance = x_f.pow(2).mean(dim=-1, keepdim=True) + out = x_f * torch.rsqrt(variance + EPS) + out = out * (1.0 + weight.float()) + return out.to(orig_dtype) + + +@pytest.mark.parametrize("shape", [(1, 512), (64, 6144), (257, 6144)]) +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +@torch.inference_mode() +def test_gemma_rmsnorm_matches_reference(shape, dtype): + torch.manual_seed(0) + x = torch.randn(*shape, device=DEVICE, dtype=dtype) + weight = torch.randn(shape[-1], device=DEVICE, dtype=torch.float32) + + got = gemma_rmsnorm(x, weight, EPS) + ref = _gemma_rmsnorm_ref(x, weight) + + torch.testing.assert_close(got, ref, atol=2e-2, rtol=2e-2) + + +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +@torch.inference_mode() +def test_gemma_rmsnorm_accepts_strided_2d_input(dtype): + torch.manual_seed(0) + base = torch.randn(128, 1024, device=DEVICE, dtype=dtype) + x = base[:, ::2] + weight = torch.randn(x.shape[-1], device=DEVICE, dtype=torch.float32) + + assert not x.is_contiguous() + got = gemma_rmsnorm(x, weight, EPS) + ref = _gemma_rmsnorm_ref(x, weight) + + torch.testing.assert_close(got, ref, atol=2e-2, rtol=2e-2) + + +@pytest.mark.parametrize("shape", [(1, 512), (64, 6144), (257, 6144)]) +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +@torch.inference_mode() +def test_gemma_fused_add_rmsnorm_matches_reference(shape, dtype): + torch.manual_seed(0) + x = torch.randn(*shape, device=DEVICE, dtype=dtype) + residual = torch.randn(*shape, device=DEVICE, dtype=dtype) + weight = torch.randn(shape[-1], device=DEVICE, dtype=torch.float32) + + got, residual_out = gemma_fused_add_rmsnorm(x, residual, weight, EPS) + ref_residual = x + residual + ref = _gemma_rmsnorm_ref(ref_residual, weight) + + torch.testing.assert_close(residual_out, ref_residual, atol=2e-2, rtol=2e-2) + torch.testing.assert_close(got, ref, atol=2e-2, rtol=2e-2) diff --git a/python/sglang/jit_kernel/tests/test_moe_topk_sigmoid.py b/python/sglang/jit_kernel/tests/test_moe_topk_sigmoid.py new file mode 100644 index 000000000..4acbd23bd --- /dev/null +++ b/python/sglang/jit_kernel/tests/test_moe_topk_sigmoid.py @@ -0,0 +1,493 @@ +""" +Correctness tests for the moe_topk_sigmoid JIT kernel. + +Validates against a pure-PyTorch reference and, when sgl_kernel is available, +cross-checks against the AOT implementation. +""" + +import itertools +import os +import sys +from typing import Optional + +import pytest +import torch + +from sglang.jit_kernel.moe_topk_sigmoid import topk_sigmoid + +try: + from sgl_kernel import topk_sigmoid as topk_sigmoid_aot + + AOT_AVAILABLE = True +except ImportError: + AOT_AVAILABLE = False + +# --------------------------------------------------------------------------- +# CI / full-range helpers +# --------------------------------------------------------------------------- + +_is_ci = ( + os.getenv("CI", "false").lower() == "true" + or os.getenv("GITHUB_ACTIONS", "false").lower() == "true" +) + +# Power-of-2 configs covered by static dispatch (num_experts 1–256) +# Plus 48 (non-power-of-2) to exercise the fallback path +NUM_TOKENS_FULL = [1, 16, 128, 512, 1024, 2048] +NUM_TOKENS_CI = [1, 128, 1024] + +NUM_EXPERTS_FULL = [16, 32, 64, 128, 256, 48] # 48 = fallback path +NUM_EXPERTS_CI = [16, 64, 48] + +TOPK_FULL = [1, 2, 4, 8] +TOPK_CI = [1, 4] + +DTYPES_FULL = [torch.float32] +DTYPES_CI = [torch.float32, torch.bfloat16] + +NUM_TOKENS = NUM_TOKENS_CI if _is_ci else NUM_TOKENS_FULL +NUM_EXPERTS = NUM_EXPERTS_CI if _is_ci else NUM_EXPERTS_FULL +TOPK_LIST = TOPK_CI if _is_ci else TOPK_FULL +DTYPES = DTYPES_CI if _is_ci else DTYPES_FULL + + +# --------------------------------------------------------------------------- +# Pure-PyTorch reference +# --------------------------------------------------------------------------- + + +def grouped_topk_gpu( + gating_output: torch.Tensor, + topk: int, + renormalize: bool, + num_expert_group: Optional[int] = None, + topk_group: Optional[int] = None, + num_fused_shared_experts: int = 0, + routed_scaling_factor: Optional[float] = None, + apply_routed_scaling_factor_on_output: Optional[bool] = False, + scoring_func: str = "softmax", +): + + # Scoring function: softmax or sigmoid + if scoring_func == "softmax": + scores = torch.softmax(gating_output, dim=-1) + elif scoring_func == "sigmoid": + scores = gating_output.sigmoid() + else: + raise ValueError(f"Unsupported scoring function: {scoring_func}") + + num_token = scores.shape[0] + num_experts = scores.shape[1] + group_scores = ( + scores.view(num_token, num_expert_group, -1).max(dim=-1).values + ) # [n, n_group] + + group_idx = torch.topk(group_scores, k=topk_group, dim=-1, sorted=False)[ + 1 + ] # [n, top_k_group] + + group_mask = torch.zeros_like(group_scores) # [n, n_group] + group_mask.scatter_(1, group_idx, 1) # [n, n_group] + score_mask = ( + group_mask.unsqueeze(-1) + .expand(num_token, num_expert_group, scores.shape[-1] // num_expert_group) + .reshape(num_token, -1) + ) # [n, e] + tmp_scores = scores.masked_fill( + ~score_mask.bool(), float("-inf") + ) # [n, e] - use -inf like VLLM + + topk_weights, topk_ids = torch.topk( + tmp_scores, + k=topk, + dim=-1, + sorted=(True if num_fused_shared_experts > 0 else True), + ) + + if num_fused_shared_experts: + topk_ids[:, -1] = torch.randint( + low=num_experts, + high=num_experts + num_fused_shared_experts, + size=(topk_ids.size(0),), + dtype=topk_ids.dtype, + device=topk_ids.device, + ) + if routed_scaling_factor is not None: + topk_weights[:, -1] = ( + topk_weights[:, :-1].sum(dim=-1) / routed_scaling_factor + ) + + if renormalize: + topk_weights_sum = ( + topk_weights.sum(dim=-1, keepdim=True) + if num_fused_shared_experts == 0 + else topk_weights[:, :-1].sum(dim=-1, keepdim=True) + ) + topk_weights = topk_weights / topk_weights_sum + if apply_routed_scaling_factor_on_output: + topk_weights *= routed_scaling_factor + + topk_weights, topk_ids = topk_weights.to(torch.float32), topk_ids.to(torch.int32) + return topk_weights, topk_ids + + +def biased_grouped_topk_impl( + gating_output: torch.Tensor, + correction_bias: torch.Tensor, + topk: int, + renormalize: bool, + num_expert_group: Optional[int] = None, + topk_group: Optional[int] = None, + num_fused_shared_experts: int = 0, + routed_scaling_factor: Optional[float] = None, + apply_routed_scaling_factor_on_output: Optional[bool] = False, +): + scores = gating_output.sigmoid() + + num_token = scores.shape[0] + num_experts = scores.shape[1] + scores_for_choice = scores.view(num_token, -1) + correction_bias.unsqueeze(0) + group_scores = ( + scores_for_choice.view(num_token, num_expert_group, -1) + .topk(2, dim=-1)[0] + .sum(dim=-1) + ) # [n, n_group] + + group_idx = torch.topk(group_scores, k=topk_group, dim=-1, sorted=False)[ + 1 + ] # [n, top_k_group] + + group_mask = torch.zeros_like(group_scores) # [n, n_group] + group_mask.scatter_(1, group_idx, 1) # [n, n_group] + score_mask = ( + group_mask.unsqueeze(-1) + .expand(num_token, num_expert_group, scores.shape[-1] // num_expert_group) + .reshape(num_token, -1) + ) # [n, e] + tmp_scores = scores_for_choice.masked_fill( + ~score_mask.bool(), float("-inf") + ) # [n, e] + + _, topk_ids = torch.topk( + tmp_scores, + k=topk, + dim=-1, + sorted=(True if num_fused_shared_experts > 0 else True), + ) + topk_weights = scores.gather(1, topk_ids) + + if num_fused_shared_experts: + topk_ids[:, -1] = torch.randint( + low=num_experts, + high=num_experts + num_fused_shared_experts, + size=(topk_ids.size(0),), + dtype=topk_ids.dtype, + device=topk_ids.device, + ) + if routed_scaling_factor is not None: + topk_weights[:, -1] = ( + topk_weights[:, :-1].sum(dim=-1) / routed_scaling_factor + ) + if renormalize: + topk_weights_sum = ( + topk_weights.sum(dim=-1, keepdim=True) + if num_fused_shared_experts == 0 + else topk_weights[:, :-1].sum(dim=-1, keepdim=True) + ) + topk_weights = topk_weights / topk_weights_sum + if apply_routed_scaling_factor_on_output: + topk_weights *= routed_scaling_factor + + topk_weights, topk_ids = topk_weights.to(torch.float32), topk_ids.to(torch.int32) + return topk_weights, topk_ids + + +def topk_sigmoid_torch_ref( + gating_output: torch.Tensor, + topk: int, + renormalize: bool, + correction_bias: torch.Tensor | None, + num_fused_shared_experts: int = 0, + routed_scaling_factor: float = 1.0, + apply_routed_scaling_factor_on_output: bool = True, +) -> tuple[torch.Tensor, torch.Tensor]: + """ + Reference: sigmoid → (add bias) → topk → (renormalize). + Indices are selected on biased scores; weights are the unbiased sigmoid values. + """ + num_experts = gating_output.shape[1] + scores = gating_output.float().sigmoid() + biased = scores if correction_bias is None else scores + correction_bias.float() + _, ref_ids = torch.topk(biased, k=topk, dim=-1) + ref_weights = scores.gather(1, ref_ids) + if num_fused_shared_experts > 0: + ref_ids[:, -1] = torch.randint( + low=num_experts, + high=num_experts + num_fused_shared_experts, + size=(ref_ids.size(0),), + dtype=ref_ids.dtype, + device=ref_ids.device, + ) + ref_weights[:, -1] = ref_weights[:, :-1].sum(dim=-1) / routed_scaling_factor + if renormalize: + topk_weights_sum = ( + ref_weights.sum(dim=-1, keepdim=True) + if num_fused_shared_experts == 0 + else ref_weights[:, :-1].sum(dim=-1, keepdim=True) + ) + ref_weights = ref_weights / topk_weights_sum + if apply_routed_scaling_factor_on_output: + ref_weights *= routed_scaling_factor + return ref_weights.float(), ref_ids.int() + + +def topk_sigmoid_grouped_ref( + gating_output: torch.Tensor, + topk: int, + renormalize: bool, + correction_bias: torch.Tensor | None, + num_fused_shared_experts: int = 0, +) -> tuple[torch.Tensor, torch.Tensor]: + if correction_bias is not None: + return biased_grouped_topk_impl( + gating_output, + correction_bias, + topk, + renormalize, + num_expert_group=1, + topk_group=1, + num_fused_shared_experts=num_fused_shared_experts, + routed_scaling_factor=1.0, + apply_routed_scaling_factor_on_output=True, + ) + else: + return grouped_topk_gpu( + gating_output, + topk, + renormalize, + num_expert_group=1, + topk_group=1, + num_fused_shared_experts=num_fused_shared_experts, + routed_scaling_factor=1.0, + apply_routed_scaling_factor_on_output=True, + scoring_func="sigmoid", + ) + + +def topk_sigmoid_ref( + gating_output: torch.Tensor, + topk: int, + renormalize: bool, + correction_bias: torch.Tensor | None, + num_fused_shared_experts: int = 0, +) -> tuple[torch.Tensor, torch.Tensor]: + return topk_sigmoid_torch_ref( + gating_output, topk, renormalize, correction_bias, num_fused_shared_experts + ) + + +# --------------------------------------------------------------------------- +# Correctness: JIT vs PyTorch reference +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize( + "num_tokens, num_experts, topk", + list(itertools.product(NUM_TOKENS, NUM_EXPERTS, TOPK_LIST)), +) +@pytest.mark.parametrize("dtype", DTYPES) +@pytest.mark.parametrize("renormalize", [False, True]) +def test_topk_sigmoid_vs_ref(num_tokens, num_experts, topk, dtype, renormalize): + if topk > num_experts: + pytest.skip("topk > num_experts") + + torch.manual_seed(num_tokens * num_experts) + gating = torch.randn((num_tokens, num_experts), dtype=dtype, device="cuda") + + topk_w = torch.empty((num_tokens, topk), dtype=torch.float32, device="cuda") + topk_i = torch.empty((num_tokens, topk), dtype=torch.int32, device="cuda") + topk_sigmoid(topk_w, topk_i, gating, renormalize=renormalize) + + ref_w, ref_i = topk_sigmoid_ref(gating, topk, renormalize, correction_bias=None) + + # Compare sorted weights (indices may differ for ties when dtype != float32) + assert torch.allclose( + topk_w.sort(dim=-1)[0], + ref_w.sort(dim=-1)[0], + atol=1e-3, + rtol=1e-3, + ), f"Weight mismatch (dtype={dtype}, n_exp={num_experts}, topk={topk}, renorm={renormalize})" + # Exact index match is only reliable for float32 (fp16/bf16 tie-breaking may differ) + if dtype == torch.float32: + assert torch.equal( + topk_i, ref_i + ), f"Index mismatch (dtype={dtype}, n_exp={num_experts}, topk={topk})" + + +# --------------------------------------------------------------------------- +# Correctness: with correction_bias +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize( + "num_tokens, num_experts, topk", + list(itertools.product(NUM_TOKENS, NUM_EXPERTS, TOPK_LIST)), +) +@pytest.mark.parametrize("renormalize", [False, True]) +def test_topk_sigmoid_with_correction_bias(num_tokens, num_experts, topk, renormalize): + if topk > num_experts: + pytest.skip("topk > num_experts") + + torch.manual_seed(num_tokens + num_experts + topk) + gating = torch.randn((num_tokens, num_experts), dtype=torch.float32, device="cuda") + bias = torch.randn(num_experts, dtype=torch.float32, device="cuda") + + topk_w = torch.empty((num_tokens, topk), dtype=torch.float32, device="cuda") + topk_i = torch.empty((num_tokens, topk), dtype=torch.int32, device="cuda") + topk_sigmoid(topk_w, topk_i, gating, renormalize=renormalize, correction_bias=bias) + + ref_w, ref_i = topk_sigmoid_ref(gating, topk, renormalize, correction_bias=bias) + + assert torch.allclose( + topk_w, ref_w, atol=1e-3, rtol=1e-3 + ), f"Weight mismatch with bias (n_exp={num_experts}, topk={topk}, renorm={renormalize})" + assert torch.equal( + topk_i, ref_i + ), f"Index mismatch with bias (n_exp={num_experts}, topk={topk})" + + +# --------------------------------------------------------------------------- +# Correctness: with fused shared experts +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize( + "num_tokens, num_experts, topk", + list(itertools.product(NUM_TOKENS, NUM_EXPERTS, TOPK_LIST)), +) +@pytest.mark.parametrize("renormalize", [False, True]) +def test_topk_sigmoid_with_fused_shared_experts( + num_tokens, num_experts, topk, renormalize +): + if topk + 1 > num_experts: + pytest.skip("topk > num_experts") + + torch.manual_seed(num_tokens + num_experts) + gating = torch.randn((num_tokens, num_experts), dtype=torch.float32, device="cuda") + bias = torch.randn(num_experts, dtype=torch.float32, device="cuda") + + topk_w = torch.empty((num_tokens, topk + 1), dtype=torch.float32, device="cuda") + topk_i = torch.empty((num_tokens, topk + 1), dtype=torch.int32, device="cuda") + topk_sigmoid( + topk_w, + topk_i, + gating, + renormalize=renormalize, + correction_bias=bias, + num_fused_shared_experts=1, + ) + + ref_w, ref_i = topk_sigmoid_ref( + gating, topk + 1, renormalize, correction_bias=bias, num_fused_shared_experts=1 + ) + + assert torch.allclose( + topk_w, ref_w, atol=1e-3, rtol=1e-3 + ), f"Weight mismatch with bias (n_exp={num_experts}, topk={topk}, renorm={renormalize})" + assert torch.equal( + topk_i, ref_i + ), f"Index mismatch with bias (n_exp={num_experts}, topk={topk})" + + +# --------------------------------------------------------------------------- +# Renormalization: weights should sum to 1 per row +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize("num_tokens, num_experts, topk", [(128, 64, 4), (1, 8, 2)]) +def test_renormalize_sums_to_one(num_tokens, num_experts, topk): + gating = torch.randn((num_tokens, num_experts), dtype=torch.float32, device="cuda") + topk_w = torch.empty((num_tokens, topk), dtype=torch.float32, device="cuda") + topk_i = torch.empty((num_tokens, topk), dtype=torch.int32, device="cuda") + topk_sigmoid(topk_w, topk_i, gating, renormalize=True) + row_sums = topk_w.sum(dim=-1) + torch.testing.assert_close( + row_sums, torch.ones(num_tokens, device="cuda"), rtol=1e-4, atol=1e-4 + ) + + +# --------------------------------------------------------------------------- +# Output shape and dtype +# --------------------------------------------------------------------------- + + +def test_output_shapes_and_dtypes(): + num_tokens, num_experts, topk = 64, 128, 4 + gating = torch.randn((num_tokens, num_experts), dtype=torch.float32, device="cuda") + topk_w = torch.empty((num_tokens, topk), dtype=torch.float32, device="cuda") + topk_i = torch.empty((num_tokens, topk), dtype=torch.int32, device="cuda") + topk_sigmoid(topk_w, topk_i, gating) + + assert topk_w.shape == (num_tokens, topk) + assert topk_i.shape == (num_tokens, topk) + assert topk_w.dtype == torch.float32 + assert topk_i.dtype == torch.int32 + + +# --------------------------------------------------------------------------- +# Fallback path (non-power-of-2 experts) +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize("num_experts", [48, 96]) +def test_fallback_non_power_of_two(num_experts): + num_tokens, topk = 64, 2 + gating = torch.randn((num_tokens, num_experts), dtype=torch.float32, device="cuda") + topk_w = torch.empty((num_tokens, topk), dtype=torch.float32, device="cuda") + topk_i = torch.empty((num_tokens, topk), dtype=torch.int32, device="cuda") + topk_sigmoid(topk_w, topk_i, gating, renormalize=True) + + # Weights should be positive and sum to 1 + assert torch.all(topk_w > 0) + torch.testing.assert_close( + topk_w.sum(dim=-1), torch.ones(num_tokens, device="cuda"), rtol=1e-4, atol=1e-4 + ) + + +# --------------------------------------------------------------------------- +# Cross-validation against AOT sgl_kernel +# --------------------------------------------------------------------------- + + +@pytest.mark.skipif(not AOT_AVAILABLE, reason="sgl_kernel not available") +@pytest.mark.parametrize( + "num_tokens, num_experts, topk", + list(itertools.product([1, 128, 1024], [8, 64, 128], [1, 4])), +) +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +@pytest.mark.parametrize("renormalize", [False, True]) +def test_topk_sigmoid_vs_aot(num_tokens, num_experts, topk, dtype, renormalize): + if topk > num_experts: + pytest.skip("topk > num_experts") + + torch.manual_seed(42) + gating = torch.randn((num_tokens, num_experts), dtype=dtype, device="cuda") + + topk_w_jit = torch.empty((num_tokens, topk), dtype=torch.float32, device="cuda") + topk_i_jit = torch.empty((num_tokens, topk), dtype=torch.int32, device="cuda") + topk_sigmoid(topk_w_jit, topk_i_jit, gating, renormalize=renormalize) + + topk_w_aot = torch.empty((num_tokens, topk), dtype=torch.float32, device="cuda") + topk_i_aot = torch.empty((num_tokens, topk), dtype=torch.int32, device="cuda") + topk_sigmoid_aot(topk_w_aot, topk_i_aot, gating, renormalize=renormalize) + + assert torch.allclose( + topk_w_jit, topk_w_aot, atol=1e-3, rtol=1e-3 + ), f"JIT vs AOT weight mismatch (dtype={dtype}, n_exp={num_experts}, topk={topk})" + assert torch.equal( + topk_i_jit, topk_i_aot + ), f"JIT vs AOT index mismatch (dtype={dtype}, n_exp={num_experts}, topk={topk})" + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v"])) diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py index 8921f71af..6801a18e5 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -127,6 +127,51 @@ def get_dsa_index_head_dim(config: PretrainedConfig) -> int: return config.index_head_dim +def is_minimax_sparse(config: PretrainedConfig) -> bool: + arch = (config.architectures or [None])[0] + return arch in ( + "MiniMaxM3SparseForCausalLM", + "MiniMaxM3SparseForConditionalGeneration", + ) + + +def get_minimax_sparse_attention_config(config: PretrainedConfig) -> dict: + text_cfg = getattr(config, "text_config", None) + cfg = ( + getattr(text_cfg, "sparse_attention_config", None) + if text_cfg is not None + else None + ) + if cfg is None: + cfg = getattr(config, "sparse_attention_config", None) + if cfg is None: + raise ValueError("Could not find sparse config. Is it MiniMax M3 Sparse model?") + return cfg + + +def get_minimax_sparse_layer_ids(sparse_cfg: dict) -> tuple[list[int], list[int]]: + sparse_freq = sparse_cfg["sparse_attention_freq"] + dense_layer_ids = [i for i, f in enumerate(sparse_freq) if f == 0] + sparse_layer_ids = [i for i, f in enumerate(sparse_freq) if f != 0] + return dense_layer_ids, sparse_layer_ids + + +def get_minimax_sparse_disable_value_layer_ids(sparse_cfg: dict) -> list[int]: + flags = sparse_cfg.get("sparse_disable_index_value") + if flags is None: + return [] + return [i for i, f in enumerate(flags) if f != 0] + + +def get_minimax_sparse_score_type(sparse_cfg: dict) -> str: + score_type = sparse_cfg.get("sparse_score_type", "max") + assert score_type in ( + "max", + "lse", + ), f"sparse_score_type must be 'max' or 'lse', got {score_type!r}" + return score_type + + def get_dsa_index_topk(config: PretrainedConfig) -> int: assert is_deepseek_dsa(config) return config.index_topk @@ -378,8 +423,15 @@ class ModelConfig: self.hf_config.architectures ) self.use_ngram_embedding = getattr(self.hf_config, "use_ngram_embedding", False) + # A multimodal arch is piecewise-incompatible until its LM prefill is validated. self.is_piecewise_cuda_graph_disabled_model = ( is_piecewise_cuda_graph_disabled_model(self.hf_config.architectures) + or ( + self.is_multimodal + and not is_multimodal_piecewise_cuda_graph_supported( + self.hf_config.architectures + ) + ) ) # Multimodal archs whose language-model prefill is verified safe to capture # under piecewise CUDA graph. ServerArgs otherwise disables prefill piecewise @@ -1215,6 +1267,7 @@ class ModelConfig: "petit_nvfp4", "quark", "mxfp4", + "mxfp8", "auto-round", "quark_int4fp8_moe", "quark_mxfp4", @@ -1321,11 +1374,16 @@ class ModelConfig: f"({self.quantization})." ) - # Check if the scale_fmt is ue8m0, and warn user if deepgemm is enabled for non-ue8m0 models on blackwell + # Warn if DeepGemm is enabled for a non-ue8m0 checkpoint on Blackwell. + # MXFP8 stores E8M0 block scales that DeepGemm consumes losslessly, so skip the warning there. self.use_scale_ue8m0 = quant_cfg.get("scale_fmt", None) == "ue8m0" from sglang.srt.layers import deep_gemm_wrapper - if not self.use_scale_ue8m0 and deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0: + if ( + not self.use_scale_ue8m0 + and deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0 + and self.quantization != "mxfp8" + ): logger.warning( "DeepGemm is enabled but the scale_fmt of checkpoint is not ue8m0. This might cause accuracy degradation on Blackwell." ) @@ -1589,6 +1647,7 @@ multimodal_model_archs = [ "Cohere2VisionForConditionalGeneration", "DeepseekVL2ForCausalLM", "Ernie4_5_VLMoeForConditionalGeneration", + "MiniMaxM3SparseForConditionalGeneration", "Gemma3ForConditionalGeneration", "Gemma3nForConditionalGeneration", "Gemma4ForConditionalGeneration", @@ -1725,6 +1784,14 @@ def is_piecewise_cuda_graph_disabled_model(model_architectures: List[str]): ) +# Multimodal archs whose LM-decoder prefill is validated under piecewise CUDA +# graph (capture wraps only the decoder; the image encoder runs eager). +multimodal_piecewise_cuda_graph_supported_archs = [ + "MiniMaxM3SparseForCausalLM", + "MiniMaxM3SparseForConditionalGeneration", +] + + def is_multimodal_piecewise_cuda_graph_supported(model_architectures: List[str]): """Whether a multimodal arch may keep prefill piecewise CUDA graph enabled.""" return any( diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 331d8eb3a..ef5c46119 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -826,6 +826,13 @@ class Envs: SGLANG_OPT_USE_JIT_KERNEL_FUSED_TOPK = EnvBool(True) SGLANG_OPT_USE_TOPK_V2 = EnvBool(True) + # MiniMax-M3 sparse decode indexer: single JIT radix-select kernel replaces the 2-stage split-K Triton topk. + SGLANG_OPT_USE_MINIMAX_DECODE_TOPK_RADIX = EnvBool(True) + + # MiniMax-M3 MXFP8 MoE experimental fusion toggles (default off; A/B only). + SGLANG_MINIMAX_M3_FUSED_SWIGLU_MXFP8 = EnvBool(False) + SGLANG_MINIMAX_M3_FUSED_MOE_COMBINE = EnvBool(False) + # GEMM / kernel fusion SGLANG_OPT_FP8_WO_A_GEMM = EnvBool(True) SGLANG_OPT_BF16_FP32_GEMM_ALGO = EnvStr("cublas") diff --git a/python/sglang/srt/layers/attention/minimax_sparse_ops/common/index.py b/python/sglang/srt/layers/attention/minimax_sparse_ops/common/index.py new file mode 100644 index 000000000..8f3b8d6de --- /dev/null +++ b/python/sglang/srt/layers/attention/minimax_sparse_ops/common/index.py @@ -0,0 +1,67 @@ +import torch + + +def topk_index_reduce(tensor: torch.Tensor, dim: int) -> torch.Tensor: + """ + Reduces a specific dimension by computing the union of all top-k indices along that dimension. + The resulting tensor will have the 'dim' removed, and the last dimension expanded. + + Example: + Input: [10, num_heads, seq_len, max_topk] with dim=0 + Output: [num_heads, seq_len, 10 * max_topk] (Left-aligned, padded with -1) + + Args: + tensor (torch.Tensor): Input tensor of shape [..., dim_size, ..., max_topk]. + dim (int): The dimension to reduce (collapse). + + Returns: + torch.Tensor: Reduced tensor with shape [..., new_max_topk]. + Where new_max_topk = dim_size * original_max_topk. + """ + # 1. Shape Transformation + # We need to merge the target 'dim' with the last dimension 'max_topk'. + # Step A: Move the target dim to the second-to-last position (-2). + # e.g., [10, H, S, K] (dim=0) -> [H, S, 10, K] + tensor_permuted = torch.movedim(tensor, source=dim, destination=-2) + + # Step B: Flatten the last two dimensions. + # e.g., [H, S, 10, K] -> [H, S, 10 * K] + combined = tensor_permuted.flatten(start_dim=-2) + + # --- The following logic is identical to 'topk_index_union' --- + + # 2. Sort row-wise. + # Groups identical values together. -1 padding sorts to the left. + sorted_vals, _ = combined.sort(dim=-1) + + # 3. Deduplication (Delta Check). + # Keep value if it differs from the previous one. + is_new_element = sorted_vals[..., 1:] != sorted_vals[..., :-1] + + # First column is always new + first_col_true = torch.ones_like(sorted_vals[..., :1], dtype=torch.bool) + non_duplicate_mask = torch.cat([first_col_true, is_new_element], dim=-1) + + # 4. Filter. + # Valid if non-duplicate AND not -1 padding. + valid_mask = non_duplicate_mask & (sorted_vals != -1) + + # 5. Packing (Left-Alignment). + # Move valid elements to the left. + sort_idx = torch.argsort((~valid_mask).int(), dim=-1, stable=True) + result = torch.gather(sorted_vals, -1, sort_idx) + + # 6. Re-masking the right side. + # Fill garbage values on the right with -1. + valid_count = valid_mask.sum(dim=-1, keepdim=True) + total_cols = result.size(-1) + + # Broadcasting check: + # valid_count shape: [..., 1] + # idx_range shape: [total_cols] + # result shape: [..., total_cols] + idx_range = torch.arange(total_cols, device=tensor.device) + + final_result = torch.where(idx_range < valid_count, result, -1) + + return final_result diff --git a/python/sglang/srt/layers/attention/minimax_sparse_ops/common/utils.py b/python/sglang/srt/layers/attention/minimax_sparse_ops/common/utils.py new file mode 100644 index 000000000..e3694cbb9 --- /dev/null +++ b/python/sglang/srt/layers/attention/minimax_sparse_ops/common/utils.py @@ -0,0 +1,262 @@ +# Copyright 2025 XunhaoLai. All rights reserved. + +import functools +from collections import deque +from typing import Any, Callable, List, Optional, Tuple + +import torch +import triton +import triton.language as tl + +_tma_keep_alive_buf = deque(maxlen=200) + +# Q is always bf16/fp16. The paged main K/V cache may be fp8 (unit-scaled) under +# --kv-cache-dtype fp8_*; the kernel widens it to the Q dtype on load (IS_FP8 +# branch). Accepted on both HIP and CUDA (the bf16->fp8 cache write is unit-scaled, +# so the widening cast is the exact inverse dequant). The bf16/fp16-only MSA +# (fmha_sm100) kernel is excluded for fp8 KV by the backend use_msa gate. +SPARSE_KV_FP8_DTYPES = ( + torch.float8_e4m3fn, + torch.float8_e5m2, + torch.float8_e4m3fnuz, +) + + +def check_sparse_kv_fp8( + q: torch.Tensor, + k_cache: torch.Tensor, + v_cache: torch.Tensor, + *, + label: str, +) -> bool: + """Validate the sparse-attention KV cache dtype contract. + + Returns True iff the K cache is fp8 (then widened to Q dtype in the kernel). + Raises AssertionError otherwise, mirroring the contract the decode and prefill + topk kernels both enforce. fp8 is accepted on both HIP and CUDA. + """ + assert q.dtype in (torch.bfloat16, torch.float16) + is_fp8 = k_cache.dtype in SPARSE_KV_FP8_DTYPES + assert k_cache.dtype == q.dtype or is_fp8, ( + f"sparse {label} expects K cache dtype == Q dtype ({q.dtype}) " + f"or fp8, got {k_cache.dtype}" + ) + assert v_cache.dtype == k_cache.dtype + return is_fp8 + + +try: + make_tensor_descriptor = tl.make_tensor_descriptor +except Exception: + make_tensor_descriptor = tl._experimental_make_tensor_descriptor + + +def robust_allocator(size: int, alignment: int, stream: int = None): + """Allocator for Triton TMA descriptors. + + We keep reference in deque to prevent GC from collecting the buffer. + """ + tensor = torch.empty(size, device="cuda", dtype=torch.uint8) + _tma_keep_alive_buf.append(tensor) + return tensor + + +def tensor_cache(maxsize: int = 8): + """ + Cache function results using identity comparison. + Zero-overhead cache hit: no hash, no DtoH, just pointer comparison. + + Args: + maxsize: Maximum number of cached entries. Supports multi-GPU scenarios + where different devices have different tensor arguments. + """ + + def decorator(fn: Callable[..., Any]) -> Callable[..., Any]: + # LRU-style cache: list of (args, kwargs, result) tuples + # Most recently used at the end + _cache: list = [] + + def _args_match(args: tuple, cached_args: tuple) -> bool: + if len(args) != len(cached_args): + return False + for i in range(len(args)): + if args[i] is not cached_args[i]: + return False + return True + + def _kwargs_match(kwargs: dict, cached_kwargs: dict) -> bool: + if not kwargs and not cached_kwargs: + return True + if kwargs.keys() != cached_kwargs.keys(): + return False + for k, v in kwargs.items(): + if v is not cached_kwargs[k]: + return False + return True + + @functools.wraps(fn) + def wrapper(*args: Any, **kwargs: Any) -> Any: + # Search cache (most recent first for better hit rate) + for i in range(len(_cache) - 1, -1, -1): + cached_args, cached_kwargs, cached_result = _cache[i] + if _args_match(args, cached_args) and _kwargs_match( + kwargs, cached_kwargs + ): + # Move to end (most recently used) + if i != len(_cache) - 1: + _cache.append(_cache.pop(i)) + return cached_result + + # Cache miss + result = fn(*args, **kwargs) + + # Add to cache + if len(_cache) >= maxsize: + _cache.pop(0) # Remove oldest + _cache.append((args, kwargs, result)) + return result + + # Expose cache for manual clearing if needed + wrapper.cache_clear = lambda: _cache.clear() + wrapper.cache_info = lambda: {"size": len(_cache), "maxsize": maxsize} + + return wrapper + + return decorator + + +@tensor_cache(maxsize=8) +def get_cu_seqblocks( + cu_seqlens: torch.Tensor, + max_seqlen: int, + block_size_q: int, + block_size_k: int, + seqlens_cpu: Optional[List[int]] = None, +) -> Tuple[torch.Tensor, int, int, torch.Tensor, int, int]: + """Compute cumulative sequence block indices for blocked sparse attention. + + Converts token-level cumulative sequence lengths to block-level indices, + which are needed for block-sparse attention kernels. + + Note: + Results are cached (maxsize=8) based on input arguments. Repeated calls + with the same cu_seqlens, max_seqlen, and block sizes will return cached + results without recomputation. + + Args: + cu_seqlens: Cumulative sequence lengths. Shape: [batch_size + 1], dtype: int32. + max_seqlen: Maximum sequence length in the batch. + block_size_q: Query block size. + block_size_k: Key-value block size. + seqlens_cpu: Optional host copy of ``torch.diff(cu_seqlens)``; when given, + ``all_seqblock_q/k`` are summed on the host to avoid a per-layer sync. + + Returns: + A tuple of 6 values: + - cu_seqblocks_q: Cumulative query block indices. Shape: [batch_size + 1] + - max_seqblock_q: Maximum number of query blocks per sequence. + - all_seqblock_q: Total number of query blocks across all sequences. + - cu_seqblocks_k: Cumulative key block indices. Shape: [batch_size + 1] + - max_seqblock_k: Maximum number of key blocks per sequence. + - all_seqblock_k: Total number of key blocks across all sequences. + """ + cu_seqblocks_q = torch.zeros_like(cu_seqlens) + cu_seqblocks_k = torch.zeros_like(cu_seqlens) + seq_lens = torch.diff(cu_seqlens) + seqblocks_q = (seq_lens + block_size_q - 1) // block_size_q + seqblocks_k = (seq_lens + block_size_k - 1) // block_size_k + max_seqblock_q = (max_seqlen + block_size_q - 1) // block_size_q + max_seqblock_k = (max_seqlen + block_size_k - 1) // block_size_k + cu_seqblocks_q[1:] = seqblocks_q + cu_seqblocks_k[1:] = seqblocks_k + cu_seqblocks_q.cumsum_(0) + cu_seqblocks_k.cumsum_(0) + if seqlens_cpu is not None: + # Bit-identical to seqblocks.sum().item() but no device->host sync. + all_seqblock_q = sum( + (s + block_size_q - 1) // block_size_q for s in seqlens_cpu + ) + all_seqblock_k = sum( + (s + block_size_k - 1) // block_size_k for s in seqlens_cpu + ) + else: + all_seqblock_q = seqblocks_q.sum().item() + all_seqblock_k = seqblocks_k.sum().item() + return ( + cu_seqblocks_q, + max_seqblock_q, + all_seqblock_q, + cu_seqblocks_k, + max_seqblock_k, + all_seqblock_k, + ) + + +# Bitonic-sort compare-and-swap primitives shared by the decode and prefill +# topk-index kernels. Identical copies previously lived in both +# decode/flash_with_topk_idx.py and prefill/flash_with_topk_idx.py. +@triton.jit +def _compare_and_swap( + x, + ids, + flip, + i: tl.constexpr, + n_dims: tl.constexpr, +): + n_outer: tl.constexpr = x.numel >> n_dims + shape: tl.constexpr = [n_outer * 2**i, 2, 2 ** (n_dims - i - 1)] + y = tl.reshape(x, shape) + # slice left/right with 'stride' 2**(n_dims - i - 1) + mask = tl.arange(0, 2)[None, :, None] + left = tl.broadcast_to(tl.sum(y * (1 - mask), 1)[:, None, :], shape).to(y.dtype) + right = tl.broadcast_to(tl.sum(y * mask, 1)[:, None, :], shape).to(y.dtype) + left = tl.reshape(left, x.shape) + right = tl.reshape(right, x.shape) + # idx + y_idx = tl.reshape(ids, shape) + left_idx = tl.broadcast_to(tl.sum(y_idx * (1 - mask), 1)[:, None, :], shape) + right_idx = tl.broadcast_to(tl.sum(y_idx * mask, 1)[:, None, :], shape) + left_idx = tl.reshape(left_idx, x.shape).to(y_idx.dtype) + right_idx = tl.reshape(right_idx, x.shape).to(y_idx.dtype) + # actual compare-and-swap + idtype = tl.core.get_int_dtype(bitwidth=x.dtype.primitive_bitwidth, signed=True) + ileft = left.to(idtype, bitcast=True) + iright = right.to(idtype, bitcast=True) + ix = x.to(idtype, bitcast=True) + + cond = (left > right) != flip + ret = ix ^ tl.where(cond, ileft ^ iright, tl.zeros_like(ix)) + new_ids = ids ^ tl.where(cond, left_idx ^ right_idx, tl.zeros_like(ids)) + return ret.to(x.dtype, bitcast=True), new_ids + + +@triton.jit +def _bitonic_merge( + x, + ids, + stage: tl.constexpr, + order: tl.constexpr, + n_dims: tl.constexpr, +): + n_outer: tl.constexpr = x.numel >> n_dims + tl.static_assert(stage <= n_dims) + # flip denotes whether to re-arrange sub-sequences of elements in ascending or + # descending order. + # if flip = 00000000... then all elements will be re-arranged ascendingly at this stage + # if flip = 00110011... then all the elements will be re-arranged alternatingly (with + # a stride of 2) at this stage + if order == 2: + shape: tl.constexpr = [ + n_outer * 2 ** (n_dims - 1 - stage), + 2, + 2**stage, + ] + flip = tl.reshape( + tl.broadcast_to(tl.arange(0, 2)[None, :, None], shape), x.shape + ) + else: + flip = order + # perform `stage` rounds of `compare-and-swap` + for i in tl.static_range(stage): + x, ids = _compare_and_swap(x, ids, flip, i + (n_dims - stage), n_dims) + return x, ids diff --git a/python/sglang/srt/layers/attention/minimax_sparse_ops/decode/flash_with_topk_idx.py b/python/sglang/srt/layers/attention/minimax_sparse_ops/decode/flash_with_topk_idx.py new file mode 100644 index 000000000..5ec46f3d5 --- /dev/null +++ b/python/sglang/srt/layers/attention/minimax_sparse_ops/decode/flash_with_topk_idx.py @@ -0,0 +1,1065 @@ +# Copyright 2025 XunhaoLai. All rights reserved. + +from typing import Optional + +import torch +import triton +import triton.language as tl + +from sglang.srt.environ import envs + +from ..common.utils import _bitonic_merge, robust_allocator + + +@triton.heuristics( + { + "BLOCK_SIZE_H": lambda args: max( + 16, triton.next_power_of_2(args["gqa_group_size"]) + ), + "BLOCK_SIZE_D": lambda args: triton.next_power_of_2(args["head_dim"]), + "BATCH_SIZE_BUCKET": lambda args: triton.next_power_of_2(args["batch_size"]), + } +) +@triton.autotune( + configs=[ + triton.Config({"BLOCK_SIZE_N": BN}, num_warps=nw, num_stages=ns) + for BN in [64, 128, 256, 512] + for nw in [4, 8, 16] + for ns in [1, 2, 3] + ], + key=[ + "BATCH_SIZE_BUCKET", + "gqa_group_size", + "head_dim", + "block_size", + "SCORE_TYPE", + ], +) +@triton.jit +def _decode_score_kernel( + q_ptr, # Q: b x qh x d + k_cache_ptr, # K paged: max_slots x kh x d + req_to_token_ptr, # req_to_token: max_reqs x max_kv_len + score_ptr, # Score: qh x b x max_seqblock + seq_lens, + slot_ids, + # shape + max_slots, + batch_size, + gqa_group_size, + head_dim, + # block size + block_size: tl.constexpr, + topk: tl.constexpr, + # sm_scale + sm_scale, + # init and local blocks + init_blocks, + local_blocks, + # stride + stride_q_b, + stride_q_h, + stride_q_d, + stride_k_s, + stride_k_h, + stride_k_d, + stride_r2t_b, + stride_s_h, + stride_s_b, + stride_s_n, + # META parameters + BATCH_SIZE_BUCKET: tl.constexpr, + BLOCK_SIZE_H: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, + BLOCK_SIZE_D: tl.constexpr, + NUM_KV_CHUNKS: tl.constexpr, + SCORE_TYPE: tl.constexpr, + SKIP_TRIVIAL_TOPK_SCORE: tl.constexpr, +): + tl.static_assert(SCORE_TYPE == "max" or SCORE_TYPE == "lse") + sm_scale_log2e = sm_scale * 1.4426950409 + tl.static_assert(BLOCK_SIZE_N >= block_size) + BLOCKS_PER_K_BLOCK: tl.constexpr = BLOCK_SIZE_N // block_size + # get batch id and head id + pid_bc, pid_kh = tl.program_id(0), tl.program_id(1) + pid_b = pid_bc % batch_size + pid_c = pid_bc // batch_size + pid_h = pid_kh * gqa_group_size + # block-aligned fixed-count chunked decode (grid independent of seq_len for cuda graph) + seq_len = tl.load(seq_lens + pid_b).to(tl.int32) + num_blocks = (seq_len + block_size - 1) // block_size + if SKIP_TRIVIAL_TOPK_SCORE: + if num_blocks <= topk: + return + chunk_size_blocks = tl.cdiv(num_blocks, NUM_KV_CHUNKS) + chunk_start_block = pid_c * chunk_size_blocks + chunk_end_block = tl.minimum(chunk_start_block + chunk_size_blocks, num_blocks) + chunk_start = chunk_start_block * block_size + chunk_end = tl.minimum(chunk_end_block * block_size, seq_len) + if chunk_start_block >= chunk_end_block: + return + sid = (tl.load(slot_ids + pid_b).to(tl.int64) + max_slots) % max_slots + # init qkv pointer + q_ptrs = tl.make_block_ptr( + base=q_ptr + pid_b * stride_q_b + pid_h * stride_q_h, + shape=(gqa_group_size, head_dim), + strides=(stride_q_h, stride_q_d), + offsets=(0, 0), + block_shape=(BLOCK_SIZE_H, BLOCK_SIZE_D), + order=(1, 0), + ) + s_ptrs = tl.make_block_ptr( + base=score_ptr + pid_b * stride_s_b + pid_h * stride_s_h, + shape=(gqa_group_size, chunk_end_block), + strides=(stride_s_h, stride_s_n), + offsets=(0, chunk_start_block), + block_shape=(BLOCK_SIZE_H, BLOCKS_PER_K_BLOCK), + order=(1, 0), + ) + # load q + q = tl.load(q_ptrs, boundary_check=(0, 1), padding_option="zero") + # init statistics + off_n = tl.arange(0, BLOCK_SIZE_N) + off_d = tl.arange(0, BLOCK_SIZE_D) + off_bpk = tl.arange(0, BLOCKS_PER_K_BLOCK) + dim_mask = off_d < head_dim + # pre-compute local_start outside the loop (constant for entire batch) + local_start = tl.maximum(0, num_blocks - local_blocks) + # prefetch first iteration's slots + r2t_base = req_to_token_ptr + sid * stride_r2t_b + prefetch_pos = chunk_start + off_n + prefetch_mask = prefetch_pos < seq_len + prefetched_slots = tl.load( + r2t_base + prefetch_pos, + mask=prefetch_mask, + other=0, + ).to(tl.int64) + # score-only: compute block scores without loading V + for i in range(chunk_start, chunk_end, BLOCK_SIZE_N): + pos_mask = prefetch_mask + slots = prefetched_slots + # prefetch next iteration's slots + next_i = i + BLOCK_SIZE_N + if next_i < chunk_end: + next_pos = next_i + off_n + prefetch_mask = next_pos < seq_len + prefetched_slots = tl.load( + r2t_base + next_pos, + mask=prefetch_mask, + other=0, + ).to(tl.int64) + slots = (slots + max_slots) % max_slots + # load K as (head_dim, BLOCK_SIZE_N) via indirect addressing + k_off = ( + slots[None, :] * stride_k_s + + pid_kh * stride_k_h + + off_d[:, None] * stride_k_d + ) + k = tl.load( + k_cache_ptr + k_off, + mask=dim_mask[:, None] & pos_mask[None, :], + other=0.0, + ) + # compute qk + qk = tl.zeros((BLOCK_SIZE_H, BLOCK_SIZE_N), dtype=tl.float32) + qk += tl.where(off_n[None, :] < chunk_end - i, 0, float("-inf")) + # [H, D], [D, N] -> [H, N] + qk += tl.dot(q, k) * sm_scale_log2e + # save qk to score + score = tl.reshape( + qk, + (BLOCK_SIZE_H, BLOCKS_PER_K_BLOCK, block_size), + can_reorder=False, + ) + sub_max = tl.max(score, axis=2) + if SCORE_TYPE == "max": + score = sub_max + else: # "lse" + score = sub_max + tl.log2( + tl.sum(tl.exp2(score - sub_max[:, :, None]), axis=2) + ) + score = tl.where(score != score, float("-inf"), score) + # apply init_blocks and local_blocks efficiently + curr_block_idx = i // block_size + off_bpk + is_init = curr_block_idx < init_blocks + is_local = (curr_block_idx >= local_start) & (curr_block_idx < num_blocks) + score = tl.where( + is_local[None, :], 1e29, tl.where(is_init[None, :], 1e30, score) + ) + tl.store(s_ptrs, score.to(score_ptr.dtype.element_ty), boundary_check=(0, 1)) + # update ptrs + s_ptrs = tl.advance(s_ptrs, (0, BLOCKS_PER_K_BLOCK)) + + +@triton.heuristics( + { + "BLOCK_SIZE_H": lambda args: max( + 16, triton.next_power_of_2(args["gqa_group_size"]) + ), + "BLOCK_SIZE_D": lambda args: triton.next_power_of_2(args["head_dim"]), + "HAS_SINK": lambda args: args["sink_ptr"] is not None, + } +) +@triton.autotune( + configs=[ + triton.Config({"BLOCK_SIZE_N": BN}, num_warps=nw, num_stages=ns) + for BN in [64, 128, 256, 512] + for nw in [4, 8, 16] + for ns in [1, 2, 3] + ], + key=[ + "gqa_group_size", + "head_dim", + "block_size", + "HAS_SINK", + "SCORE_TYPE", + ], +) +@triton.jit +def _decode_score_attn_kernel( + q_ptr, # Q: b x qh x d + sink_ptr, # Sink: qh x d + k_cache_ptr, # K paged: max_slots x kh x d + v_cache_ptr, # V paged: max_slots x kh x d + req_to_token_ptr, # req_to_token: max_reqs x max_kv_len + o_ptr, # O: c x b x qh x d + lse_ptr, # lse: c x b x qh + score_ptr, # Score: qh x b x max_seqblock + seq_lens, + slot_ids, + # shape + max_slots, + batch_size, + gqa_group_size, + head_dim, + # block size + block_size: tl.constexpr, + topk: tl.constexpr, + # sm_scale + sm_scale, + # init and local blocks + init_blocks, + local_blocks, + # stride + stride_q_b, + stride_q_h, + stride_q_d, + stride_sink_h, + stride_sink_d, + stride_k_s, + stride_k_h, + stride_k_d, + stride_v_s, + stride_v_h, + stride_v_d, + stride_r2t_b, + stride_o_c, + stride_o_b, + stride_o_h, + stride_o_d, + stride_l_c, + stride_l_b, + stride_l_h, + stride_s_h, + stride_s_b, + stride_s_n, + # META parameters + BLOCK_SIZE_H: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, + BLOCK_SIZE_D: tl.constexpr, + NUM_KV_CHUNKS: tl.constexpr, + HAS_SINK: tl.constexpr, + SCORE_TYPE: tl.constexpr, + SKIP_TRIVIAL_TOPK_SCORE: tl.constexpr, +): + tl.static_assert(SCORE_TYPE == "max" or SCORE_TYPE == "lse") + sm_scale_log2e = sm_scale * 1.4426950409 + tl.static_assert(BLOCK_SIZE_N >= block_size) + BLOCKS_PER_K_BLOCK: tl.constexpr = BLOCK_SIZE_N // block_size + # get batch id and head id + pid_bc, pid_kh = tl.program_id(0), tl.program_id(1) + pid_b = pid_bc % batch_size + pid_c = pid_bc // batch_size + pid_h = pid_kh * gqa_group_size + # block-aligned fixed-count chunked decode (grid independent of seq_len for cuda graph) + seq_len = tl.load(seq_lens + pid_b).to(tl.int32) + num_blocks = (seq_len + block_size - 1) // block_size + chunk_size_blocks = tl.cdiv(num_blocks, NUM_KV_CHUNKS) + chunk_start_block = pid_c * chunk_size_blocks + chunk_end_block = tl.minimum(chunk_start_block + chunk_size_blocks, num_blocks) + chunk_start = chunk_start_block * block_size + chunk_end = tl.minimum(chunk_end_block * block_size, seq_len) + if chunk_start_block >= chunk_end_block: + return + sid = ( + tl.load(slot_ids + pid_b).to(tl.int64) + max_slots + ) % max_slots # to avoid bugs when slot_ids is negative + # init qkv pointer + q_ptrs = tl.make_block_ptr( + base=q_ptr + pid_b * stride_q_b + pid_h * stride_q_h, + shape=(gqa_group_size, head_dim), + strides=(stride_q_h, stride_q_d), + offsets=(0, 0), + block_shape=(BLOCK_SIZE_H, BLOCK_SIZE_D), + order=(1, 0), + ) + s_ptrs = tl.make_block_ptr( + base=score_ptr + pid_b * stride_s_b + pid_h * stride_s_h, + shape=(gqa_group_size, chunk_end_block), + strides=(stride_s_h, stride_s_n), + offsets=(0, chunk_start_block), + block_shape=(BLOCK_SIZE_H, BLOCKS_PER_K_BLOCK), + order=(1, 0), + ) + # load q + q = tl.load(q_ptrs, boundary_check=(0, 1), padding_option="zero") + # init statistics + off_n = tl.arange(0, BLOCK_SIZE_N) + off_d = tl.arange(0, BLOCK_SIZE_D) + off_bpk = tl.arange(0, BLOCKS_PER_K_BLOCK) + dim_mask = off_d < head_dim + if HAS_SINK: + if pid_c == 0: + sink_ptrs = tl.make_block_ptr( + base=sink_ptr + pid_h * stride_sink_h, + shape=(gqa_group_size, head_dim), + strides=(stride_sink_h, stride_sink_d), + offsets=(0, 0), + block_shape=(BLOCK_SIZE_H, BLOCK_SIZE_D), + order=(1, 0), + ) + sink = tl.load(sink_ptrs, boundary_check=(0, 1), padding_option="zero").to( + tl.float32 + ) + qsink = ( + tl.sum(q.to(tl.float32) * sink, axis=1) * sm_scale_log2e + ) # (BLOCK_SIZE_H,) + m_i = qsink + l_i = tl.full((BLOCK_SIZE_H,), 1.0, dtype=tl.float32) + else: + m_i = tl.full((BLOCK_SIZE_H,), float("-inf"), dtype=tl.float32) + l_i = tl.full((BLOCK_SIZE_H,), 0.0, dtype=tl.float32) + else: + m_i = tl.full((BLOCK_SIZE_H,), float("-inf"), dtype=tl.float32) + l_i = tl.full((BLOCK_SIZE_H,), 0.0, dtype=tl.float32) + acc_o = tl.full((BLOCK_SIZE_H, BLOCK_SIZE_D), 0, dtype=tl.float32) + # pre-compute local_start outside the loop (constant for entire batch) + local_start = tl.maximum(0, num_blocks - local_blocks) + # full attention or causal attention + for i in range(chunk_start, chunk_end, BLOCK_SIZE_N): + pos = i + off_n + pos_mask = pos < seq_len + # resolve slot per position via req_to_token + slots = tl.load( + req_to_token_ptr + sid * stride_r2t_b + pos, + mask=pos_mask, + other=0, + ).to(tl.int64) + slots = (slots + max_slots) % max_slots # safety against negative + # load K as (head_dim, BLOCK_SIZE_N) via indirect addressing + k_off = ( + slots[None, :] * stride_k_s + + pid_kh * stride_k_h + + off_d[:, None] * stride_k_d + ) + k = tl.load( + k_cache_ptr + k_off, + mask=dim_mask[:, None] & pos_mask[None, :], + other=0.0, + ) + # load V as (BLOCK_SIZE_N, head_dim) via indirect addressing + v_off = ( + slots[:, None] * stride_v_s + + pid_kh * stride_v_h + + off_d[None, :] * stride_v_d + ) + v = tl.load( + v_cache_ptr + v_off, + mask=pos_mask[:, None] & dim_mask[None, :], + other=0.0, + ) + # compute qk + qk = tl.zeros((BLOCK_SIZE_H, BLOCK_SIZE_N), dtype=tl.float32) + qk += tl.where(off_n[None, :] < chunk_end - i, 0, float("-inf")) + # [H, D], [D, N] -> [H, N] + qk += tl.dot(q, k) * sm_scale_log2e + # save qk to score + score = tl.reshape( + qk, + (BLOCK_SIZE_H, BLOCKS_PER_K_BLOCK, block_size), + can_reorder=False, + ) + sub_max = tl.max(score, axis=2) + if SCORE_TYPE == "max": + score = sub_max + else: # "lse" + score = sub_max + tl.log2( + tl.sum(tl.exp2(score - sub_max[:, :, None]), axis=2) + ) + score = tl.where(score != score, float("-inf"), score) + # apply init_blocks and local_blocks efficiently + # current block indices: [BLOCKS_PER_K_BLOCK] + curr_block_idx = i // block_size + off_bpk + # Combined condition with single nested tl.where: + # - local_blocks (1e29) takes priority over init_blocks (1e30) + # - When init_blocks=0: curr_block_idx < 0 is always False + # - When local_blocks=0: local_start=num_blocks, so curr_block_idx >= num_blocks is always False + is_init = curr_block_idx < init_blocks + is_local = (curr_block_idx >= local_start) & (curr_block_idx < num_blocks) + # Apply in one fused operation: local > init > original + score = tl.where( + is_local[None, :], 1e29, tl.where(is_init[None, :], 1e30, score) + ) + if SKIP_TRIVIAL_TOPK_SCORE: + if num_blocks > topk: + tl.store( + s_ptrs, score.to(score_ptr.dtype.element_ty), boundary_check=(0, 1) + ) + else: + tl.store( + s_ptrs, score.to(score_ptr.dtype.element_ty), boundary_check=(0, 1) + ) + # max-of-max == max(qk), avoids re-scanning qk + m_ij = tl.maximum(m_i, tl.max(sub_max, axis=1)) + p = tl.exp2(qk - m_ij[:, None]) + l_ij = tl.sum(p, axis=1) + acc_o_scale = tl.exp2(m_i - m_ij) + acc_o = acc_o * acc_o_scale[:, None] + # [H, N], [N, D] -> [H, D] + acc_o += tl.dot(p.to(v.dtype), v) + m_i = m_ij + l_i = l_i * acc_o_scale + l_ij + # update ptrs + s_ptrs = tl.advance(s_ptrs, (0, BLOCKS_PER_K_BLOCK)) + # final scale + acc_o = acc_o / l_i[:, None] + # compute lse in log2 scale for merge kernel + lse_i = m_i + tl.log2(l_i) + # save output + o_ptrs = tl.make_block_ptr( + base=o_ptr + pid_c * stride_o_c + pid_b * stride_o_b + pid_h * stride_o_h, + shape=(gqa_group_size, head_dim), + strides=(stride_o_h, stride_o_d), + offsets=(0, 0), + block_shape=(BLOCK_SIZE_H, BLOCK_SIZE_D), + order=(1, 0), + ) + tl.store(o_ptrs, acc_o.to(o_ptr.dtype.element_ty), boundary_check=(0, 1)) + # save lse (in log2 scale, will be used directly in merge kernel) + lse_ptrs = tl.make_block_ptr( + base=lse_ptr + pid_c * stride_l_c + pid_b * stride_l_b + pid_h * stride_l_h, + shape=(gqa_group_size,), + strides=(stride_l_h,), + offsets=(0,), + block_shape=(BLOCK_SIZE_H,), + order=(0,), + ) + tl.store(lse_ptrs, lse_i.to(lse_ptr.dtype.element_ty), boundary_check=(0,)) + + +@triton.heuristics( + { + "BLOCK_SIZE_D": lambda args: triton.next_power_of_2(args["head_dim"]), + } +) +@triton.autotune( + configs=[ + triton.Config({}, num_warps=nw, num_stages=ns) + for nw in [4, 8] + for ns in [2, 3, 4] + ], + key=["BLOCK_SIZE_D"], + restore_value=["o_ptr"], +) +@triton.jit +def _merge_attn_out_kernel( + o_ptr, + lse_ptr, + # shape + seq_lens, + head_dim, + block_size: tl.constexpr, + # stride + stride_o_c, + stride_o_b, + stride_o_h, + stride_o_d, + stride_l_c, + stride_l_b, + stride_l_h, + NUM_KV_CHUNKS: tl.constexpr, + BLOCK_SIZE_D: tl.constexpr, +): + pid_b, pid_h = tl.program_id(0), tl.program_id(1) + # recompute chunk layout (must match _decode_score_attn_kernel) + seq_len = tl.load(seq_lens + pid_b) + num_blocks = (seq_len + block_size - 1) // block_size + chunk_size_blocks = tl.cdiv(num_blocks, NUM_KV_CHUNKS) + valid_chunks = tl.cdiv(num_blocks, chunk_size_blocks) + # ptrs + off_c = tl.arange(0, NUM_KV_CHUNKS) + off_d = tl.arange(0, BLOCK_SIZE_D) + o_ptrs = tl.make_block_ptr( + base=o_ptr + pid_b * stride_o_b + pid_h * stride_o_h, + shape=(valid_chunks, head_dim), + strides=(stride_o_c, stride_o_d), + offsets=(0, 0), + block_shape=(NUM_KV_CHUNKS, BLOCK_SIZE_D), + order=(1, 0), + ) + lse_ptrs = lse_ptr + pid_b * stride_l_b + pid_h * stride_l_h + off_c * stride_l_c + # load o and lse + o = tl.load(o_ptrs, boundary_check=(0, 1), padding_option="zero") + lse = tl.load(lse_ptrs, mask=off_c < valid_chunks, other=float("-inf")) + # merge o (lse is in log2 scale, use exp2/log2 for efficiency) + lse_shifted = lse - tl.max(lse, axis=0) + log_sum = tl.log2(tl.sum(tl.exp2(lse_shifted), axis=0)) + scale = tl.exp2(lse_shifted - log_sum) + o = tl.sum(o * scale[:, None], axis=0) + # save o + o_ptrs = o_ptr + pid_b * stride_o_b + pid_h * stride_o_h + off_d * stride_o_d + tl.store(o_ptrs, o.to(o_ptr.dtype.element_ty), mask=off_d < head_dim) + + +@triton.heuristics( + { + "BLOCK_SIZE_T": lambda args: triton.next_power_of_2(args["topk"]), + } +) +@triton.autotune( + configs=[ + triton.Config({"BLOCK_SIZE_K": 256}, num_warps=8, num_stages=2), + triton.Config({"BLOCK_SIZE_K": 256}, num_warps=4, num_stages=2), + triton.Config({"BLOCK_SIZE_K": 128}, num_warps=4, num_stages=2), + triton.Config({"BLOCK_SIZE_K": 128}, num_warps=4, num_stages=3), + triton.Config({"BLOCK_SIZE_K": 64}, num_warps=2, num_stages=2), + ], + key=["topk"], +) +@triton.jit +def _topk_index_partial_kernel( + s_ptr, # Score: [num_q_heads, BS, max_seqblock] + ts_partial_ptr, # Partial scores out: [NUM_TOPK_CHUNKS, num_q_heads, BS, BLOCK_SIZE_T] + ti_partial_ptr, # Partial idx out (1-indexed global, 0=invalid): same shape + seq_lens, + block_size: tl.constexpr, + topk: tl.constexpr, + chunk_blocks: tl.constexpr, # how many score-blocks each chunk owns + # strides + stride_s_h, + stride_s_b, + stride_s_k, + stride_ts_c, + stride_ts_h, + stride_ts_b, + stride_ts_t, + stride_ti_c, + stride_ti_h, + stride_ti_b, + stride_ti_t, + # META + BLOCK_SIZE_K: tl.constexpr, + BLOCK_SIZE_T: tl.constexpr, +): + tl.static_assert(BLOCK_SIZE_K > topk) + pid_b = tl.program_id(0) + pid_h = tl.program_id(1) + pid_chunk = tl.program_id(2) + + seq_len = tl.load(seq_lens + pid_b) + num_blocks = (seq_len + block_size - 1) // block_size + + # Slice this chunk owns within [0, num_blocks). + chunk_start = pid_chunk * chunk_blocks + chunk_end = tl.minimum(chunk_start + chunk_blocks, num_blocks) + chunk_actual = tl.maximum(chunk_end - chunk_start, 0) + + off_k = tl.arange(0, BLOCK_SIZE_K) + off_t = tl.arange(0, BLOCK_SIZE_T) + + # Pointer starts at chunk_start. + s_ptrs = ( + s_ptr + + pid_b * stride_s_b + + pid_h * stride_s_h + + (chunk_start + off_k) * stride_s_k + ) + + topk_score = tl.full((BLOCK_SIZE_K,), -1e30, dtype=tl.float32) + topk_idx = tl.full((BLOCK_SIZE_K,), 0, dtype=tl.int32) + left_half_mask = tl.arange(0, BLOCK_SIZE_K) < BLOCK_SIZE_K // 2 + + # Streaming top-K within this chunk. tl.range(0, 0) is a no-op so empty + # chunks (chunk_actual == 0) skip the body and store sentinel -1e30 / 0. + for i in tl.range(0, chunk_actual, BLOCK_SIZE_K): + mask = off_k < chunk_actual - i + score = tl.load(s_ptrs, mask=mask, other=-1e30).to(tl.float32) + score = tl.where(score != score, -1e30, score) + s_ptrs = s_ptrs + stride_s_k * BLOCK_SIZE_K + topk_score, last_topk_score = score, topk_score + topk_idx, last_topk_idx = ( + tl.where(mask, chunk_start + i + off_k + 1, 0), # 1-indexed global + topk_idx, + ) + n_dims: tl.constexpr = tl.standard._log2(BLOCK_SIZE_K) + for j in tl.static_range(1, n_dims): + topk_score, topk_idx = _bitonic_merge( + topk_score, topk_idx.to(tl.int32), j, 2, n_dims + ) + if i != 0: + topk_score, topk_idx = _bitonic_merge( + topk_score, topk_idx.to(tl.int32), n_dims, False, n_dims + ) + topk_score_new = last_topk_score * left_half_mask + topk_score * ( + 1 - left_half_mask + ) + topk_idx_new = last_topk_idx * left_half_mask + topk_idx * ( + 1 - left_half_mask + ) + topk_score, topk_idx = _bitonic_merge( + topk_score_new, topk_idx_new.to(tl.int32), n_dims, True, n_dims + ) + else: + topk_score, topk_idx = _bitonic_merge( + topk_score, topk_idx.to(tl.int32), n_dims, True, n_dims + ) + + # Extract first BLOCK_SIZE_T entries (top-K of this chunk after descending sort). + topk_mask_extract = tl.arange(0, BLOCK_SIZE_K // BLOCK_SIZE_T) == 0 + final_score = tl.sum( + topk_mask_extract[:, None] + * tl.reshape(topk_score, [BLOCK_SIZE_K // BLOCK_SIZE_T, BLOCK_SIZE_T]), + axis=0, + ) + final_idx = tl.sum( + topk_mask_extract[:, None] + * tl.reshape(topk_idx, [BLOCK_SIZE_K // BLOCK_SIZE_T, BLOCK_SIZE_T]), + axis=0, + ) + + # Store partial. Always write all BLOCK_SIZE_T slots — invalid slots carry + # -1e30 / 0 sentinels and lose to real scores in the merge stage. + ts_ptrs = ( + ts_partial_ptr + + pid_chunk * stride_ts_c + + pid_b * stride_ts_b + + pid_h * stride_ts_h + + off_t * stride_ts_t + ) + ti_ptrs = ( + ti_partial_ptr + + pid_chunk * stride_ti_c + + pid_b * stride_ti_b + + pid_h * stride_ti_h + + off_t * stride_ti_t + ) + tl.store(ts_ptrs, final_score) + tl.store(ti_ptrs, final_idx) + + +@triton.heuristics( + { + "BLOCK_SIZE_T": lambda args: triton.next_power_of_2(args["topk"]), + "BLOCK_SIZE_K": lambda args: triton.next_power_of_2( + args["NUM_TOPK_CHUNKS"] * triton.next_power_of_2(args["topk"]) + ), + } +) +@triton.jit +def _topk_index_merge_kernel( + ts_partial_ptr, # Partial scores: [NUM_TOPK_CHUNKS, num_q_heads, BS, BLOCK_SIZE_T] + ti_partial_ptr, # Partial idx (1-indexed global, 0=invalid): same shape + ti_final_ptr, # Final idx (0-indexed global, -1=invalid): [num_q_heads, BS, topk] + seq_lens, + block_size: tl.constexpr, + topk: tl.constexpr, + # strides + stride_ts_c, + stride_ts_h, + stride_ts_b, + stride_ts_t, + stride_ti_c, + stride_ti_h, + stride_ti_b, + stride_ti_t, + stride_tif_h, + stride_tif_b, + stride_tif_t, + # META + NUM_TOPK_CHUNKS: tl.constexpr, + BLOCK_SIZE_K: tl.constexpr, + BLOCK_SIZE_T: tl.constexpr, +): + pid_b = tl.program_id(0) + pid_h = tl.program_id(1) + + seq_len = tl.load(seq_lens + pid_b) + num_blocks = (seq_len + block_size - 1) // block_size + + # Layout: load NUM_TOPK_CHUNKS * BLOCK_SIZE_T candidates, padded to BLOCK_SIZE_K. + # candidate at flat position p comes from chunk = p // BLOCK_SIZE_T, + # in_chunk = p % BLOCK_SIZE_T. + off = tl.arange(0, BLOCK_SIZE_K) + chunk_idx = off // BLOCK_SIZE_T + in_chunk_idx = off % BLOCK_SIZE_T + valid = chunk_idx < NUM_TOPK_CHUNKS + + score_offset = ( + chunk_idx * stride_ts_c + + pid_h * stride_ts_h + + pid_b * stride_ts_b + + in_chunk_idx * stride_ts_t + ) + idx_offset = ( + chunk_idx * stride_ti_c + + pid_h * stride_ti_h + + pid_b * stride_ti_b + + in_chunk_idx * stride_ti_t + ) + + score = tl.load(ts_partial_ptr + score_offset, mask=valid, other=-1e30).to( + tl.float32 + ) + score = tl.where(score != score, -1e30, score) + idx = tl.load(ti_partial_ptr + idx_offset, mask=valid, other=0).to(tl.int32) + + # Full bitonic descending sort of BLOCK_SIZE_K items. + n_dims: tl.constexpr = tl.standard._log2(BLOCK_SIZE_K) + for j in tl.static_range(1, n_dims): + score, idx = _bitonic_merge(score, idx.to(tl.int32), j, 2, n_dims) + score, idx = _bitonic_merge(score, idx.to(tl.int32), n_dims, True, n_dims) + + # Extract first BLOCK_SIZE_T positions — these are the global top-K. + extract_mask = tl.arange(0, BLOCK_SIZE_K // BLOCK_SIZE_T) == 0 + topk_idx_final = tl.sum( + extract_mask[:, None] + * tl.reshape(idx - 1, [BLOCK_SIZE_K // BLOCK_SIZE_T, BLOCK_SIZE_T]), + axis=0, + ) + + off_t = tl.arange(0, BLOCK_SIZE_T) + tif_ptrs = ( + ti_final_ptr + + pid_h * stride_tif_h + + pid_b * stride_tif_b + + off_t * stride_tif_t + ) + topk_idx_final = tl.where(off_t < tl.minimum(topk, num_blocks), topk_idx_final, -1) + tl.store(tif_ptrs, topk_idx_final.to(ti_final_ptr.dtype.element_ty)) + + +@torch.no_grad() +def flash_decode_with_topk_idx( + q: torch.Tensor, # [batch_size, num_heads, head_dim] + sink: Optional[torch.Tensor], + k_cache: torch.Tensor, # [max_slots, num_kv_heads, head_dim] (paged) + v_cache: Optional[torch.Tensor], # paged; ignored when disable_index_value=True + req_to_token: torch.Tensor, # [max_reqs, max_kv_len] + seq_lens: torch.Tensor, # [batch_size, ] + max_seqlen: int, + slot_ids: torch.Tensor, # [batch_size, ] + block_size: int, + topk: int, + init_blocks: int, + local_blocks: int, + sm_scale: Optional[float] = None, + use_tma: bool = True, + score_type: str = "max", + disable_index_value: bool = False, + use_dense_main_attn: bool = False, # NOTE: need transform idx in this case + page_size: int = 1, +) -> torch.Tensor: + assert score_type in ( + "max", + "lse", + ), f"score_type must be 'max' or 'lse', got {score_type!r}" + triton.set_allocator(robust_allocator) + # dtype check + assert ( + q.dtype == torch.bfloat16 + or q.dtype == torch.float16 + and k_cache.dtype == q.dtype + ) + if not disable_index_value: + assert v_cache is not None + # shape + batch_size, num_q_heads, head_dim = q.shape + max_slots, num_kv_heads, _ = k_cache.shape + max_kv_len = req_to_token.shape[1] + assert slot_ids.shape[0] == batch_size and seq_lens.shape[0] == batch_size + # gqa + assert num_q_heads % num_kv_heads == 0 + gqa_group_size = num_q_heads // num_kv_heads + # sm scale + if sm_scale is None: + sm_scale = head_dim**-0.5 + # NUM_KV_CHUNKS controls how many parallel chunks each (batch, kv_head) gets. + # Total CTAs = batch_size * NUM_KV_CHUNKS * num_kv_heads. + # TARGET_GRID is the desired total CTA count; NUM_KV_CHUNKS is derived by: + # NUM_KV_CHUNKS = clamp(TARGET_GRID // (BS * num_kv_heads), 1, MAX_NUM_KV_CHUNKS) + # then rounded down to the nearest power of 2 (merge kernel requires it). + # Must only depend on cuda-graph-constant quantities (BS, num_kv_heads), not seq_len. + # Empty chunks early-return cheaply, so over-chunking is nearly free. + # E.g. with num_kv_heads=1: BS=1→NKC=256,CTAs=256; BS=32→NKC=128,CTAs=4096. + TARGET_GRID = 4096 + MAX_NUM_KV_CHUNKS = 256 + target = max( + 1, + min(MAX_NUM_KV_CHUNKS, TARGET_GRID // max(1, batch_size * num_kv_heads)), + ) + NUM_KV_CHUNKS = 1 << (target.bit_length() - 1) + score_kv_len = min(max_seqlen, max_kv_len) + # The score producers below write every valid block column + # [0, ceil(seq_len / block_size)) for each (head, batch) row. All consumers + # clamp their scan to the same per-row valid block count, so columns beyond + # seq_len are never read. Avoid a full-tensor -inf memset here; on CUDA graph + # capture the static score shape can be much larger than the live context. + score = torch.empty( + (num_q_heads, batch_size, triton.cdiv(score_kv_len, block_size)), + dtype=torch.float32, + device=q.device, + ) + use_jit_topk = ( + envs.SGLANG_OPT_USE_MINIMAX_DECODE_TOPK_RADIX.get() + and score.shape[2] <= 4096 + and topk <= 32 + ) + # If the live context has <= topk sparse blocks, the downstream dense + # page-table/JIT top-k kernels select every block from seq_lens directly + # without reading score. Keep this gate in sync with the consumers below: + # _skip_block_topk/use_dense_main_attn and use_jit_topk special-case + # num_blocks <= topk, while the Triton fallback reads score and must not + # skip these writes. + skip_trivial_topk_score = use_dense_main_attn or use_jit_topk + + grid = (batch_size * NUM_KV_CHUNKS, num_kv_heads) + if disable_index_value: + _decode_score_kernel[grid]( + q, + k_cache, + req_to_token, + score, + seq_lens, + slot_ids, + max_slots, + batch_size, + gqa_group_size, + head_dim, + block_size, + topk, + sm_scale, + init_blocks, + local_blocks, + q.stride(0), + q.stride(1), + q.stride(2), + k_cache.stride(0), + k_cache.stride(1), + k_cache.stride(2), + req_to_token.stride(0), + score.stride(0), + score.stride(1), + score.stride(2), + NUM_KV_CHUNKS=NUM_KV_CHUNKS, + SCORE_TYPE=score_type, + SKIP_TRIVIAL_TOPK_SCORE=skip_trivial_topk_score, + ) + else: + assert v_cache is not None + o = torch.empty( + NUM_KV_CHUNKS, + batch_size, + num_q_heads, + head_dim, + dtype=q.dtype, + device=q.device, + ) + lse = torch.empty( + NUM_KV_CHUNKS, batch_size, num_q_heads, dtype=torch.float32, device=q.device + ) + _decode_score_attn_kernel[grid]( + q, + sink, + k_cache, + v_cache, + req_to_token, + o, + lse, + score, + seq_lens, + slot_ids, + max_slots, + batch_size, + gqa_group_size, + head_dim, + block_size, + topk, + sm_scale, + init_blocks, + local_blocks, + q.stride(0), + q.stride(1), + q.stride(2), + sink.stride(0) if sink is not None else 0, + sink.stride(1) if sink is not None else 0, + k_cache.stride(0), + k_cache.stride(1), + k_cache.stride(2), + v_cache.stride(0), + v_cache.stride(1), + v_cache.stride(2), + req_to_token.stride(0), + o.stride(0), + o.stride(1), + o.stride(2), + o.stride(3), + lse.stride(0), + lse.stride(1), + lse.stride(2), + score.stride(0), + score.stride(1), + score.stride(2), + NUM_KV_CHUNKS=NUM_KV_CHUNKS, + SCORE_TYPE=score_type, + SKIP_TRIVIAL_TOPK_SCORE=skip_trivial_topk_score, + ) + # Fused top-k + page-table transform: emit the dense backend's page table + # directly (page-size-aware) instead of block ids, skipping a separate gather. + # The page table + per-query effective KV length are allocated and returned. + real_seq_lens = None + if use_dense_main_attn: + from sglang.jit_kernel.minimax_decode_topk import ( + minimax_decode_topk_page_table, + ) + + # num_q_heads (= num index/kv heads) may be > 1 for DP attention; the + # kernel emits a flattened [batch*num_q_heads] page table + seq_lens with + # the kv head head-encoded (head-minor) into the page index. + topk_idx, real_seq_lens = minimax_decode_topk_page_table( + score, seq_lens, req_to_token, slot_ids, block_size, topk, page_size + ) + if disable_index_value: + return None, topk_idx, real_seq_lens + # fall through to the idx_o (index value) merge, then return + _skip_block_topk = True + else: + _skip_block_topk = False + + # get topk index + if not _skip_block_topk: + topk_idx = torch.empty( + (num_q_heads, batch_size, topk), + device=score.device, + dtype=torch.int32, + ) + + if _skip_block_topk: + pass + elif use_jit_topk: + # Single-stage JIT radix-select: one kernel, no intermediate buffers. + # Equivalent output to the 2-stage path (set of block ids, front-packed, + # -1 padded); ~2-16x faster for long context. See + # sglang/jit_kernel/minimax_decode_topk.py. + from sglang.jit_kernel.minimax_decode_topk import minimax_decode_topk + + minimax_decode_topk(score, seq_lens, block_size, topk, out=topk_idx) + else: + # 2-stage split-K Triton fallback. + # Choose NUM_TOPK_CHUNKS to add parallelism over the seqblock dim. + # Same constraints as flash_decode: must be deterministic from values + # constant within a cuda graph (BS, num_q_heads), pow2, capped so the + # merge kernel's BLOCK_SIZE_K = pow2(NUM_TOPK_CHUNKS * pow2(topk)) + # stays reasonable. With topk=32 and cap 16, merge sorts 512 items. + TOPK_TARGET_GRID = 64 + MAX_NUM_TOPK_CHUNKS = 16 + topk_target = max( + 1, + min( + MAX_NUM_TOPK_CHUNKS, + TOPK_TARGET_GRID // max(1, batch_size * num_q_heads), + ), + ) + NUM_TOPK_CHUNKS = 1 << (topk_target.bit_length() - 1) + BLOCK_SIZE_T = triton.next_power_of_2(topk) + max_seqblock = score.shape[2] + chunk_blocks = (max_seqblock + NUM_TOPK_CHUNKS - 1) // NUM_TOPK_CHUNKS + topk_score_partial = torch.empty( + NUM_TOPK_CHUNKS, + num_q_heads, + batch_size, + BLOCK_SIZE_T, + dtype=torch.float32, + device=score.device, + ) + topk_idx_partial = torch.empty( + NUM_TOPK_CHUNKS, + num_q_heads, + batch_size, + BLOCK_SIZE_T, + dtype=torch.int32, + device=score.device, + ) + # stream 0 (default): topk partial → topk merge + grid = (batch_size, num_q_heads, NUM_TOPK_CHUNKS) + _topk_index_partial_kernel[grid]( + score, + topk_score_partial, + topk_idx_partial, + seq_lens, + block_size, + topk, + chunk_blocks, + score.stride(0), + score.stride(1), + score.stride(2), + topk_score_partial.stride(0), + topk_score_partial.stride(1), + topk_score_partial.stride(2), + topk_score_partial.stride(3), + topk_idx_partial.stride(0), + topk_idx_partial.stride(1), + topk_idx_partial.stride(2), + topk_idx_partial.stride(3), + ) + grid = (batch_size, num_q_heads) + _topk_index_merge_kernel[grid]( + topk_score_partial, + topk_idx_partial, + topk_idx, + seq_lens, + block_size, + topk, + topk_score_partial.stride(0), + topk_score_partial.stride(1), + topk_score_partial.stride(2), + topk_score_partial.stride(3), + topk_idx_partial.stride(0), + topk_idx_partial.stride(1), + topk_idx_partial.stride(2), + topk_idx_partial.stride(3), + topk_idx.stride(0), + topk_idx.stride(1), + topk_idx.stride(2), + NUM_TOPK_CHUNKS=NUM_TOPK_CHUNKS, + ) + if disable_index_value: + return None, topk_idx, real_seq_lens + # attn output merge (default stream) + grid = (batch_size, num_q_heads) + _merge_attn_out_kernel[grid]( + o, + lse, + seq_lens, + head_dim, + block_size, + o.stride(0), + o.stride(1), + o.stride(2), + o.stride(3), + lse.stride(0), + lse.stride(1), + lse.stride(2), + NUM_KV_CHUNKS=NUM_KV_CHUNKS, + ) + o = o[0].contiguous() + return o, topk_idx, real_seq_lens diff --git a/python/sglang/srt/layers/attention/minimax_sparse_ops/decode/topk_sparse.py b/python/sglang/srt/layers/attention/minimax_sparse_ops/decode/topk_sparse.py new file mode 100644 index 000000000..e2d902604 --- /dev/null +++ b/python/sglang/srt/layers/attention/minimax_sparse_ops/decode/topk_sparse.py @@ -0,0 +1,419 @@ +# Copyright 2025 XunhaoLai. All rights reserved. + +from typing import Optional + +import torch +import triton +import triton.language as tl + +from ..common.utils import check_sparse_kv_fp8, robust_allocator + + +@triton.heuristics( + { + "BLOCK_SIZE_H": lambda args: max( + 16, triton.next_power_of_2(args["gqa_group_size"]) + ), + "BLOCK_SIZE_D": lambda args: triton.next_power_of_2(args["head_dim"]), + "BLOCK_SIZE_T": lambda args: triton.next_power_of_2(args["max_topk"]), + "HAS_SINK": lambda args: args["sink_ptr"] is not None, + "BATCH_SIZE_BUCKET": lambda args: triton.next_power_of_2(args["batch_size"]), + } +) +@triton.autotune( + configs=[ + triton.Config({}, num_warps=nw, num_stages=ns) + for nw in [4, 8] + for ns in [2, 3, 4, 5] + ], + key=["BATCH_SIZE_BUCKET", "gqa_group_size", "head_dim", "block_size", "HAS_SINK"], +) +@triton.jit +def _gqa_share_sparse_decode_kernel( + q_ptr, # Q: b x qh x d + sink_ptr, # Sink: qh x d + k_cache_ptr, # K paged: max_slots x kh x d + v_cache_ptr, # V paged: max_slots x kh x d + req_to_token_ptr, # req_to_token: max_reqs x max_kv_len + idx_ptr, # topk index: qh x b x topk + o_ptr, # O partial: c x b x qh x d + lse_ptr, # lse partial: c x b x qh + seq_lens, + slot_ids, + # shape + max_slots, + batch_size, + gqa_group_size, + head_dim, + max_topk, + max_kv_len, + # sm_scale + sm_scale, + # stride + stride_q_b, + stride_q_h, + stride_q_d, + stride_sink_h, + stride_sink_d, + stride_k_s, + stride_k_h, + stride_k_d, + stride_v_s, + stride_v_h, + stride_v_d, + stride_r2t_b, + stride_ti_h, + stride_ti_b, + stride_ti_t, + stride_o_c, + stride_o_b, + stride_o_h, + stride_o_d, + stride_l_c, + stride_l_b, + stride_l_h, + # META parameters + BATCH_SIZE_BUCKET: tl.constexpr, + BLOCK_SIZE_H: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, + BLOCK_SIZE_D: tl.constexpr, + BLOCK_SIZE_T: tl.constexpr, + NUM_TOPK_CHUNKS: tl.constexpr, + HAS_SINK: tl.constexpr, + IS_FP8: tl.constexpr, +): + # decode program ids: split-K over the topk dimension to give every SM + # something to do at small batch. pid(0) folds (batch, chunk) together so + # the grid size = batch_size * NUM_TOPK_CHUNKS. + pid_bc, pid_kh = tl.program_id(0), tl.program_id(1) + pid_b = pid_bc % batch_size + pid_c = pid_bc // batch_size + pid_h = pid_kh * gqa_group_size + # per-chunk topk range. chunk_size is *runtime* (depends on max_topk which + # is a runtime arg, not constexpr), so don't annotate as tl.constexpr — + # doing so produces undefined behavior in Triton. + chunk_size_topk = (max_topk + NUM_TOPK_CHUNKS - 1) // NUM_TOPK_CHUNKS + chunk_start_topk = pid_c * chunk_size_topk + chunk_end_topk_compiletime = chunk_start_topk + chunk_size_topk + # get q k start and len after rmpad + seq_len = tl.minimum(tl.load(seq_lens + pid_b), max_kv_len) + sid = ( + tl.load(slot_ids + pid_b).to(tl.int64) + max_slots + ) % max_slots # to avoid bugs when slot_ids is negative + # get real topk + off_t = tl.arange(0, BLOCK_SIZE_T) + idx_base = idx_ptr + pid_kh * stride_ti_h + pid_b * stride_ti_b + topk_idx = tl.load(idx_base + off_t * stride_ti_t, mask=off_t < max_topk, other=-1) + valid_idx = tl.where(topk_idx >= 0, off_t, -1) + real_topk = tl.sum(valid_idx != -1, axis=0) + chunk_end_topk = tl.minimum(chunk_end_topk_compiletime, real_topk) + # init pointer + off_n = tl.arange(0, BLOCK_SIZE_N) + off_d = tl.arange(0, BLOCK_SIZE_D) + dim_mask = off_d < head_dim + # init statistics — kept at -inf so empty chunks (chunk_start >= real_topk) + # naturally fall out as weight=0 in the merge step. + if HAS_SINK and pid_c == 0: + q_ptrs = tl.make_block_ptr( + base=q_ptr + pid_b * stride_q_b + pid_h * stride_q_h, + shape=(gqa_group_size, head_dim), + strides=(stride_q_h, stride_q_d), + offsets=(0, 0), + block_shape=(BLOCK_SIZE_H, BLOCK_SIZE_D), + order=(1, 0), + ) + q = tl.load(q_ptrs, boundary_check=(0, 1), padding_option="zero") + sink_ptrs = tl.make_block_ptr( + base=sink_ptr + pid_h * stride_sink_h, + shape=(gqa_group_size, head_dim), + strides=(stride_sink_h, stride_sink_d), + offsets=(0, 0), + block_shape=(BLOCK_SIZE_H, BLOCK_SIZE_D), + order=(1, 0), + ) + sink = tl.load(sink_ptrs, boundary_check=(0, 1), padding_option="zero").to( + tl.float32 + ) + qsink = tl.sum(q.to(tl.float32) * sink, axis=1) * sm_scale # (BLOCK_SIZE_H,) + m_i = qsink + lse_i = qsink + else: + m_i = tl.full((BLOCK_SIZE_H,), float("-inf"), dtype=tl.float32) + lse_i = tl.full((BLOCK_SIZE_H,), float("-inf"), dtype=tl.float32) + q_ptrs = tl.make_block_ptr( + base=q_ptr + pid_b * stride_q_b + pid_h * stride_q_h, + shape=(gqa_group_size, head_dim), + strides=(stride_q_h, stride_q_d), + offsets=(0, 0), + block_shape=(BLOCK_SIZE_H, BLOCK_SIZE_D), + order=(1, 0), + ) + q = tl.load(q_ptrs, boundary_check=(0, 1), padding_option="zero") + acc_o = tl.full((BLOCK_SIZE_H, BLOCK_SIZE_D), 0, dtype=tl.float32) + # only iterate over this chunk's topk slice. the load must respect the + # per-chunk start offset. + cur_idx_ptr = idx_base + chunk_start_topk * stride_ti_t + for _ in tl.range(chunk_start_topk, chunk_end_topk): + # load index + c = tl.load(cur_idx_ptr).to(tl.int32) * BLOCK_SIZE_N + cur_idx_ptr = cur_idx_ptr + stride_ti_t + # resolve slots for this block via req_to_token + pos = c + off_n + pos_mask = pos < seq_len + slots = tl.load( + req_to_token_ptr + sid * stride_r2t_b + pos, + mask=pos_mask, + other=0, + ).to(tl.int64) + slots = (slots + max_slots) % max_slots # safety against negative + # load K as (head_dim, BLOCK_SIZE_N) via indirect addressing + k_off = ( + slots[None, :] * stride_k_s + + pid_kh * stride_k_h + + off_d[:, None] * stride_k_d + ) + k = tl.load( + k_cache_ptr + k_off, + mask=dim_mask[:, None] & pos_mask[None, :], + other=0.0, + ) + if IS_FP8: + # fp8 KV cache is unit-scaled (set_kv_buffer casts bf16->fp8 with no + # scale), so dequant is just a widening cast to the Q compute dtype + # before the tl.dot. Matches the bf16 path bit-for-bit when the cache + # is bf16 (IS_FP8 False -> this branch is compiled out). + k = k.to(q.dtype) + # load V as (BLOCK_SIZE_N, head_dim) via indirect addressing + v_off = ( + slots[:, None] * stride_v_s + + pid_kh * stride_v_h + + off_d[None, :] * stride_v_d + ) + v = tl.load( + v_cache_ptr + v_off, + mask=pos_mask[:, None] & dim_mask[None, :], + other=0.0, + ) + if IS_FP8: + # Widen V before the P@V dot. This also makes the `p.to(v.dtype)` + # below cast P to the compute dtype (not to fp8, which would be + # catastrophic precision loss on the attention weights). + v = v.to(q.dtype) + # compute qk + qk = tl.zeros((BLOCK_SIZE_H, BLOCK_SIZE_N), dtype=tl.float32) + qk += tl.where(off_n[None, :] < seq_len - c, 0, float("-inf")) + # [H, D], [D, N] -> [H, N] + qk += tl.dot(q, k) * sm_scale + # compute m_ij and l_ij + m_ij = tl.maximum(m_i, tl.max(qk, axis=1)) + p = tl.exp(qk - m_ij[:, None]) + l_ij = tl.sum(p, axis=1) + # scale acc_o + acc_o_scale = tl.exp(m_i - m_ij) + acc_o = acc_o * acc_o_scale[:, None] + # load v and update acc_o + p = p.to(v.dtype) + # [H, N], [N, D] -> [H, D] + acc_o += tl.dot(p.to(v.dtype), v) + # update statistics + m_i = m_ij + lse_i = m_ij + tl.log(tl.exp(lse_i - m_ij) + l_ij) + # final scale (matches the old non-split kernel for chunks where lse_i>-inf). + # For empty chunks (chunk_start_topk >= real_topk) the inner loop never + # runs, so m_i = lse_i = -inf and naive `tl.exp(m_i - lse_i)` would compute + # exp(-inf - (-inf)) = exp(NaN) = NaN, then 0 * NaN = NaN poisons o_partial + # and the merge result. Gate the scale with tl.where so empty chunks emit a + # clean zero (lse_i stays -inf which the merge correctly turns into weight=0). + scale = tl.where( + lse_i > float("-inf"), + tl.exp(m_i - lse_i), + tl.zeros_like(lse_i), + ) + acc_o = acc_o * scale[:, None] + # save partial output and lse for the merge step + o_ptrs = tl.make_block_ptr( + base=o_ptr + pid_c * stride_o_c + pid_b * stride_o_b + pid_h * stride_o_h, + shape=(gqa_group_size, head_dim), + strides=(stride_o_h, stride_o_d), + offsets=(0, 0), + block_shape=(BLOCK_SIZE_H, BLOCK_SIZE_D), + order=(1, 0), + ) + tl.store(o_ptrs, acc_o.to(o_ptr.dtype.element_ty), boundary_check=(0, 1)) + lse_ptrs = tl.make_block_ptr( + base=lse_ptr + pid_c * stride_l_c + pid_b * stride_l_b + pid_h * stride_l_h, + shape=(gqa_group_size,), + strides=(stride_l_h,), + offsets=(0,), + block_shape=(BLOCK_SIZE_H,), + order=(0,), + ) + tl.store(lse_ptrs, lse_i.to(lse_ptr.dtype.element_ty), boundary_check=(0,)) + + +@triton.heuristics( + { + "BLOCK_SIZE_D": lambda args: triton.next_power_of_2(args["head_dim"]), + } +) +@triton.jit +def _merge_topk_attn_out_kernel( + o_ptr, # [NUM_TOPK_CHUNKS, BS, NQH, D] — partials in, merged out at chunk 0 + lse_ptr, # [NUM_TOPK_CHUNKS, BS, NQH] + head_dim, + stride_o_c, + stride_o_b, + stride_o_h, + stride_o_d, + stride_l_c, + stride_l_b, + stride_l_h, + NUM_TOPK_CHUNKS: tl.constexpr, + BLOCK_SIZE_D: tl.constexpr, +): + pid_b, pid_h = tl.program_id(0), tl.program_id(1) + off_c = tl.arange(0, NUM_TOPK_CHUNKS) + off_d = tl.arange(0, BLOCK_SIZE_D) + o_ptrs = tl.make_block_ptr( + base=o_ptr + pid_b * stride_o_b + pid_h * stride_o_h, + shape=(NUM_TOPK_CHUNKS, head_dim), + strides=(stride_o_c, stride_o_d), + offsets=(0, 0), + block_shape=(NUM_TOPK_CHUNKS, BLOCK_SIZE_D), + order=(1, 0), + ) + lse_ptrs = lse_ptr + pid_b * stride_l_b + pid_h * stride_l_h + off_c * stride_l_c + o = tl.load(o_ptrs, boundary_check=(0, 1), padding_option="zero") + lse = tl.load(lse_ptrs) # empty chunks contribute -inf -> weight 0 + # standard flash-decoding merge in linear (not log2) space, matching the + # decode kernel which uses tl.exp / tl.log. + lse_max = tl.max(lse, axis=0) + weights = tl.exp(lse - lse_max) + weights = weights / tl.sum(weights, axis=0) + o_merged = tl.sum(o * weights[:, None], axis=0) + o_out_ptrs = o_ptr + pid_b * stride_o_b + pid_h * stride_o_h + off_d * stride_o_d + tl.store(o_out_ptrs, o_merged.to(o_ptr.dtype.element_ty), mask=off_d < head_dim) + + +@torch.no_grad() +def flash_decode_with_gqa_share_sparse( + q: torch.Tensor, # [batch_size, num_q_heads, head_dim] + sink: Optional[torch.Tensor], + k_cache: torch.Tensor, # [max_slots, num_kv_heads, head_dim] (paged) + v_cache: torch.Tensor, # [max_slots, num_kv_heads, head_dim] (paged) + req_to_token: torch.Tensor, # [max_reqs, max_kv_len] + seq_lens: torch.Tensor, # [batch_size, ] + slot_ids: torch.Tensor, # [batch_size, ] + block_size: int, + topk_idx: torch.Tensor, # [num_kv_heads, batch_size, topk] + sm_scale: Optional[float] = None, + use_tma: bool = True, +) -> torch.Tensor: + triton.set_allocator(robust_allocator) + is_fp8 = check_sparse_kv_fp8(q, k_cache, v_cache, label="decode") + # shape + batch_size, num_q_heads, head_dim = q.shape + max_slots, num_kv_heads, _ = k_cache.shape + assert slot_ids.shape[0] == batch_size and seq_lens.shape[0] == batch_size + assert topk_idx.shape[0] == num_kv_heads + assert ( + triton.next_power_of_2(block_size) == block_size + ), f"block_size must be a power of 2, but got {block_size}" + # assert slot_ids.max() < max_slots, f"get slot_ids {slot_ids}, but kv_cache shape is {kv_cache.shape}" + max_kv_len = req_to_token.shape[1] + # gqa + assert num_q_heads % num_kv_heads == 0 + gqa_group_size = num_q_heads // num_kv_heads + max_topk = topk_idx.shape[2] + # sm scale + if sm_scale is None: + sm_scale = head_dim**-0.5 + # Pick NUM_TOPK_CHUNKS so total grid ≈ TARGET_GRID. Same constraints as + # flash_decode_with_topk_idx: must be power of 2 (Triton arange) and must + # only depend on shape constants (so grid is fixed within a cuda graph). + # Capped by max_topk because chunks beyond real_topk early-fall-through to + # the merge-as-zero path; capping avoids wasting blocks at tiny topk. + TARGET_GRID = 256 + target = max( + 1, + min(max_topk, TARGET_GRID // max(1, batch_size * num_kv_heads)), + ) + NUM_TOPK_CHUNKS = 1 << (target.bit_length() - 1) + # output tensor: split-K partials, merged into chunk 0 by the merge kernel + o_partial = torch.empty( + NUM_TOPK_CHUNKS, + batch_size, + num_q_heads, + head_dim, + dtype=q.dtype, + device=q.device, + ) + lse_partial = torch.empty( + NUM_TOPK_CHUNKS, + batch_size, + num_q_heads, + dtype=torch.float32, + device=q.device, + ) + # launch attention kernel + grid = (batch_size * NUM_TOPK_CHUNKS, num_kv_heads) + _gqa_share_sparse_decode_kernel[grid]( + q, + sink, + k_cache, + v_cache, + req_to_token, + topk_idx, + o_partial, + lse_partial, + seq_lens, + slot_ids, + max_slots, + batch_size, + gqa_group_size, + head_dim, + max_topk, + max_kv_len, + sm_scale, + q.stride(0), + q.stride(1), + q.stride(2), + sink.stride(0) if sink is not None else 0, + sink.stride(1) if sink is not None else 0, + k_cache.stride(0), + k_cache.stride(1), + k_cache.stride(2), + v_cache.stride(0), + v_cache.stride(1), + v_cache.stride(2), + req_to_token.stride(0), + topk_idx.stride(0), + topk_idx.stride(1), + topk_idx.stride(2), + o_partial.stride(0), + o_partial.stride(1), + o_partial.stride(2), + o_partial.stride(3), + lse_partial.stride(0), + lse_partial.stride(1), + lse_partial.stride(2), + BLOCK_SIZE_N=block_size, + NUM_TOPK_CHUNKS=NUM_TOPK_CHUNKS, + IS_FP8=is_fp8, + ) + # merge partials into chunk 0 + merge_grid = (batch_size, num_q_heads) + _merge_topk_attn_out_kernel[merge_grid]( + o_partial, + lse_partial, + head_dim, + o_partial.stride(0), + o_partial.stride(1), + o_partial.stride(2), + o_partial.stride(3), + lse_partial.stride(0), + lse_partial.stride(1), + lse_partial.stride(2), + NUM_TOPK_CHUNKS=NUM_TOPK_CHUNKS, + ) + return o_partial[0].contiguous() diff --git a/python/sglang/srt/layers/attention/minimax_sparse_ops/minimax_sparse.py b/python/sglang/srt/layers/attention/minimax_sparse_ops/minimax_sparse.py new file mode 100644 index 000000000..a4ee9fcf2 --- /dev/null +++ b/python/sglang/srt/layers/attention/minimax_sparse_ops/minimax_sparse.py @@ -0,0 +1,234 @@ +# Copyright 2025 XunhaoLai. All rights reserved. + +from typing import Callable, List, Optional, Tuple + +import torch + +from .common.index import topk_index_reduce +from .common.utils import get_cu_seqblocks +from .decode.flash_with_topk_idx import flash_decode_with_topk_idx +from .decode.topk_sparse import flash_decode_with_gqa_share_sparse +from .prefill.flash_with_topk_idx import flash_prefill_with_topk_index +from .prefill.topk_sparse import flash_prefill_with_gqa_share_sparse + + +def minimax_sparse_prefill( + q: torch.Tensor, # [total_extend_tokens, num_q_heads, qk_head_dim] + k_cache: torch.Tensor, # [max_slots, num_kv_heads, head_dim] (paged main) + v_cache: torch.Tensor, # [max_slots, num_kv_heads, head_dim] (paged main) + sink: Optional[torch.Tensor], # [num_q_heads, qk_head_dim] + idx_q: torch.Tensor, # [total_extend_tokens, num_idx_heads, idx_head_dim] + idx_k_cache: torch.Tensor, # [max_slots, 1, idx_head_dim] (paged index) + idx_v_cache: Optional[ + torch.Tensor + ], # [max_slots, 1, idx_head_dim] (paged index); None when disable_index_value + idx_sink: Optional[torch.Tensor], # [num_idx_heads, idx_head_dim] + req_to_token: torch.Tensor, # [max_reqs, max_kv_len] + slot_ids: torch.Tensor, # [batch_size, ] + cu_seqlens: torch.Tensor, # [batch_size + 1, ] (Q-side cumulative) + seq_lens: torch.Tensor, # [batch_size, ] total K length (prefix + chunk) + prefix_lens: torch.Tensor, # [batch_size, ] + max_seqlen_q: int, + max_seqlen_k: int, + block_size_q: int, + block_size_k: int, + topk: int, + init_blocks: int, + local_blocks: int, + sm_scale: Optional[float] = None, + idx_sm_scale: Optional[float] = None, + score_type: str = "max", + disable_index_value: bool = False, + use_msa: bool = False, + cu_seqblocks_q: Optional[torch.Tensor] = None, + max_seqblock_q: Optional[int] = None, + all_seqblock_q: Optional[int] = None, + seqlens_cpu: Optional[List[int]] = None, +): + """Run MiniMax-M3 sparse prefill. + + ``cu_seqblocks_q``, ``max_seqblock_q``, and ``all_seqblock_q`` are optional + precomputed query-block metadata shared by the index and value sparse + kernels. Supplying them avoids recomputing the same block layout twice. + ``seqlens_cpu`` (host copy of ``torch.diff(cu_seqlens)``) is forwarded to + ``get_cu_seqblocks`` to avoid a per-layer device sync when it recomputes. + """ + if cu_seqblocks_q is None or max_seqblock_q is None or all_seqblock_q is None: + cu_seqblocks_q, max_seqblock_q, all_seqblock_q, _, _, _ = get_cu_seqblocks( + cu_seqlens, max_seqlen_q, block_size_q, block_size_k, seqlens_cpu + ) + + # All seqlen is less than topk, use full attention + # Step 1: Flash attention with topk index (using index head) + idx_o, topk_idx = flash_prefill_with_topk_index( + q=idx_q, + k_cache=idx_k_cache, + v_cache=idx_v_cache, + sink=idx_sink, + req_to_token=req_to_token, + slot_ids=slot_ids, + cu_seqlens=cu_seqlens, + seq_lens=seq_lens, + prefix_lens=prefix_lens, + max_seqlen_q=max_seqlen_q, + max_seqlen_k=max_seqlen_k, + block_size_q=block_size_q, + block_size_k=block_size_k, + topk=topk, + init_blocks=init_blocks, + local_blocks=local_blocks, + sm_scale=idx_sm_scale, + score_type=score_type, + disable_index_value=disable_index_value, + cu_seqblocks_q=cu_seqblocks_q, + max_seqblock_q=max_seqblock_q, + all_seqblock_q=all_seqblock_q, + ) + # Step 2: Reduce topk idx if num_idx_heads > num_kv_heads + num_idx_heads = idx_q.shape[1] + num_kv_heads = k_cache.shape[1] + idx_group_size = num_idx_heads // num_kv_heads + if idx_group_size > 1: + topk_idx = topk_index_reduce( + topk_idx.view(num_kv_heads, idx_group_size, -1, topk), dim=1 + ) + # Step 3: Sparse attention using topk index (main head). The MSA path only + # replaces this step; the indexer above is unchanged. MSA has no attn-sink + # input, so keep the Triton path when sink is present. + if use_msa and sink is None: + from .msa import msa_sparse_prefill_main + + o = msa_sparse_prefill_main( + q=q, + k_cache=k_cache, + v_cache=v_cache, + topk_idx=topk_idx, + req_to_token=req_to_token, + slot_ids=slot_ids, + cu_seqlens=cu_seqlens, + seq_lens=seq_lens, + prefix_lens=prefix_lens, + block_size_k=block_size_k, + sm_scale=sm_scale, + ) + else: + o = flash_prefill_with_gqa_share_sparse( + q=q, + k_cache=k_cache, + v_cache=v_cache, + sink=sink, + req_to_token=req_to_token, + slot_ids=slot_ids, + topk_idx=topk_idx, + block_size_q=block_size_q, + block_size_k=block_size_k, + cu_seqlens=cu_seqlens, + seq_lens=seq_lens, + prefix_lens=prefix_lens, + max_seqlen_q=max_seqlen_q, + sm_scale=sm_scale, + cu_seqblocks_q=cu_seqblocks_q, + max_seqblock_q=max_seqblock_q, + ) + return idx_o, o + + +def minimax_sparse_decode( + q: torch.Tensor, # [batch_size, num_q_heads, qk_head_dim] + sink: Optional[torch.Tensor], # [num_q_heads, qk_head_dim] + k_cache: torch.Tensor, # [max_slots, num_kv_heads, head_dim] (paged) + v_cache: torch.Tensor, # [max_slots, num_kv_heads, head_dim] (paged) + idx_q: torch.Tensor, # [batch_size, num_idx_heads, idx_head_dim], num_idx_heads >= num_kv_heads + idx_sink: Optional[torch.Tensor], # [num_idx_heads, idx_head_dim] + idx_k_cache: torch.Tensor, # [max_slots, 1, idx_head_dim] (paged) + idx_v_cache: Optional[ + torch.Tensor + ], # [max_slots, 1, idx_head_dim] (paged); None when disable_index_value + req_to_token: torch.Tensor, # [max_reqs, max_kv_len] + slot_ids: torch.Tensor, # [batch_size, ] + seq_lens: torch.Tensor, # [batch_size, ] + max_seqlen: int, # max of seq_lens, passed from caller to avoid sync during CUDA graph capture + block_size_q: int, # useless for now, will always be 1 + block_size_k: int, + topk: int, + init_blocks: int, + local_blocks: int, + sm_scale: Optional[float] = None, + idx_sm_scale: Optional[float] = None, + score_type: str = "max", + disable_index_value: bool = False, + dense_main_attn_fn: Optional[Callable] = None, + page_size: int = 1, + use_msa: bool = False, + msa_kv_indices: Optional[ + torch.Tensor + ] = None, # per-forward MSA page table (cached) + msa_plan=None, # per-forward MSA fmha_sm100 plan (cached) +) -> Tuple[torch.Tensor, torch.Tensor]: + # Step 1: Flash decode with topk index (using index head). When the dense main + # attention is used, the indexer emits the page table directly (fused + # transform) instead of block ids, plus the per-query effective KV length. + idx_o, topk_idx, real_seq_lens = flash_decode_with_topk_idx( + q=idx_q, + sink=idx_sink, + k_cache=idx_k_cache, + v_cache=idx_v_cache, + req_to_token=req_to_token, + seq_lens=seq_lens, + max_seqlen=max_seqlen, + slot_ids=slot_ids, + block_size=block_size_k, + topk=topk, + init_blocks=init_blocks, + local_blocks=local_blocks, + sm_scale=idx_sm_scale, + score_type=score_type, + disable_index_value=disable_index_value, + use_dense_main_attn=dense_main_attn_fn is not None, + page_size=page_size, + ) + num_idx_heads = idx_q.shape[1] + num_kv_heads = k_cache.shape[1] + idx_group_size = num_idx_heads // num_kv_heads + if dense_main_attn_fn is not None: + # topk_idx is the page table; real_seq_lens is the per-query cache_seqlens + assert idx_group_size == 1 + o = dense_main_attn_fn(q, topk_idx, real_seq_lens) + else: + # Step 2: Reduce topk idx if num_idx_heads > num_kv_heads + if idx_group_size > 1: + topk_idx = topk_index_reduce( + topk_idx.view(num_kv_heads, idx_group_size, -1, topk), dim=1 + ) + # Step 3: Sparse attention using topk index (main head). The MSA path + # only replaces this step; keep the Triton path when sink is present. + if use_msa and sink is None: + from .msa import msa_sparse_decode_main + + o = msa_sparse_decode_main( + q=q, + k_cache=k_cache, + v_cache=v_cache, + topk_idx=topk_idx, + req_to_token=req_to_token, + slot_ids=slot_ids, + seq_lens=seq_lens, + block_size_k=block_size_k, + sm_scale=sm_scale, + kv_indices=msa_kv_indices, + plan=msa_plan, + ) + else: + o = flash_decode_with_gqa_share_sparse( + q=q, + sink=sink, + k_cache=k_cache, + v_cache=v_cache, + req_to_token=req_to_token, + seq_lens=seq_lens, + slot_ids=slot_ids, + block_size=block_size_k, + topk_idx=topk_idx, + sm_scale=sm_scale, + ) + return idx_o, o diff --git a/python/sglang/srt/layers/attention/minimax_sparse_ops/msa.py b/python/sglang/srt/layers/attention/minimax_sparse_ops/msa.py new file mode 100644 index 000000000..b79fc3227 --- /dev/null +++ b/python/sglang/srt/layers/attention/minimax_sparse_ops/msa.py @@ -0,0 +1,353 @@ +# MSA (fmha_sm100) drop-in for the MiniMax-M3 main sparse-attention step. +# +# Replaces only step 3 of MiniMax sparse prefill/decode. The lightning indexer +# (steps 1-2) is unchanged and still produces `topk_idx`. +# NVIDIA Blackwell (SM100/sm_103) only; callers gate on `msa_available()`. + +from __future__ import annotations + +import functools +from typing import Optional + +import torch + + +@functools.lru_cache(maxsize=1) +def msa_available() -> bool: + """True iff the fmha_sm100 sparse kernels are importable on this device.""" + try: + cap = torch.cuda.get_device_capability() + except Exception: + return False + # SM100 family: B200 (10,0) and B300 (10,3). fmha_sm100/jit.py emits both + # sm_100a and sm_103a; the kernels run on either. + if cap[0] != 10 or cap[1] not in (0, 3): + return False + try: + import fmha_sm100 # noqa: F401 + + return True + except Exception: + return False + + +def _build_page_table( + req_to_token: torch.Tensor, # [max_reqs, max_kv_len], physical slot per logical pos + slot_ids: torch.Tensor, # [batch] + seq_lens: torch.Tensor, # [batch] total K length (prefix + chunk) + page_size: int, +) -> torch.Tensor: + """Flattened physical page ids per request (MSA `kv_indices`). + + sglang's paged allocator stores page_size contiguous physical slots per page, so the + physical page of logical position p is ``req_to_token[req, p] // page_size`` and is the + same for every p within a page. We read one slot per logical page to recover the table. + + Vectorized (no per-request Python loop): pages are packed contiguously by request in the + same order MSA's planner expects (``kv_page_indptr = cumsum(ceil(seq_lens/page_size))``). + ``searchsorted`` maps each packed page slot back to its request; one ``.item()`` recovers + the total page count (this runs eagerly, outside CUDA-graph capture). + """ + P = page_size + n_pages = (seq_lens.to(torch.int64) + (P - 1)) // P # [batch] + offsets = ( + torch.cumsum(n_pages, 0) - n_pages + ) # [batch] exclusive page offset per request + total = int(n_pages.sum().item()) + idx = torch.arange( + total, device=req_to_token.device + ) # packed page slot -> (req, page) + req = torch.searchsorted( + offsets + n_pages, idx, right=True + ) # request id per packed slot + logical_first = (idx - offsets[req]) * P # first logical position of that page + rows = slot_ids[req].to(torch.int64) + return (req_to_token[rows, logical_first] // P).to(torch.int32) + + +def msa_sparse_prefill_main( + q: torch.Tensor, # [total_q, num_q_heads, head_dim] + k_cache: torch.Tensor, # [max_slots, num_kv_heads, head_dim] (slot-major NHD) + v_cache: torch.Tensor, # [max_slots, num_kv_heads, head_dim] + topk_idx: torch.Tensor, # [num_kv_heads, total_q, topk] (0-based, -1 pad) -- step1/2 output + req_to_token: torch.Tensor, # [max_reqs, max_kv_len] + slot_ids: torch.Tensor, # [batch] + cu_seqlens: torch.Tensor, # [batch+1] cumulative Q lengths + seq_lens: torch.Tensor, # [batch] total K length (prefix + chunk) + prefix_lens: torch.Tensor, # [batch] + block_size_k: int, # == page_size == 128 for M3 + sm_scale: Optional[float] = None, +) -> torch.Tensor: + """Drop-in for flash_prefill_with_gqa_share_sparse using MSA fmha_sm100. + + Returns o [total_q, num_q_heads, head_dim]. + """ + from fmha_sm100 import fmha_sm100, fmha_sm100_plan + + max_slots, num_kv_heads, head_dim = k_cache.shape + num_q_heads = q.shape[1] + P = block_size_k + topk = topk_idx.shape[-1] + if max_slots % P != 0: + raise ValueError(f"max_slots={max_slots} not divisible by page_size={P}") + if sm_scale is None: + sm_scale = head_dim**-0.5 + + # Whole pool as MSA paged KV: [num_phys_pages, num_kv_heads, P, head_dim]. + n_phys_pages = max_slots // P + k_paged = k_cache.view(n_phys_pages, P, num_kv_heads, head_dim).permute(0, 2, 1, 3) + v_paged = v_cache.view(n_phys_pages, P, num_kv_heads, head_dim).permute(0, 2, 1, 3) + + # Per-request Q lengths (extend) and physical page table. + qo_segment_lens = (cu_seqlens[1:] - cu_seqlens[:-1]).to(torch.int32) + kv_indices = _build_page_table(req_to_token, slot_ids, seq_lens, P) + + # topk_idx [Hkv, total_q, topk] -> kv_block_indexes [total_q, Hkv, topk]. + kv_block_indexes = topk_idx.permute(1, 0, 2).contiguous().to(torch.int32) + + plan = fmha_sm100_plan( + qo_segment_lens, + seq_lens.to(torch.int32), + num_q_heads, + num_kv_heads=num_kv_heads, + page_size=P, + kv_block_num=topk, + causal=True, + qo_offset=prefix_lens.to(torch.int32), + ) + o, _ = fmha_sm100( + q, + k_paged, + v_paged, + plan, + sm_scale=sm_scale, + kv_indices=kv_indices, + kv_block_indexes=kv_block_indexes, + ) + return o + + +def build_msa_decode_meta( + k_cache: torch.Tensor, # [max_slots, num_kv_heads, head_dim] + req_to_token: torch.Tensor, + slot_ids: torch.Tensor, # [batch] + seq_lens: torch.Tensor, # [batch] cached KV length per request + num_q_heads: int, + block_size_k: int, + topk: int, +): + """Per-forward MSA decode metadata (page table + fmha plan), shared across layers. + + Within one decode forward every sparse layer has the same batch, seq_lens, page size + and topk, so the physical page table and the fmha_sm100 plan are identical for all of + them — only ``kv_block_indexes`` (the per-layer top-k selection) changes. Building these + once per forward instead of once per layer removes the dominant host-side overhead + (page-table build + the host-side ``fmha_sm100_plan``) from the 57-layer decode loop. + + Used only by the standalone parity harnesses; the serving backend builds eager-decode + metadata via ``build_msa_decode_cg_plan`` + ``update_msa_decode_cg_meta`` instead. + """ + from fmha_sm100 import fmha_sm100_plan + + max_slots, num_kv_heads, _ = k_cache.shape + P = block_size_k + if max_slots % P != 0: + raise ValueError(f"max_slots={max_slots} not divisible by page_size={P}") + B = slot_ids.shape[0] + kv_indices = _build_page_table(req_to_token, slot_ids, seq_lens, P) + seq_lens_i32 = seq_lens.to(torch.int32) + plan = fmha_sm100_plan( + torch.ones(B, dtype=torch.int32), + seq_lens_i32, + num_q_heads, + num_kv_heads=num_kv_heads, + page_size=P, + kv_block_num=topk, + causal=False, + qo_offset=seq_lens_i32 - 1, # decode query sits at the last cached position + ) + return kv_indices, plan + + +# --------------------------------------------------------------------------- +# Eager-only MSA decode plan (NOT used under CUDA graph) +# +# WARNING: the fmha_sm100 sparse decode kernel is NOT cuda-graph-safe — captured +# and replayed it returns silently wrong results that compound across replays +# (~14% GSM8K loss on B200). The backend routes decode to the cuda-graph-safe +# Triton sparse path whenever decode runs under a CUDA graph (see +# MiniMaxSparseAttnBackend._use_msa_decode); this plan is reachable ONLY in eager +# decode (no decode CUDA graph), where there is no capture/replay. Do NOT wire it +# back into a captured graph — that reintroduces the ~14% regression. +# +# The build-once / replay-update structure below (refreshing the four length +# tensors ``{kv_segment_lens, kv_segment_offsets, kv_page_indptr, qo_offset}`` and +# the page table in place) is a leftover from the abandoned capture-once attempt; +# it is kept only because eager decode reuses one per-forward plan across layers. +# --------------------------------------------------------------------------- + +_MSA_CG_LEN_KEYS = ( + "kv_segment_lens", + "kv_segment_offsets", + "kv_page_indptr", + "qo_offset", +) + + +def _check_cg_plan_layout(plan) -> None: + """Fail fast if fmha_sm100's plan layout drifted from what replay-update + assumes (dict at tuple index 3 holding the four length tensors) — these are + undocumented fmha_sm100 internals.""" + if not (isinstance(plan, tuple) and len(plan) > 3 and isinstance(plan[3], dict)): + raise RuntimeError( + "fmha_sm100_plan no longer returns a tuple with a metadata dict at index 3; " + "the MSA CUDA-graph decode path must be revalidated against this fmha_sm100 " + "version. Set SGLANG_DISABLE_MSA=1 to use the Triton path meanwhile." + ) + missing = [k for k in _MSA_CG_LEN_KEYS if not torch.is_tensor(plan[3].get(k))] + if missing: + raise RuntimeError( + f"fmha_sm100 plan is missing length tensors {missing}; the MSA CUDA-graph " + "decode path must be revalidated against this fmha_sm100 version. " + "Set SGLANG_DISABLE_MSA=1 to use the Triton path meanwhile." + ) + + +def build_msa_decode_cg_plan( + num_q_heads: int, + num_kv_heads: int, + block_size_k: int, + topk: int, + batch_size: int, + device: Optional[torch.device] = None, +): + """Persistent fmha_sm100 decode plan for one batch size (CUDA-graph stable). + + Built once per captured batch size. The worklist is length-independent (it uses + ``topk * page_size`` internally), so the reference KV length here only has to make + every topk block valid; the length-dependent tensors are overwritten each step by + ``update_msa_decode_cg_meta``. Returns the plan tuple to pass to ``fmha_sm100``. + """ + from fmha_sm100 import fmha_sm100_plan + + P = block_size_k + ref_len = topk * P # length at which all topk blocks exist -> full worklist + qo = torch.ones(batch_size, dtype=torch.int32) + kv = torch.full((batch_size,), ref_len, dtype=torch.int32) + plan = fmha_sm100_plan( + qo, + kv, + num_q_heads, + num_kv_heads=num_kv_heads, + page_size=P, + kv_block_num=topk, + causal=False, + qo_offset=kv - 1, + device=device, + ) + _check_cg_plan_layout(plan) + return plan + + +def update_msa_decode_cg_meta( + plan, + kv_indices_buf: torch.Tensor, # persistent page-table buffer [batch * max_pages] + req_to_token: torch.Tensor, + slot_ids: torch.Tensor, # [batch] + seq_lens: torch.Tensor, # [batch] cached KV length per request + block_size_k: int, + topk: int, + num_q_heads: int, + num_kv_heads: int, +): + """Refresh the persistent decode plan's length-dependent tensors + page table IN PLACE. + + Host-side (calls fmha_sm100_plan and one ``.item()``); MUST run outside CUDA-graph + capture — i.e. only from ``init_forward_metadata_out_graph``. The captured graph then + reads the same plan-tensor and ``kv_indices_buf`` addresses on replay. + """ + from fmha_sm100 import fmha_sm100_plan + + P = block_size_k + B = seq_lens.shape[0] + seq_lens_i32 = seq_lens.to(torch.int32) + # Fresh plan for the real lengths; copy only its four length-dependent tensors into the + # persistent plan (same shapes — they depend on batch size, not length). The fresh + # worklist is identical to the persistent one (topk*P based) and is discarded. + # qo_offset is clamped: graph replay pads the batch with seq_len==0 slots + # (masked via kv_segment_lens==0, but seq_len-1 would be -1). + fresh = fmha_sm100_plan( + torch.ones(B, dtype=torch.int32), + seq_lens_i32, + num_q_heads, + num_kv_heads=num_kv_heads, + page_size=P, + kv_block_num=topk, + causal=False, + qo_offset=(seq_lens_i32 - 1).clamp_min(0), + device=seq_lens.device, + ) + _check_cg_plan_layout(fresh) + pd, fd = plan[3], fresh[3] + for k in _MSA_CG_LEN_KEYS: + pd[k].copy_(fd[k]) + table = _build_page_table(req_to_token, slot_ids, seq_lens, P) + kv_indices_buf[: table.numel()].copy_(table) + + +def msa_sparse_decode_main( + q: torch.Tensor, # [batch, num_q_heads, head_dim] (1 query token per request) + k_cache: torch.Tensor, # [max_slots, num_kv_heads, head_dim] (slot-major NHD) + v_cache: torch.Tensor, # [max_slots, num_kv_heads, head_dim] + topk_idx: torch.Tensor, # [num_kv_heads, batch, topk] (0-based, -1 pad) + req_to_token: torch.Tensor, # [max_reqs, max_kv_len] + slot_ids: torch.Tensor, # [batch] + seq_lens: torch.Tensor, # [batch] cached KV length per request + block_size_k: int, # == page_size == 128 + sm_scale: Optional[float] = None, + kv_indices: Optional[ + torch.Tensor + ] = None, # precomputed page table (per-forward cache) + plan=None, # precomputed fmha_sm100 plan (per-forward cache) +) -> torch.Tensor: + """Drop-in for flash_decode_with_gqa_share_sparse using MSA fmha_sm100. + + Each request is one decode query at absolute position seq_len-1 attending to its + cached KV through the topk selected 128-blocks. Returns o [batch, num_q_heads, head_dim]. + + ``kv_indices`` / ``plan`` are shared across all sparse layers of a forward; the serving + backend builds them once via ``build_msa_decode_cg_plan`` + ``update_msa_decode_cg_meta`` + (eager decode only) and passes them in. When omitted (only the standalone parity + harnesses) they are built here via ``build_msa_decode_meta``. + """ + from fmha_sm100 import fmha_sm100 + + max_slots, num_kv_heads, head_dim = k_cache.shape + H = q.shape[1] + P = block_size_k + topk = topk_idx.shape[-1] + if max_slots % P != 0: + raise ValueError(f"max_slots={max_slots} not divisible by page_size={P}") + if sm_scale is None: + sm_scale = head_dim**-0.5 + + n_phys_pages = max_slots // P + k_paged = k_cache.view(n_phys_pages, P, num_kv_heads, head_dim).permute(0, 2, 1, 3) + v_paged = v_cache.view(n_phys_pages, P, num_kv_heads, head_dim).permute(0, 2, 1, 3) + + if kv_indices is None or plan is None: + kv_indices, plan = build_msa_decode_meta( + k_cache, req_to_token, slot_ids, seq_lens, H, P, topk + ) + kv_block_indexes = topk_idx.permute(1, 0, 2).contiguous().to(torch.int32) + + o, _ = fmha_sm100( + q, + k_paged, + v_paged, + plan, + sm_scale=sm_scale, + kv_indices=kv_indices, + kv_block_indexes=kv_block_indexes, + ) + return o diff --git a/python/sglang/srt/layers/attention/minimax_sparse_ops/naive/flash_with_topk_idx.py b/python/sglang/srt/layers/attention/minimax_sparse_ops/naive/flash_with_topk_idx.py new file mode 100644 index 000000000..379ffa31e --- /dev/null +++ b/python/sglang/srt/layers/attention/minimax_sparse_ops/naive/flash_with_topk_idx.py @@ -0,0 +1,84 @@ +# Copyright 2025 XunhaoLai. All rights reserved. + +from typing import Optional + +import torch +from einops import einsum, rearrange + + +def naive_flash_decode_with_topk_idx( + q: torch.Tensor, # [batch_size, num_heads, head_dim] + sink: Optional[torch.Tensor], # [num_heads, head_dim] + kv_cache: torch.Tensor, # [max_slots, 2, max_len, num_heads, head_dim] + seq_lens: torch.Tensor, # [batch_size, ] + max_seqlen: int, + slot_ids: torch.Tensor, # [batch_size, ] + block_size: int, + topk: int, + sm_scale: Optional[float] = None, + init_blocks: int = 0, + local_blocks: int = 0, +): + assert ( + kv_cache.shape[2] % block_size == 0 + ), "max cache len must be divisible by block size" + if sm_scale is None: + sm_scale = q.shape[-1] ** -0.5 + original_dtype = q.dtype + batch_size = q.shape[0] + num_q_heads = q.shape[1] + num_kv_heads = kv_cache.shape[3] + q = rearrange(q.float(), "b (h g) d -> b h g d", h=num_kv_heads) + kv_cache_float = kv_cache.float() + qk = ( + einsum(q, kv_cache_float[slot_ids, 0, ...], "b h g d, b n h d -> b h g n") + * sm_scale + ) + mask = torch.arange(kv_cache.shape[2], device=q.device) < seq_lens[:, None] + qk = qk.masked_fill(~mask[:, None, None, :], float("-inf")) + # get score + score = qk.clone().to(torch.float32) + score = rearrange(score, "b h g (n s) -> b (h g) n s", s=block_size) + score = score.max(dim=-1).values # [batch_size, num_q_heads, num_blocks] + # post-process score for init_blocks and local_blocks + INIT_SCORE = 1e30 + LOCAL_SCORE = 1e29 + if init_blocks > 0: + score[:, :, :init_blocks] = INIT_SCORE + if local_blocks > 0: + num_blocks_per_batch = (seq_lens + block_size - 1) // block_size + for b in range(batch_size): + num_blks = num_blocks_per_batch[b].item() + local_start = max(0, num_blks - local_blocks) + score[b, :, local_start:num_blks] = LOCAL_SCORE + # compute topk indices per (batch, head) + # score shape: [batch_size, num_q_heads, num_blocks] + topk_idx = torch.full( + (num_q_heads, batch_size, topk), + fill_value=-1, + device=score.device, + dtype=torch.int32, + ) + num_blocks_per_batch = (seq_lens + block_size - 1) // block_size + for b in range(batch_size): + num_blks = num_blocks_per_batch[b].item() + actual_topk = min(topk, num_blks) + for h in range(num_q_heads): + # get topk indices for this (batch, head) + _, indices = torch.topk(score[b, h, :num_blks], k=actual_topk, dim=-1) + topk_idx[h, b, :actual_topk] = indices.to(torch.int32) + # compute attention output with sink + if sink is not None: + # sink: [num_q_heads, head_dim] -> reshape to match q: [b, h, g, d] + sink_reshaped = rearrange(sink.float(), "(h g) d -> h g d", h=num_kv_heads) + qsink = ( + einsum(q, sink_reshaped, "b h g d, h g d -> b h g") * sm_scale + ) # [b, h, g] + qk_with_sink = torch.cat([qsink[..., None], qk], dim=-1) # [b, h, g, n+1] + attn = qk_with_sink.softmax(dim=-1, dtype=torch.float32) + attn = attn[..., 1:] # remove sink score + else: + attn = qk.softmax(dim=-1, dtype=torch.float32) + o = einsum(attn, kv_cache_float[slot_ids, 1, ...], "b h g n, b n h d -> b h g d") + o = rearrange(o, "b h g d -> b (h g) d", h=num_kv_heads) + return o.to(original_dtype), topk_idx diff --git a/python/sglang/srt/layers/attention/minimax_sparse_ops/naive/topk_sparse.py b/python/sglang/srt/layers/attention/minimax_sparse_ops/naive/topk_sparse.py new file mode 100644 index 000000000..6fefac38e --- /dev/null +++ b/python/sglang/srt/layers/attention/minimax_sparse_ops/naive/topk_sparse.py @@ -0,0 +1,88 @@ +# Copyright 2025 XunhaoLai. All rights reserved. + +from typing import Optional + +import torch + + +def naive_flash_decode_with_gqa_share_sparse( + q: torch.Tensor, # [batch_size, num_q_heads, head_dim] + sink: Optional[torch.Tensor], # [num_q_heads, head_dim] + kv_cache: torch.Tensor, # [max_slots, 2, max_len, num_kv_heads, head_dim] + seq_lens: torch.Tensor, # [batch_size, ] + slot_ids: torch.Tensor, # [batch_size, ] + block_size: int, + topk_idx: torch.Tensor, # [num_kv_heads, batch_size, topk] + sm_scale: Optional[float] = None, +) -> torch.Tensor: + """ + Naive implementation of sparse attention with GQA group sharing. + Each GQA group (all q heads sharing the same kv head) uses the same topk_idx. + """ + if sm_scale is None: + sm_scale = q.shape[-1] ** -0.5 + + original_dtype = q.dtype + batch_size, num_q_heads, head_dim = q.shape + max_slots, _, max_kv_len, num_kv_heads, _ = kv_cache.shape + gqa_group_size = num_q_heads // num_kv_heads + # Output tensor + o = torch.zeros(batch_size, num_q_heads, head_dim, dtype=q.dtype, device=q.device) + + for b in range(batch_size): + seq_len = seq_lens[b].item() + sid = slot_ids[b].item() % max_slots + + for kh in range(num_kv_heads): + # Get topk indices for this (kv_head, batch) + block_indices = topk_idx[kh, b, :].tolist() + + # Collect selected K and V blocks + selected_k_blocks = [] + selected_v_blocks = [] + for block_idx in block_indices: + if block_idx < 0: + continue # Invalid index + start = block_idx * block_size + end = min(start + block_size, seq_len) + if start >= seq_len: + continue + # K: [block_len, head_dim] + k_block = kv_cache[sid, 0, start:end, kh, :] + # V: [block_len, head_dim] + v_block = kv_cache[sid, 1, start:end, kh, :] + selected_k_blocks.append(k_block) + selected_v_blocks.append(v_block) + + if len(selected_k_blocks) == 0: + continue + + # Concatenate selected blocks: [total_selected_len, head_dim] + k_selected = torch.cat(selected_k_blocks, dim=0) + v_selected = torch.cat(selected_v_blocks, dim=0) + + # Compute attention for all q heads in this GQA group + for g in range(gqa_group_size): + qh = kh * gqa_group_size + g + # q_vec: [head_dim] + q_vec = q[b, qh, :].float() + + # Compute attention scores: [total_selected_len] + scores = torch.matmul(q_vec, k_selected.float().T) * sm_scale + + # Add sink to softmax normalization if present + if sink is not None: + sink_vec = sink[qh, :].float() + qsink = torch.dot(q_vec, sink_vec) * sm_scale + # Concatenate sink score with regular scores for softmax + scores_with_sink = torch.cat([qsink.unsqueeze(0), scores], dim=0) + attn_weights_with_sink = torch.softmax(scores_with_sink, dim=-1) + # Remove sink weight (only used for normalization) + attn_weights = attn_weights_with_sink[1:] + else: + attn_weights = torch.softmax(scores, dim=-1) + + # Compute output: [head_dim] + o[b, qh, :] = torch.matmul(attn_weights.to(original_dtype), v_selected) + + return o diff --git a/python/sglang/srt/layers/attention/minimax_sparse_ops/prefill/flash_with_topk_idx.py b/python/sglang/srt/layers/attention/minimax_sparse_ops/prefill/flash_with_topk_idx.py new file mode 100644 index 000000000..6be92ff67 --- /dev/null +++ b/python/sglang/srt/layers/attention/minimax_sparse_ops/prefill/flash_with_topk_idx.py @@ -0,0 +1,563 @@ +# Copyright 2025 XunhaoLai. All rights reserved. + +from typing import Optional + +import torch +import triton +import triton.language as tl + +from ..common.utils import _bitonic_merge, get_cu_seqblocks, robust_allocator + + +@triton.heuristics( + { + "BLOCK_SIZE_KD": lambda args: triton.next_power_of_2(args["qk_head_dim"]), + "BLOCK_SIZE_VD": lambda args: triton.next_power_of_2(args["v_head_dim"]), + "HAS_SINK": lambda args: args["sink_ptr"] is not None, + } +) +@triton.autotune( + configs=[ + # Small block (64x64): low shared mem, can use higher num_stages + triton.Config( + {"BLOCK_SIZE_Q": 64, "BLOCK_SIZE_K": 64}, num_warps=4, num_stages=2 + ), + triton.Config( + {"BLOCK_SIZE_Q": 64, "BLOCK_SIZE_K": 64}, num_warps=4, num_stages=3 + ), + triton.Config( + {"BLOCK_SIZE_Q": 64, "BLOCK_SIZE_K": 64}, num_warps=4, num_stages=4 + ), + # Medium block (64x128, 128x64): moderate shared mem, ns=2,3 + triton.Config( + {"BLOCK_SIZE_Q": 64, "BLOCK_SIZE_K": 128}, num_warps=8, num_stages=2 + ), + triton.Config( + {"BLOCK_SIZE_Q": 64, "BLOCK_SIZE_K": 128}, num_warps=8, num_stages=3 + ), + triton.Config( + {"BLOCK_SIZE_Q": 128, "BLOCK_SIZE_K": 64}, num_warps=8, num_stages=2 + ), + triton.Config( + {"BLOCK_SIZE_Q": 128, "BLOCK_SIZE_K": 64}, num_warps=8, num_stages=3 + ), + # Large block (128x128): high shared mem, ns=2,3 only, nw=8 + triton.Config( + {"BLOCK_SIZE_Q": 128, "BLOCK_SIZE_K": 128}, num_warps=8, num_stages=2 + ), + triton.Config( + {"BLOCK_SIZE_Q": 128, "BLOCK_SIZE_K": 128}, num_warps=8, num_stages=3 + ), + ], + key=[ + "qk_head_dim", + "v_head_dim", + "block_size", + "use_gumbel_topk", + "SCORE_TYPE", + "DISABLE_INDEX_VALUE", + ], +) +@triton.jit +def _flash_attn_fwd_with_block_score_kernel( + q_ptr, # Q: n x h x d + k_cache_ptr, # K paged: max_slots x kh x d + v_cache_ptr, # V paged: max_slots x kh x d + sink_ptr, # Sink: h x d + o_ptr, # O: n x h x d + score_ptr, # Score: h x n x max_seqblock + req_to_token_ptr, # req_to_token: max_reqs x max_kv_len + # seqlens + cu_seqlens, + seq_lens, + prefix_lens, + slot_ids, + # shape + max_slots, + num_heads, + gqa_group_size, + qk_head_dim, + v_head_dim, + block_size: tl.constexpr, + # sm_scale + sm_scale, + # gumbel topk + use_gumbel_topk: tl.constexpr, + gumbel_seed, + # stride + stride_q_n, + stride_q_h, + stride_q_d, + stride_k_s, + stride_k_h, + stride_k_d, + stride_v_s, + stride_v_h, + stride_v_d, + stride_sink_h, + stride_sink_d, + stride_o_n, + stride_o_h, + stride_o_d, + stride_s_h, + stride_s_q, + stride_s_k, + stride_r2t_b, + # META parameters + BLOCK_SIZE_Q: tl.constexpr, # q block size + BLOCK_SIZE_K: tl.constexpr, # k block size + BLOCK_SIZE_KD: tl.constexpr, + BLOCK_SIZE_VD: tl.constexpr, + # has sink + HAS_SINK: tl.constexpr, + SCORE_TYPE: tl.constexpr, + DISABLE_INDEX_VALUE: tl.constexpr, +): + tl.static_assert(SCORE_TYPE == "max" or SCORE_TYPE == "lse") + sm_scale_log2e = sm_scale * 1.4426950409 + tl.static_assert(BLOCK_SIZE_K >= block_size) + BLOCKS_PER_K_BLOCK: tl.constexpr = BLOCK_SIZE_K // block_size + # get batch id and head id + pid_q, pid_bh = tl.program_id(0), tl.program_id(1) + pid_b = pid_bh // num_heads + pid_h = pid_bh % num_heads + pid_kh = pid_h // gqa_group_size + # get q k start and len after rmpad + seq_start = tl.load(cu_seqlens + pid_b) + q_len = tl.load(cu_seqlens + pid_b + 1) - seq_start + seq_len = tl.load(seq_lens + pid_b) + prefix_len = tl.load(prefix_lens + pid_b) + sid = ( + tl.load(slot_ids + pid_b).to(tl.int64) + max_slots + ) % max_slots # safety against negative + if BLOCK_SIZE_Q * pid_q >= q_len: + return + block_num = (seq_len + block_size - 1) // block_size + # init qkv pointer + q_ptrs = tl.make_block_ptr( + base=q_ptr + seq_start * stride_q_n + pid_h * stride_q_h, + shape=(q_len, qk_head_dim), + strides=(stride_q_n, stride_q_d), + offsets=(pid_q * BLOCK_SIZE_Q, 0), + block_shape=(BLOCK_SIZE_Q, BLOCK_SIZE_KD), + order=(1, 0), + ) + s_ptrs = tl.make_block_ptr( + base=score_ptr + seq_start * stride_s_q + pid_h * stride_s_h, + shape=(q_len, block_num), + strides=(stride_s_q, stride_s_k), + offsets=(pid_q * BLOCK_SIZE_Q, 0), + block_shape=(BLOCK_SIZE_Q, BLOCKS_PER_K_BLOCK), + order=(1, 0), + ) + # load q + q = tl.load(q_ptrs, boundary_check=(0, 1), padding_option="zero") + if HAS_SINK: + off_d = tl.arange(0, BLOCK_SIZE_KD) + sink = tl.load( + sink_ptr + pid_h * stride_sink_h + off_d * stride_sink_d, + mask=off_d < qk_head_dim, + other=0, + ) + # init statistics + off_q = tl.arange(0, BLOCK_SIZE_Q) + pid_q * BLOCK_SIZE_Q + prefix_len + off_k = tl.arange(0, BLOCK_SIZE_K) + off_kd = tl.arange(0, BLOCK_SIZE_KD) + off_vd = tl.arange(0, BLOCK_SIZE_VD) + off_bpk = tl.arange(0, BLOCKS_PER_K_BLOCK) + kd_mask = off_kd < qk_head_dim + vd_mask = off_vd < v_head_dim + if HAS_SINK: + m_i = tl.zeros((BLOCK_SIZE_Q,), dtype=tl.float32) + lse_i = tl.zeros((BLOCK_SIZE_Q,), dtype=tl.float32) + qsink = tl.sum(q * sink[None, :], axis=1) * sm_scale_log2e # (BLOCK_SIZE_Q,) + m_i += qsink + lse_i += qsink + else: + m_i = tl.full((BLOCK_SIZE_Q,), float("-inf"), dtype=tl.float32) + lse_i = tl.full((BLOCK_SIZE_Q,), float("-inf"), dtype=tl.float32) + acc_o = tl.full((BLOCK_SIZE_Q, BLOCK_SIZE_VD), 0, dtype=tl.float32) + # attention + diag_start = (prefix_len + pid_q * BLOCK_SIZE_Q) // BLOCK_SIZE_K * BLOCK_SIZE_K + hi = min(seq_len, prefix_len + (pid_q + 1) * BLOCK_SIZE_Q) + for i in tl.range(0, hi, BLOCK_SIZE_K): + # paged load K via req_to_token: pos -> slot -> k_cache + pos = i + off_k + pos_mask = pos < seq_len + slots = tl.load( + req_to_token_ptr + sid * stride_r2t_b + pos, + mask=pos_mask, + other=0, + ).to(tl.int64) + slots = (slots + max_slots) % max_slots # safety against negative + # k shape: [BLOCK_SIZE_KD, BLOCK_SIZE_K] (transposed for tl.dot) + k = tl.load( + k_cache_ptr + + slots[None, :] * stride_k_s + + pid_kh * stride_k_h + + off_kd[:, None] * stride_k_d, + mask=kd_mask[:, None] & pos_mask[None, :], + other=0.0, + ) + # compute qk + qk = tl.dot(q, k) * sm_scale_log2e + if i >= diag_start: + qk = tl.where(off_q[:, None] >= (i + off_k)[None, :], qk, float("-inf")) + # K boundary mask: positions beyond seq_len contribute -inf + qk += tl.where(pos_mask[None, :], 0, float("-inf")) + # save score + score = tl.reshape( + qk, (BLOCK_SIZE_Q, BLOCKS_PER_K_BLOCK, block_size), can_reorder=False + ) + sub_max = tl.max(score, axis=2) + if SCORE_TYPE == "max": + score = sub_max + else: # "lse" + # fully-masked sub-blocks produce NaN via -inf - (-inf); clamp + # back to -inf so downstream bitonic sort sees a clean sentinel. + score = sub_max + tl.log2( + tl.sum(tl.exp2(score - sub_max[:, :, None]), axis=2) + ) + score = tl.where(score != score, float("-inf"), score) + if use_gumbel_topk: + # generate non-conflicting offset for random generation + # noise_offset shape: (BLOCK_SIZE_Q, BLOCKS_PER_K_BLOCK) + # random seed include head id, batch id and gumbel seed + # (Head low 7 bits | Batch middle 12 bits | Other high bits) + local_seed = (pid_h | (pid_b << 7) | (gumbel_seed << 19)).to(tl.int32) + # noise offset include q index and k block index + # [31-13: Q (19bits)] | [12-0: K_Block (13bits)] + noise_offset = (off_q << 13)[:, None] | (off_bpk + i // block_size)[None, :] + # gumbel noise (scaled to log2 scale to match sm_scale_log2e) + noise = tl.rand(local_seed, offset=noise_offset) + noise = tl.clamp(noise, min=1e-9, max=1 - 1e-9) # avoid log(0) + noise = -tl.log(-tl.log(noise)) * 1.4426950409 + score += noise + tl.store(s_ptrs, score.to(score_ptr.dtype.element_ty), boundary_check=(0, 1)) + if not DISABLE_INDEX_VALUE: + # compute m_ij and l_ij + m_ij = tl.maximum(m_i, tl.max(qk, axis=1)) + p = tl.exp2(qk - m_ij[:, None]) + l_ij = tl.sum(p, axis=1) + # scale acc_o + acc_o_scale = tl.exp2(m_i - m_ij) + acc_o = acc_o * acc_o_scale[:, None] + # paged load V + v = tl.load( + v_cache_ptr + + slots[:, None] * stride_v_s + + pid_kh * stride_v_h + + off_vd[None, :] * stride_v_d, + mask=pos_mask[:, None] & vd_mask[None, :], + other=0.0, + ) + p = p.to(v.dtype) + acc_o += tl.dot(p, v) + # update statistics + m_i = m_ij + lse_i = m_ij + tl.log2(tl.exp2(lse_i - m_ij) + l_ij) + # update ptrs + s_ptrs = tl.advance(s_ptrs, (0, BLOCKS_PER_K_BLOCK)) + if not DISABLE_INDEX_VALUE: + # final scale + acc_o = acc_o * tl.exp2(m_i - lse_i)[:, None] + # save output + o_ptrs = tl.make_block_ptr( + base=o_ptr + seq_start * stride_o_n + pid_h * stride_o_h, + shape=(q_len, v_head_dim), + strides=(stride_o_n, stride_o_d), + offsets=(pid_q * BLOCK_SIZE_Q, 0), + block_shape=(BLOCK_SIZE_Q, BLOCK_SIZE_VD), + order=(1, 0), + ) + tl.store(o_ptrs, acc_o.to(o_ptr.dtype.element_ty), boundary_check=(0, 1)) + + +@triton.heuristics({"BLOCK_SIZE_T": lambda args: triton.next_power_of_2(args["topk"])}) +@triton.autotune( + configs=[ + # Large configs for H200/B200 to support larger topk + triton.Config({"BLOCK_SIZE_K": 2048}, num_warps=8, num_stages=2), + triton.Config({"BLOCK_SIZE_K": 1024}, num_warps=8, num_stages=2), + triton.Config({"BLOCK_SIZE_K": 512}, num_warps=8, num_stages=2), + triton.Config({"BLOCK_SIZE_K": 256}, num_warps=8, num_stages=2), + triton.Config({"BLOCK_SIZE_K": 256}, num_warps=4, num_stages=2), + triton.Config({"BLOCK_SIZE_K": 128}, num_warps=4, num_stages=2), + triton.Config({"BLOCK_SIZE_K": 128}, num_warps=4, num_stages=3), + triton.Config({"BLOCK_SIZE_K": 64}, num_warps=2, num_stages=2), + ], + key=[ + "BLOCK_SIZE_T" + ], # use BLOCK_SIZE_T instead of topk to reduce autotune frequency +) +@triton.jit +def _topk_index_kernel( + s_ptr, # Score: h x n x max_seqblock + ti_ptr, # topk_idx: h x n x topk + # size + sample_interval: tl.constexpr, + block_size: tl.constexpr, + # seqlens + cu_seqlens, + cu_seqblocks_q, + prefix_lens, + # shape + topk, # not constexpr to avoid recompilation when topk changes + init_blocks: tl.constexpr, + local_blocks: tl.constexpr, + # stride + stride_s_h, + stride_s_n, + stride_s_k, + stride_ti_h, + stride_ti_n, + stride_ti_t, + # META parameters + BLOCK_SIZE_K: tl.constexpr, + BLOCK_SIZE_T: tl.constexpr, + MASK_INIT: tl.constexpr, + MASK_LOCAL: tl.constexpr, +): + tl.static_assert( + BLOCK_SIZE_K > BLOCK_SIZE_T + ) # use BLOCK_SIZE_T instead of topk (stricter but safe) + # get batch id and head id + pid_q = tl.program_id(0) + pid_b = tl.program_id(1) + pid_h = tl.program_id(2) + # get q k start and len after rmpad + seq_start = tl.load(cu_seqlens + pid_b) + block_start = tl.load(cu_seqblocks_q + pid_b) + block_num = tl.load(cu_seqblocks_q + pid_b + 1) - block_start + prefix_len = tl.load(prefix_lens + pid_b) + if pid_q >= block_num: + return + # offsets + off_k = tl.arange(0, BLOCK_SIZE_K) + off_t = tl.arange(0, BLOCK_SIZE_T) + # init qkv pointer + s_ptrs = ( + s_ptr + + (seq_start + pid_q * sample_interval) * stride_s_n + + pid_h * stride_s_h + + off_k * stride_s_k + ) + # init statistics + topk_score = tl.full((BLOCK_SIZE_K,), -1e30, dtype=tl.float32) + topk_idx = tl.full((BLOCK_SIZE_K,), 0, dtype=tl.int32) + left_half_mask = tl.arange(0, BLOCK_SIZE_K) < BLOCK_SIZE_K // 2 + # compute topk + valid_blocks = (prefix_len + pid_q * sample_interval + block_size) // block_size + for i in tl.range(0, valid_blocks, BLOCK_SIZE_K): + # masks + causal_mask = i + off_k < valid_blocks + local_mask = i + off_k >= max(0, valid_blocks - local_blocks) + init_mask = i + off_k < init_blocks + # load score + score = tl.load(s_ptrs, mask=causal_mask, other=-1e30).to(tl.float32) + # handle NaN: NaN inputs cause bitonic sort to fail, resulting in invalid indices (-2) + # appearing in the topk list. We replace NaN with -inf to maintain sort order. + score = tl.where(score != score, -1e30, score) + s_ptrs = s_ptrs + stride_s_k * BLOCK_SIZE_K + # fill init and local part, make sure init part is always in topk + # and at the first position. Note: must use causal_mask to protect + # init_mask to avoid selecting blocks outside causal window + if MASK_INIT: + score = tl.where(causal_mask & init_mask, score - 1e29, score) + else: + score = tl.where(causal_mask & init_mask, 1e30, score) + if MASK_LOCAL: + score = tl.where(causal_mask & local_mask, score - 1e28, score) + else: + score = tl.where(causal_mask & local_mask, 1e29, score) + # bitonic merge + topk_score, last_topk_score = score, topk_score + topk_idx, last_topk_idx = (tl.where(causal_mask, i + off_k + 1, 0), topk_idx) + n_dims: tl.constexpr = tl.standard._log2(BLOCK_SIZE_K) + for j in tl.static_range(1, n_dims): + topk_score, topk_idx = _bitonic_merge( + topk_score, topk_idx.to(tl.int32), j, 2, n_dims + ) + if i != 0: + topk_score, topk_idx = _bitonic_merge( + topk_score, topk_idx.to(tl.int32), n_dims, False, n_dims + ) + topk_score_new = last_topk_score * left_half_mask + topk_score * ( + 1 - left_half_mask + ) + topk_idx_new = last_topk_idx * left_half_mask + topk_idx * ( + 1 - left_half_mask + ) + topk_score, topk_idx = _bitonic_merge( + topk_score_new, topk_idx_new.to(tl.int32), n_dims, True, n_dims + ) + else: + topk_score, topk_idx = _bitonic_merge( + topk_score, topk_idx.to(tl.int32), n_dims, True, n_dims + ) + # get topk, shape: [BLOCK_SIZE_T,] + topk_mask = tl.arange(0, BLOCK_SIZE_K // BLOCK_SIZE_T) == 0 + topk_idx = tl.sum( + topk_mask[:, None] + * tl.reshape(topk_idx - 1, [BLOCK_SIZE_K // BLOCK_SIZE_T, BLOCK_SIZE_T]), + axis=0, + ) + # save topk + ti_ptrs = ( + ti_ptr + + (block_start + pid_q) * stride_ti_n + + pid_h * stride_ti_h + + off_t * stride_ti_t + ) + topk_mask = tl.arange(0, BLOCK_SIZE_T) < min(topk, valid_blocks) + tl.store(ti_ptrs, topk_idx.to(ti_ptrs.dtype.element_ty), mask=topk_mask) + + +@torch.no_grad() +def flash_prefill_with_topk_index( + q: torch.Tensor, + k_cache: torch.Tensor, # paged + v_cache: Optional[torch.Tensor], # paged; ignored when disable_index_value=True + sink: Optional[torch.Tensor], + req_to_token: torch.Tensor, + slot_ids: torch.Tensor, + cu_seqlens: torch.Tensor, + seq_lens: torch.Tensor, + prefix_lens: torch.Tensor, + max_seqlen_q: int, + max_seqlen_k: int, + block_size_q: int, + block_size_k: int, + topk: int, + init_blocks: int = 1, + local_blocks: int = 2, + sm_scale: Optional[float] = None, + use_tma: bool = False, + score_type: str = "max", + disable_index_value: bool = False, + cu_seqblocks_q: Optional[torch.Tensor] = None, + max_seqblock_q: Optional[int] = None, + all_seqblock_q: Optional[int] = None, +): + assert score_type in ( + "max", + "lse", + ), f"score_type must be 'max' or 'lse', got {score_type!r}" + triton.set_allocator(robust_allocator) + # dtype check + assert q.dtype == torch.bfloat16 or q.dtype == torch.float16 + assert k_cache.dtype == q.dtype + assert cu_seqlens.dtype == torch.int32 + # shape + total_q, num_heads, qk_head_dim = q.shape + max_slots, num_kv_heads, _ = k_cache.shape + if disable_index_value: + # placeholder for BLOCK_SIZE_VD; V is never loaded + v_head_dim = qk_head_dim + else: + assert v_cache is not None and v_cache.dtype == q.dtype + assert v_cache.shape[1] == k_cache.shape[1] + v_head_dim = v_cache.shape[-1] + gqa_group_size = num_heads // num_kv_heads + batch_size = cu_seqlens.shape[0] - 1 + assert qk_head_dim <= 256 and v_head_dim <= 256, "head_dim must be less than 256" + if sink is not None: + assert sink.shape[0] == num_heads and sink.shape[1] == qk_head_dim + assert ( + init_blocks + local_blocks <= topk + ), "init_blocks + local_blocks must be less than topk" + if sm_scale is None: + sm_scale = qk_head_dim**-0.5 + if cu_seqblocks_q is None or max_seqblock_q is None or all_seqblock_q is None: + cu_seqblocks_q, max_seqblock_q, all_seqblock_q, _, _, _ = get_cu_seqblocks( + cu_seqlens, max_seqlen_q, block_size_q, block_size_k + ) + max_seqblock_k = triton.cdiv(max_seqlen_k, block_size_k) + if disable_index_value: + o = None + else: + o = torch.empty(total_q, num_heads, v_head_dim, dtype=q.dtype, device=q.device) + score = torch.full( + (num_heads, total_q, max_seqblock_k), + float("-inf"), + dtype=torch.float32, + device=q.device, + ) + + # launch kernel + def grid(META): + return (triton.cdiv(max_seqlen_q, META["BLOCK_SIZE_Q"]), batch_size * num_heads) + + _flash_attn_fwd_with_block_score_kernel[grid]( + q, + k_cache, + v_cache, + sink, + o, + score, + req_to_token, + cu_seqlens, + seq_lens, + prefix_lens, + slot_ids, + max_slots, + num_heads, + gqa_group_size, + qk_head_dim, + v_head_dim, + block_size_k, + sm_scale, + False, + 1, + q.stride(0), + q.stride(1), + q.stride(2), + k_cache.stride(0), + k_cache.stride(1), + k_cache.stride(2), + v_cache.stride(0) if v_cache is not None else 0, + v_cache.stride(1) if v_cache is not None else 0, + v_cache.stride(2) if v_cache is not None else 0, + sink.stride(0) if sink is not None else 0, + sink.stride(1) if sink is not None else 0, + o.stride(0) if o is not None else 0, + o.stride(1) if o is not None else 0, + o.stride(2) if o is not None else 0, + score.stride(0), + score.stride(1), + score.stride(2), + req_to_token.stride(0), + SCORE_TYPE=score_type, + DISABLE_INDEX_VALUE=disable_index_value, + ) + + # topk extraction kernel + topk_idx = torch.full( + (num_heads, all_seqblock_q, topk), + fill_value=-1, + device=score.device, + dtype=torch.int32, + ) + # launch kernel + grid = (max_seqblock_q, batch_size, num_heads) + _topk_index_kernel[grid]( + score, + topk_idx, + block_size_q, + block_size_k, + cu_seqlens, + cu_seqblocks_q, + prefix_lens, + topk, + init_blocks, + local_blocks, + score.stride(0), + score.stride(1), + score.stride(2), + topk_idx.stride(0), + topk_idx.stride(1), + topk_idx.stride(2), + MASK_INIT=False, + MASK_LOCAL=False, + ) + return o, topk_idx diff --git a/python/sglang/srt/layers/attention/minimax_sparse_ops/prefill/topk_sparse.py b/python/sglang/srt/layers/attention/minimax_sparse_ops/prefill/topk_sparse.py new file mode 100644 index 000000000..2de840a84 --- /dev/null +++ b/python/sglang/srt/layers/attention/minimax_sparse_ops/prefill/topk_sparse.py @@ -0,0 +1,356 @@ +# Copyright 2025 XunhaoLai. All rights reserved. + +from typing import Optional + +import torch +import triton +import triton.language as tl + +from ..common.utils import check_sparse_kv_fp8, get_cu_seqblocks, robust_allocator + + +@triton.heuristics( + { + "BLOCK_SIZE_KD": lambda args: triton.next_power_of_2(args["qk_head_dim"]), + "BLOCK_SIZE_VD": lambda args: triton.next_power_of_2(args["v_head_dim"]), + "BLOCK_SIZE_H": lambda args: triton.next_power_of_2( + max( + 16 // args["BLOCK_SIZE_Q"], + triton.next_power_of_2(args["gqa_group_size"]), + ) + ), + "BLOCK_SIZE_T": lambda args: triton.next_power_of_2(args["max_topk"]), + "BLOCK_SIZE_QH": lambda args: args["BLOCK_SIZE_Q"] * args["BLOCK_SIZE_H"], + "HAS_SINK": lambda args: args["sink_ptr"] is not None, + } +) +@triton.autotune( + # Configs that fail to compile on the target arch are skipped, so widening + # the num_warps x num_stages grid only adds candidates, never a bad kernel. + configs=[ + triton.Config({}, num_warps=nw, num_stages=ns) + for nw in (2, 4, 8) + for ns in (2, 3, 4) + ], + key=[ + "BLOCK_SIZE_Q", + "BLOCK_SIZE_K", + "qk_head_dim", + "v_head_dim", + "gqa_group_size", + ], +) +@triton.jit +def _gqa_share_sparse_fwd_kernel( + q_ptr, # Q: n x h x d + k_cache_ptr, # K paged: max_slots x kh x d + v_cache_ptr, # V paged: max_slots x kh x d + sink_ptr, # Sink: h x d + t_ptr, # topk_idx: kh x n x k + o_ptr, # O: n x h x d + req_to_token_ptr, # req_to_token: max_reqs x max_kv_len + # seqlens + cu_seqlens_q, + cu_seqblocks_q, + seq_lens, + prefix_lens, + slot_ids, + # shape + max_slots, + num_kv_heads, + gqa_group_size, + qk_head_dim, + v_head_dim, + max_topk, + # q loop num + num_q_loop, + # sm_scale + sm_scale, + # stride + stride_qn, + stride_qh, + stride_qd, + stride_ks, + stride_kh, + stride_kd, + stride_vs, + stride_vh, + stride_vd, + stride_sh, + stride_sd, + stride_th, + stride_tn, + stride_tk, + stride_on, + stride_oh, + stride_od, + stride_r2t_b, + # META parameters + BLOCK_SIZE_Q: tl.constexpr, # q block size + BLOCK_SIZE_K: tl.constexpr, # k block size + BLOCK_SIZE_KD: tl.constexpr, + BLOCK_SIZE_VD: tl.constexpr, + BLOCK_SIZE_H: tl.constexpr, + BLOCK_SIZE_T: tl.constexpr, + BLOCK_SIZE_QH: tl.constexpr, + # has sink + HAS_SINK: tl.constexpr, + USE_TMA: tl.constexpr, + IS_FP8: tl.constexpr, +): + sm_scale_log2e = sm_scale * 1.4426950409 + # get batch id and head id + pid_q = tl.program_id(0) + pid_kh = tl.program_id(1) + pid_b = tl.program_id(2) + pid_h = pid_kh * gqa_group_size + # get q k start and len after rmpad + q_start = tl.load(cu_seqlens_q + pid_b) + q_len = tl.load(cu_seqlens_q + pid_b + 1) - q_start + q_block_start = tl.load(cu_seqblocks_q + pid_b) + q_block_len = tl.load(cu_seqblocks_q + pid_b + 1) - q_block_start + seq_len = tl.load(seq_lens + pid_b) + prefix_len = tl.load(prefix_lens + pid_b) + sid = ( + tl.load(slot_ids + pid_b).to(tl.int64) + max_slots + ) % max_slots # safety against negative + if pid_q * num_q_loop >= q_block_len: + return + real_q_loop = min(num_q_loop, q_block_len - pid_q * num_q_loop) + if HAS_SINK: + sink_ptrs = tl.make_block_ptr( + base=sink_ptr + pid_h * stride_sh, + shape=(gqa_group_size, qk_head_dim), + strides=(stride_sh, stride_sd), + offsets=(0, 0), + block_shape=(BLOCK_SIZE_H, BLOCK_SIZE_KD), + order=(1, 0), + ) + sink = tl.load(sink_ptrs, boundary_check=(0, 1), padding_option="zero").to( + tl.float32 + ) + # offsets for paged K/V load + off_n = tl.arange(0, BLOCK_SIZE_K) + off_kd = tl.arange(0, BLOCK_SIZE_KD) + off_vd = tl.arange(0, BLOCK_SIZE_VD) + kd_mask = off_kd < qk_head_dim + vd_mask = off_vd < v_head_dim + for j in range(real_q_loop): + pid_q_j = pid_q * num_q_loop + j + # init topk idx pointer + t_ptr_j = t_ptr + (q_block_start + pid_q_j) * stride_tn + pid_kh * stride_th + # we assume that the topk_idx is right padded with -1 + off_t = tl.arange(0, BLOCK_SIZE_T) + topk_idx = tl.load(t_ptr_j + off_t * stride_tk, mask=off_t < max_topk, other=-1) + valid_idx = tl.where(topk_idx >= 0, off_t, -1) + real_topk = tl.sum(valid_idx != -1, axis=0) + # init qkv pointer + q_ptrs = tl.make_block_ptr( + base=q_ptr + q_start * stride_qn + pid_h * stride_qh, + shape=(q_len, gqa_group_size, qk_head_dim), + strides=(stride_qn, stride_qh, stride_qd), + offsets=(pid_q_j * BLOCK_SIZE_Q, 0, 0), + block_shape=(BLOCK_SIZE_Q, BLOCK_SIZE_H, BLOCK_SIZE_KD), + order=(2, 1, 0), + ) + # load q, shape: [BLOCK_SIZE_Q, BLOCK_SIZE_H, BLOCK_SIZE_D] -> [BLOCK_SIZE_QH, BLOCK_SIZE_D] + q = tl.load(q_ptrs, boundary_check=(0, 1, 2), padding_option="zero") + # init statistics + off_q_k = ( + tl.arange(0, BLOCK_SIZE_Q)[:, None] + + pid_q_j * BLOCK_SIZE_Q + + prefix_len + - tl.arange(0, BLOCK_SIZE_K)[None, :] + ) + if HAS_SINK: + m_i = tl.zeros((BLOCK_SIZE_Q, BLOCK_SIZE_H), dtype=tl.float32) + lse_i = tl.zeros((BLOCK_SIZE_Q, BLOCK_SIZE_H), dtype=tl.float32) + qsink = ( + tl.sum(q.to(tl.float32) * sink[None, :, :], axis=2) * sm_scale_log2e + ) # (BLOCK_SIZE_Q, BLOCK_SIZE_H) + m_i += qsink + lse_i += qsink + m_i = tl.reshape(m_i, BLOCK_SIZE_QH) + lse_i = tl.reshape(lse_i, BLOCK_SIZE_QH) + else: + m_i = tl.full((BLOCK_SIZE_QH,), float("-inf"), dtype=tl.float32) + lse_i = tl.full((BLOCK_SIZE_QH,), float("-inf"), dtype=tl.float32) + acc_o = tl.full((BLOCK_SIZE_QH, BLOCK_SIZE_VD), 0, dtype=tl.float32) + q = tl.reshape(q, BLOCK_SIZE_QH, BLOCK_SIZE_KD) + # sparse attention + for i in range(real_topk): + # get current block start index (absolute K position) + c = tl.load(t_ptr_j).to(tl.int32) * BLOCK_SIZE_K + t_ptr_j = t_ptr_j + stride_tk + # paged load K via req_to_token: pos -> slot -> k_cache + pos = c + off_n + pos_mask = pos < seq_len + slots = tl.load( + req_to_token_ptr + sid * stride_r2t_b + pos, + mask=pos_mask, + other=0, + ).to(tl.int64) + slots = (slots + max_slots) % max_slots # safety against negative + # k shape: [BLOCK_SIZE_KD, BLOCK_SIZE_K] (transposed for tl.dot) + k = tl.load( + k_cache_ptr + + slots[None, :] * stride_ks + + pid_kh * stride_kh + + off_kd[:, None] * stride_kd, + mask=kd_mask[:, None] & pos_mask[None, :], + other=0.0, + ) + if IS_FP8: + # fp8 main K cache is unit-scaled; widen to the Q compute dtype + # before the tl.dot (compiled out when the cache is bf16). + k = k.to(q.dtype) + # compute qk + qk = tl.zeros((BLOCK_SIZE_Q, BLOCK_SIZE_H, BLOCK_SIZE_K), dtype=tl.float32) + # causal mask + qk += tl.where(off_q_k[:, None, :] >= c, 0, float("-inf")) + qk = tl.reshape(qk, BLOCK_SIZE_QH, BLOCK_SIZE_K) + # [BLOCK_SIZE_QH, qk_head_dim] @ [qk_head_dim, BLOCK_SIZE_K] + # -> [BLOCK_SIZE_QH, BLOCK_SIZE_K] + qk += tl.dot(q, k) * sm_scale_log2e + # K boundary mask: positions beyond seq_len contribute -inf + qk += tl.where(pos_mask[None, :], 0, float("-inf")) + # compute m_ij and l_ij + m_ij = tl.maximum(m_i, tl.max(qk, axis=1)) + p = tl.exp2(qk - m_ij[:, None]) + l_ij = tl.sum(p, axis=1) + # scale acc_o + acc_o_scale = tl.exp2(m_i - m_ij) + acc_o = acc_o * acc_o_scale[:, None] + # paged load V + v = tl.load( + v_cache_ptr + + slots[:, None] * stride_vs + + pid_kh * stride_vh + + off_vd[None, :] * stride_vd, + mask=pos_mask[:, None] & vd_mask[None, :], + other=0.0, + ) + if IS_FP8: + # Widen V so `p.to(v.dtype)` casts P to the compute dtype rather + # than to fp8 (which would wreck attention-weight precision). + v = v.to(q.dtype) + p = p.to(v.dtype) + acc_o += tl.dot(p, v) + # update statistics + m_i = m_ij + lse_i = m_ij + tl.log2(tl.exp2(lse_i - m_ij) + l_ij) + # final scale + acc_o = acc_o * tl.exp2(m_i - lse_i)[:, None] + # save output + acc_o = tl.reshape(acc_o, BLOCK_SIZE_Q, BLOCK_SIZE_H, BLOCK_SIZE_VD) + o_ptrs = tl.make_block_ptr( + base=o_ptr + q_start * stride_on + pid_h * stride_oh, + shape=(q_len, gqa_group_size, v_head_dim), + strides=(stride_on, stride_oh, stride_od), + offsets=(pid_q_j * BLOCK_SIZE_Q, 0, 0), + block_shape=(BLOCK_SIZE_Q, BLOCK_SIZE_H, BLOCK_SIZE_VD), + order=(2, 1, 0), + ) + tl.store(o_ptrs, acc_o.to(o_ptr.dtype.element_ty), boundary_check=(0, 1, 2)) + + +@torch.no_grad() +def flash_prefill_with_gqa_share_sparse( + q: torch.Tensor, + k_cache: torch.Tensor, + v_cache: torch.Tensor, + sink: Optional[torch.Tensor], + req_to_token: torch.Tensor, + slot_ids: torch.Tensor, + topk_idx: torch.Tensor, + block_size_q: int, + block_size_k: int, + cu_seqlens: torch.Tensor, + seq_lens: torch.Tensor, + prefix_lens: torch.Tensor, + max_seqlen_q: int, + sm_scale: Optional[float] = None, + use_tma: bool = True, + cu_seqblocks_q: Optional[torch.Tensor] = None, + max_seqblock_q: Optional[int] = None, +) -> torch.Tensor: + triton.set_allocator(robust_allocator) + is_fp8 = check_sparse_kv_fp8(q, k_cache, v_cache, label="prefill") + assert block_size_q in {1, 2, 4, 8, 16, 32, 64} + assert block_size_k in {16, 32, 64, 128} + # shape + total_q, num_q_heads, qk_head_dim = q.shape + max_slots, num_k_heads, _ = k_cache.shape + _, num_v_heads, v_head_dim = v_cache.shape + batch_size = cu_seqlens.shape[0] - 1 + topk = topk_idx.shape[-1] + assert topk_idx.shape[0] == num_k_heads + # gqa + assert num_k_heads == num_v_heads + assert num_q_heads % num_k_heads == 0 + gqa_group_size = num_q_heads // num_k_heads + assert gqa_group_size * block_size_q <= 128 + if sm_scale is None: + sm_scale = qk_head_dim**-0.5 + if cu_seqblocks_q is None or max_seqblock_q is None: + cu_seqblocks_q, max_seqblock_q, _, _, _, _ = get_cu_seqblocks( + cu_seqlens, max_seqlen_q, block_size_q, block_size_k + ) + # output tensor + o = torch.empty(total_q, num_q_heads, v_head_dim, device=q.device, dtype=q.dtype) + # launch kernel + num_q_loop = ( + max_seqblock_q // 131072 + 1 + ) # calculate multiple queries in one kernel if seqlence length is too long + BLOCK_SIZE_Q = triton.next_power_of_2(block_size_q) + BLOCK_SIZE_K = triton.next_power_of_2(block_size_k) + grid = ( + triton.cdiv(triton.cdiv(max_seqlen_q, block_size_q), num_q_loop), + num_k_heads, + batch_size, + ) + _gqa_share_sparse_fwd_kernel[grid]( + q, + k_cache, + v_cache, + sink, + topk_idx, + o, + req_to_token, + cu_seqlens, + cu_seqblocks_q, + seq_lens, + prefix_lens, + slot_ids, + max_slots, + num_k_heads, + gqa_group_size, + qk_head_dim, + v_head_dim, + topk, + num_q_loop, + sm_scale, + q.stride(0), + q.stride(1), + q.stride(2), + k_cache.stride(0), + k_cache.stride(1), + k_cache.stride(2), + v_cache.stride(0), + v_cache.stride(1), + v_cache.stride(2), + sink.stride(0) if sink is not None else 0, + sink.stride(1) if sink is not None else 0, + topk_idx.stride(0), + topk_idx.stride(1), + topk_idx.stride(2), + o.stride(0), + o.stride(1), + o.stride(2), + req_to_token.stride(0), + BLOCK_SIZE_Q=BLOCK_SIZE_Q, + BLOCK_SIZE_K=BLOCK_SIZE_K, + USE_TMA=use_tma, + IS_FP8=is_fp8, + ) + return o diff --git a/python/sglang/srt/layers/attention/minimax_sparse_ops/tests/test_flash_with_topk_idx.py b/python/sglang/srt/layers/attention/minimax_sparse_ops/tests/test_flash_with_topk_idx.py new file mode 100644 index 000000000..4438a5632 --- /dev/null +++ b/python/sglang/srt/layers/attention/minimax_sparse_ops/tests/test_flash_with_topk_idx.py @@ -0,0 +1,480 @@ +import sys + +import pytest +import torch + +from sglang.srt.environ import envs +from sglang.srt.layers.attention.minimax_sparse_ops.decode.flash_with_topk_idx import ( + flash_decode_with_topk_idx, +) + +DEVICE = "cuda" +RTOL_VS_REF = 5e-3 +ATOL_VS_REF = 5e-3 + + +# --------------------------------------------------------------------------- +# Reference & helpers +# --------------------------------------------------------------------------- + + +def pytorch_reference( + q, + sink, + k_cache, + v_cache, + req_to_token, + seq_lens, + slot_ids, + block_size, + topk, + init_blocks, + local_blocks, + sm_scale=None, + score_type="max", +): + batch_size, num_q_heads, head_dim = q.shape + num_kv_heads = k_cache.shape[1] + gqa_group_size = num_q_heads // num_kv_heads + if sm_scale is None: + sm_scale = head_dim**-0.5 + max_sl = seq_lens.max().item() + + # Gather K/V for all batches: [BS, max_sl, kv_heads, hd] -> [BS, q_heads, max_sl, hd] + all_slots = req_to_token[:, :max_sl].long() + k = ( + k_cache[all_slots] + .float() + .permute(0, 2, 1, 3) + .repeat_interleave(gqa_group_size, dim=1) + ) + v = ( + v_cache[all_slots] + .float() + .permute(0, 2, 1, 3) + .repeat_interleave(gqa_group_size, dim=1) + ) + + # Batched QK: [BS, q_heads, max_sl] + qk = (q.float().unsqueeze(2) @ k.transpose(-1, -2)).squeeze(2) * sm_scale + seq_mask = torch.arange(max_sl, device=q.device).unsqueeze(0) < seq_lens.unsqueeze( + 1 + ) + qk = qk.masked_fill(~seq_mask.unsqueeze(1), float("-inf")) + + # Attention output (full attention on all tokens) + if sink is not None: + sink_score = (q.float() * sink.float().unsqueeze(0)).sum( + dim=-1, keepdim=True + ) * sm_scale + attn = torch.softmax(torch.cat([sink_score, qk], dim=-1), dim=-1) + o = (attn[:, :, 1:].unsqueeze(2) @ v).squeeze(2) + else: + attn = torch.softmax(qk, dim=-1) + o = (attn.unsqueeze(2) @ v).squeeze(2) + + # Block scores + topk + topk_idx = torch.full( + (num_q_heads, batch_size, topk), -1, dtype=torch.int32, device=q.device + ) + max_num_blocks = (max_sl + block_size - 1) // block_size + padded_len = max_num_blocks * block_size + qk_padded = torch.full( + (batch_size, num_q_heads, padded_len), float("-inf"), device=q.device + ) + qk_padded[:, :, :max_sl] = qk + block_scores = qk_padded.reshape( + batch_size, num_q_heads, max_num_blocks, block_size + ) + if score_type == "max": + block_scores = block_scores.max(dim=-1).values + else: + bmax = block_scores.max(dim=-1, keepdim=True).values + block_scores = bmax.squeeze(-1) + torch.log( + torch.sum(torch.exp(block_scores - bmax), dim=-1) + ) + block_scores = torch.where(block_scores.isnan(), float("-inf"), block_scores) + + for b in range(batch_size): + sl = seq_lens[b].item() + num_blocks = (sl + block_size - 1) // block_size + bs_b = block_scores[b, :, :num_blocks].clone() + if init_blocks > 0: + bs_b[:, :init_blocks] = 1e30 + if local_blocks > 0: + local_start = max(0, num_blocks - local_blocks) + bs_b[:, local_start:num_blocks] = 1e29 + actual_topk = min(topk, num_blocks) + _, tidx = bs_b.topk(actual_topk, dim=-1) + topk_idx[:, b, :actual_topk] = tidx.to(torch.int32) + + return o, topk_idx + + +def build_inputs( + batch_size, + num_q_heads, + num_kv_heads, + head_dim, + seq_lens_list, + max_kv_len=None, + with_sink=False, + dtype=torch.bfloat16, +): + if max_kv_len is None: + max_kv_len = max(seq_lens_list) + max_slots = batch_size * max_kv_len + k_cache = torch.randn(max_slots, num_kv_heads, head_dim, dtype=dtype, device=DEVICE) + v_cache = torch.randn(max_slots, num_kv_heads, head_dim, dtype=dtype, device=DEVICE) + q = torch.randn(batch_size, num_q_heads, head_dim, dtype=dtype, device=DEVICE) + req_to_token = torch.zeros(batch_size, max_kv_len, dtype=torch.int32, device=DEVICE) + slot_ids = torch.zeros(batch_size, dtype=torch.int64, device=DEVICE) + seq_lens = torch.tensor(seq_lens_list, dtype=torch.int32, device=DEVICE) + for i in range(batch_size): + base = i * max_kv_len + slot_ids[i] = i + req_to_token[i, :max_kv_len] = torch.arange( + base, base + max_kv_len, device=DEVICE + ) + sink = ( + torch.randn(num_q_heads, head_dim, dtype=dtype, device=DEVICE) + if with_sink + else None + ) + return q, sink, k_cache, v_cache, req_to_token, seq_lens, max_kv_len, slot_ids + + +def make_seq_lens(pattern, batch_size, block_size): + if pattern == "aligned": + sl = [1024] * batch_size + elif pattern == "unaligned": + base = [513, 1023, 257, 769] + sl = (base * ((batch_size + len(base) - 1) // len(base)))[:batch_size] + elif pattern == "mixed": + base = [64, 2048, 512, 128] + sl = (base * ((batch_size + len(base) - 1) // len(base)))[:batch_size] + elif pattern == "few_blocks": + base = [block_size, block_size * 2] + sl = (base * ((batch_size + len(base) - 1) // len(base)))[:batch_size] + elif pattern == "long": + sl = [524288] * batch_size + return sl, max(sl) + + +# --------------------------------------------------------------------------- +# Test cases: compact set covering all code paths + long sequence. +# --------------------------------------------------------------------------- + + +def _case(bs, nqh, nkh, hd, blk, sink, seq_pat, ib, lb, tk): + tag = ( + f"bs{bs}_gqa{nqh}:{nkh}_hd{hd}_blk{blk}" + f"_{'sink' if sink else 'nosink'}_{seq_pat}" + f"_init{ib}_local{lb}_topk{tk}" + ) + return pytest.param(bs, nqh, nkh, hd, blk, sink, seq_pat, ib, lb, tk, id=tag) + + +# fmt: off +CASES = [ + # -- Core code paths -- + _case(2, 8, 1, 128, 64, False, "aligned", 0, 0, 16), # baseline + _case(2, 8, 1, 128, 64, False, "unaligned", 2, 4, 16), # unaligned + init+local + _case(2, 8, 1, 128, 32, False, "mixed", 0, 0, 16), # block_size=32 + mixed + _case(2, 8, 1, 128, 64, True, "aligned", 0, 0, 16), # sink + _case(8, 8, 1, 128, 64, True, "unaligned", 2, 4, 16), # sink + init+local + _case(8, 8, 1, 128, 64, False, "aligned", 0, 0, 1), # topk=1 + _case(32, 8, 1, 128, 64, False, "few_blocks", 0, 0, 32), # topk > num_blocks + _case(32, 32, 8, 128, 64, False, "aligned", 0, 0, 16), # GQA 32:8 + _case(128, 8, 4, 64, 64, False, "aligned", 0, 0, 16), # head_dim=64 + _case(128, 8, 4, 128, 64, False, "few_blocks", 0, 0, 1), # single batch, 1 block, topk=1 + # -- Long sequence (512k) -- + _case(1, 8, 1, 128, 64, False, "long", 0, 0, 16), # 512k baseline + _case(1, 8, 1, 128, 64, False, "long", 2, 4, 16), # 512k + init+local + _case(1, 8, 1, 128, 64, True, "long", 0, 0, 16), # 512k + sink +] +# fmt: on + + +# --------------------------------------------------------------------------- +# Tests +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize("score_type", ["max", "lse"]) +@pytest.mark.parametrize( + "bs,nqh,nkh,hd,blk,with_sink,seq_pat,ib,lb,tk", + CASES, +) +def test_flash_decode_with_topk_idx( + bs, nqh, nkh, hd, blk, with_sink, seq_pat, ib, lb, tk, score_type +): + torch.manual_seed(42) + seq_lens, mkl = make_seq_lens(seq_pat, bs, blk) + + q, sink, k_cache, v_cache, req_to_token, seq_lens_t, mkl, slot_ids = build_inputs( + bs, + nqh, + nkh, + hd, + seq_lens, + max_kv_len=mkl, + with_sink=with_sink, + ) + + o_new, topk_new, _ = flash_decode_with_topk_idx( + q, + sink, + k_cache, + v_cache, + req_to_token, + seq_lens_t, + mkl, + slot_ids, + blk, + tk, + ib, + lb, + score_type=score_type, + ) + o_ref, topk_ref = pytorch_reference( + q, + sink, + k_cache, + v_cache, + req_to_token, + seq_lens_t, + slot_ids, + blk, + tk, + ib, + lb, + score_type=score_type, + ) + o_ref = o_ref.to(o_new.dtype) + + # --- attention output vs reference --- + assert torch.allclose( + o_new.float(), o_ref.float(), rtol=RTOL_VS_REF, atol=ATOL_VS_REF + ), f"vs ref max abs diff {(o_new.float() - o_ref.float()).abs().max().item():.4e}" + + # --- topk set match vs reference --- + for h in range(nqh): + for b in range(bs): + sl = seq_lens[b] + num_blocks = (sl + blk - 1) // blk + actual_k = min(tk, num_blocks) + set_new = set(topk_new[h, b, :actual_k].tolist()) + set_ref = set(topk_ref[h, b, :actual_k].tolist()) + assert ( + set_new == set_ref + ), f"topk mismatch at h={h} b={b}: kernel={set_new} ref={set_ref}" + + # --- topk sentinel: invalid positions must be -1 --- + for b in range(bs): + sl = seq_lens[b] + num_blocks = (sl + blk - 1) // blk + actual_k = min(tk, num_blocks) + if actual_k < tk: + invalid = topk_new[:, b, actual_k:] + assert ( + invalid == -1 + ).all(), f"sentinel fail at b={b}: expected -1, got {invalid[invalid != -1].tolist()}" + + +@pytest.mark.parametrize("score_type", ["max", "lse"]) +@pytest.mark.parametrize( + "bs,nqh,nkh,hd,blk,with_sink,seq_pat,ib,lb,tk", + CASES, +) +def test_flash_decode_score_only( + bs, nqh, nkh, hd, blk, with_sink, seq_pat, ib, lb, tk, score_type +): + torch.manual_seed(42) + seq_lens, mkl = make_seq_lens(seq_pat, bs, blk) + + q, sink, k_cache, v_cache, req_to_token, seq_lens_t, mkl, slot_ids = build_inputs( + bs, + nqh, + nkh, + hd, + seq_lens, + max_kv_len=mkl, + with_sink=with_sink, + ) + + o_new, topk_new, _ = flash_decode_with_topk_idx( + q, + sink, + k_cache, + v_cache, + req_to_token, + seq_lens_t, + mkl, + slot_ids, + blk, + tk, + ib, + lb, + disable_index_value=True, + score_type=score_type, + ) + _, topk_ref = pytorch_reference( + q, + sink, + k_cache, + v_cache, + req_to_token, + seq_lens_t, + slot_ids, + blk, + tk, + ib, + lb, + score_type=score_type, + ) + + assert o_new is None, "expected None output when disable_index_value=True" + + for h in range(nqh): + for b in range(bs): + sl = seq_lens[b] + num_blocks = (sl + blk - 1) // blk + actual_k = min(tk, num_blocks) + set_new = set(topk_new[h, b, :actual_k].tolist()) + set_ref = set(topk_ref[h, b, :actual_k].tolist()) + assert ( + set_new == set_ref + ), f"topk mismatch at h={h} b={b}: kernel={set_new} ref={set_ref}" + + for b in range(bs): + sl = seq_lens[b] + num_blocks = (sl + blk - 1) // blk + actual_k = min(tk, num_blocks) + if actual_k < tk: + invalid = topk_new[:, b, actual_k:] + assert ( + invalid == -1 + ).all(), f"sentinel fail at b={b}: expected -1, got {invalid[invalid != -1].tolist()}" + + +def test_flash_decode_jit_topk_trivial_rows_skip_score_writes(): + torch.manual_seed(123) + bs, nqh, nkh, hd, blk, tk = 4, 8, 1, 128, 64, 32 + seq_lens, mkl = make_seq_lens("few_blocks", bs, blk) + + q, sink, k_cache, v_cache, req_to_token, seq_lens_t, mkl, slot_ids = build_inputs( + bs, + nqh, + nkh, + hd, + seq_lens, + max_kv_len=mkl, + ) + + with envs.SGLANG_OPT_USE_MINIMAX_DECODE_TOPK_RADIX.override(True): + o_new, topk_new, real_seq_lens = flash_decode_with_topk_idx( + q, + sink, + k_cache, + v_cache, + req_to_token, + seq_lens_t, + mkl, + slot_ids, + blk, + tk, + 0, + 0, + ) + o_ref, topk_ref = pytorch_reference( + q, + sink, + k_cache, + v_cache, + req_to_token, + seq_lens_t, + slot_ids, + blk, + tk, + 0, + 0, + ) + + assert real_seq_lens is None + assert torch.allclose( + o_new.float(), o_ref.float(), rtol=RTOL_VS_REF, atol=ATOL_VS_REF + ) + for h in range(nqh): + for b in range(bs): + sl = seq_lens[b] + num_blocks = (sl + blk - 1) // blk + actual_k = min(tk, num_blocks) + assert set(topk_new[h, b, :actual_k].tolist()) == set( + topk_ref[h, b, :actual_k].tolist() + ) + assert (topk_new[h, b, actual_k:] == -1).all() + + +def test_flash_decode_dense_page_table_trivial_rows_skip_score_writes(): + torch.manual_seed(321) + bs, nqh, nkh, hd, blk, tk, page_size = 3, 4, 1, 128, 64, 32, 1 + seq_lens, mkl = make_seq_lens("few_blocks", bs, blk) + + q, sink, k_cache, v_cache, req_to_token, seq_lens_t, mkl, slot_ids = build_inputs( + bs, + nqh, + nkh, + hd, + seq_lens, + max_kv_len=mkl, + ) + + o_new, page_table, real_seq_lens = flash_decode_with_topk_idx( + q, + sink, + k_cache, + v_cache, + req_to_token, + seq_lens_t, + mkl, + slot_ids, + blk, + tk, + 0, + 0, + use_dense_main_attn=True, + page_size=page_size, + ) + o_ref, _ = pytorch_reference( + q, + sink, + k_cache, + v_cache, + req_to_token, + seq_lens_t, + slot_ids, + blk, + tk, + 0, + 0, + ) + + assert torch.allclose( + o_new.float(), o_ref.float(), rtol=RTOL_VS_REF, atol=ATOL_VS_REF + ) + assert page_table.shape == (bs * nqh, tk * blk // page_size) + assert torch.equal(real_seq_lens.cpu(), seq_lens_t.repeat_interleave(nqh).cpu()) + + page_table_cpu = page_table.cpu() + req_to_token_cpu = req_to_token.cpu() + for b, seq_len in enumerate(seq_lens): + valid_pages = seq_len // page_size + for h in range(nqh): + row = b * nqh + h + expected = req_to_token_cpu[b, :seq_len:page_size] // page_size * nqh + h + assert torch.equal(page_table_cpu[row, :valid_pages], expected) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/python/sglang/srt/layers/attention/minimax_sparse_ops/tests/test_sparse_gqa.py b/python/sglang/srt/layers/attention/minimax_sparse_ops/tests/test_sparse_gqa.py new file mode 100644 index 000000000..0b2ec7e76 --- /dev/null +++ b/python/sglang/srt/layers/attention/minimax_sparse_ops/tests/test_sparse_gqa.py @@ -0,0 +1,362 @@ +"""Unit tests for flash_decode_with_gqa_share_sparse (sparse GQA attention). + +Tests the Triton sparse GQA kernel against a PyTorch reference that computes +attention only on the topk blocks via standard softmax, covering GQA ratios, +sink tokens, paged KV (randperm), variable seq_lens, and edge cases. +""" + +import sys + +import pytest +import torch + +from sglang.srt.layers.attention.minimax_sparse_ops.decode.topk_sparse import ( + flash_decode_with_gqa_share_sparse, +) + +DEVICE = "cuda" +RTOL = 5e-3 +ATOL = 5e-3 + + +def pytorch_sparse_gqa_reference( + q, + sink, + k_cache, + v_cache, + req_to_token, + seq_lens, + block_size, + topk_idx, + sm_scale=None, +): + """PyTorch reference: gather topk block tokens, then batched attention.""" + batch_size, num_q_heads, head_dim = q.shape + num_kv_heads = k_cache.shape[1] + gqa_group_size = num_q_heads // num_kv_heads + topk = topk_idx.shape[2] + if sm_scale is None: + sm_scale = head_dim**-0.5 + + # Build per-batch token positions from topk block indices, padded to uniform length + max_tokens = topk * block_size + all_slots = torch.zeros(batch_size, max_tokens, dtype=torch.long, device=q.device) + mask = torch.zeros(batch_size, max_tokens, dtype=torch.bool, device=q.device) + + for b in range(batch_size): + sl = seq_lens[b].item() + offset = 0 + for t in range(topk): + bi = topk_idx[0, b, t].item() + if bi < 0: + continue + start = bi * block_size + end = min(start + block_size, sl) + n = end - start + positions = torch.arange(start, end, device=q.device) + all_slots[b, offset : offset + n] = req_to_token[b, positions].long() + mask[b, offset : offset + n] = True + offset += n + + # Gather K/V: [BS, max_tokens, num_kv_heads, hd] -> [BS, num_q_heads, max_tokens, hd] + k = k_cache[all_slots].float() # [BS, max_tokens, num_kv_heads, hd] + v = v_cache[all_slots].float() + k = k.permute(0, 2, 1, 3).repeat_interleave( + gqa_group_size, dim=1 + ) # [BS, num_q_heads, max_tokens, hd] + v = v.permute(0, 2, 1, 3).repeat_interleave(gqa_group_size, dim=1) + + # QK: [BS, num_q_heads, 1, hd] @ [BS, num_q_heads, hd, max_tokens] -> [BS, num_q_heads, max_tokens] + qk = (q.float().unsqueeze(2) @ k.transpose(-1, -2)).squeeze(2) * sm_scale + # Mask invalid positions + qk = qk.masked_fill(~mask.unsqueeze(1), float("-inf")) + + if sink is not None: + sink_score = (q.float() * sink.float().unsqueeze(0)).sum( + dim=-1, keepdim=True + ) * sm_scale # [BS, num_q_heads, 1] + qk = torch.cat([sink_score, qk], dim=-1) # [BS, num_q_heads, 1+max_tokens] + attn = torch.softmax(qk, dim=-1) + o = (attn[:, :, 1:].unsqueeze(2) @ v).squeeze(2) # [BS, num_q_heads, hd] + else: + attn = torch.softmax(qk, dim=-1) + o = (attn.unsqueeze(2) @ v).squeeze(2) + + return o + + +def build_inputs( + batch_size, + num_q_heads, + num_kv_heads, + head_dim, + seq_lens_list, + block_size, + topk, + with_sink=False, + paged=True, + dtype=torch.bfloat16, +): + max_kv_len = max(seq_lens_list) + max_slots = batch_size * max_kv_len + + q = torch.randn(batch_size, num_q_heads, head_dim, dtype=dtype, device=DEVICE) + k_cache = torch.randn(max_slots, num_kv_heads, head_dim, dtype=dtype, device=DEVICE) + v_cache = torch.randn(max_slots, num_kv_heads, head_dim, dtype=dtype, device=DEVICE) + req_to_token = torch.zeros(batch_size, max_kv_len, dtype=torch.int32, device=DEVICE) + slot_ids = torch.zeros(batch_size, dtype=torch.int64, device=DEVICE) + seq_lens = torch.tensor(seq_lens_list, dtype=torch.int32, device=DEVICE) + + for i in range(batch_size): + base = i * max_kv_len + slot_ids[i] = i + if paged: + req_to_token[i, :max_kv_len] = ( + torch.randperm(max_kv_len, device=DEVICE) + base + ).to(torch.int32) + else: + req_to_token[i, :max_kv_len] = torch.arange( + base, base + max_kv_len, device=DEVICE + ).to(torch.int32) + + num_blocks_list = [(sl + block_size - 1) // block_size for sl in seq_lens_list] + actual_topk = min(topk, min(num_blocks_list)) + topk_idx = torch.zeros( + num_kv_heads, batch_size, topk, dtype=torch.int32, device=DEVICE + ) + for kh in range(num_kv_heads): + for b in range(batch_size): + nb = num_blocks_list[b] + ak = min(topk, nb) + perm = torch.randperm(nb, device=DEVICE)[:ak] + topk_idx[kh, b, :ak] = perm.to(torch.int32) + if ak < topk: + topk_idx[kh, b, ak:] = -1 + + sink = ( + torch.randn(num_q_heads, head_dim, dtype=dtype, device=DEVICE) + if with_sink + else None + ) + + return q, sink, k_cache, v_cache, req_to_token, seq_lens, slot_ids, topk_idx + + +def _case(bs, nqh, nkh, hd, blk, tk, sink, seq_pat, paged=True): + tag = ( + f"bs{bs}_gqa{nqh}:{nkh}_hd{hd}_blk{blk}_topk{tk}" + f"_{'sink' if sink else 'nosink'}_{seq_pat}" + f"{'_paged' if paged else '_contig'}" + ) + return pytest.param(bs, nqh, nkh, hd, blk, tk, sink, seq_pat, paged, id=tag) + + +def make_seq_lens(pattern, batch_size, block_size): + if pattern == "aligned": + return [1024] * batch_size + elif pattern == "unaligned": + base = [513, 1023, 257, 769] + return (base * ((batch_size + len(base) - 1) // len(base)))[:batch_size] + elif pattern == "short": + return [block_size * 2] * batch_size + elif pattern == "mixed": + base = [block_size, block_size * 4, block_size * 16, block_size * 2] + return (base * ((batch_size + len(base) - 1) // len(base)))[:batch_size] + elif pattern == "long": + return [block_size * 512] * batch_size + + +CASES = [ + # Baseline: GQA 8:1 + _case(1, 8, 1, 128, 64, 16, False, "aligned"), + _case(2, 8, 1, 128, 64, 16, False, "aligned"), + _case(4, 8, 1, 128, 64, 32, False, "aligned"), + # Unaligned seq_lens + _case(4, 8, 1, 128, 64, 16, False, "unaligned"), + # Short sequences (topk ~ num_blocks) + _case(2, 8, 1, 128, 64, 16, False, "short"), + # topk=1 + _case(2, 8, 1, 128, 64, 1, False, "aligned"), + # topk=32 (production config) + _case(2, 8, 1, 128, 64, 32, False, "aligned"), + # With sink + _case(2, 8, 1, 128, 64, 16, True, "aligned"), + _case(4, 8, 1, 128, 64, 16, True, "unaligned"), + # GQA 32:8 + _case(2, 32, 8, 128, 64, 16, False, "aligned"), + # GQA 16:1 (production) + _case(2, 16, 1, 128, 64, 32, False, "aligned"), + # head_dim=64 + _case(2, 8, 4, 64, 64, 16, False, "aligned"), + # block_size=32 + _case(2, 8, 1, 128, 32, 16, False, "aligned"), + # Large BS + _case(32, 8, 1, 128, 64, 16, False, "aligned"), + _case(128, 8, 1, 128, 64, 32, False, "short"), + # Contiguous (non-paged) KV + _case(2, 8, 1, 128, 64, 16, False, "aligned", paged=False), + # Mixed seq_lens in batch + _case(4, 8, 1, 128, 64, 8, False, "mixed"), + # Long sequence + _case(1, 8, 1, 128, 64, 32, False, "long"), +] + + +@pytest.mark.parametrize("bs,nqh,nkh,hd,blk,tk,with_sink,seq_pat,paged", CASES) +def test_sparse_gqa_vs_reference(bs, nqh, nkh, hd, blk, tk, with_sink, seq_pat, paged): + """Kernel output must match PyTorch reference (attend only to topk blocks).""" + torch.manual_seed(42) + seq_lens_list = make_seq_lens(seq_pat, bs, blk) + + q, sink, k_cache, v_cache, req_to_token, seq_lens, slot_ids, topk_idx = ( + build_inputs( + bs, + nqh, + nkh, + hd, + seq_lens_list, + blk, + tk, + with_sink=with_sink, + paged=paged, + ) + ) + + o_kernel = flash_decode_with_gqa_share_sparse( + q, + sink, + k_cache, + v_cache, + req_to_token, + seq_lens, + slot_ids, + blk, + topk_idx, + ) + o_ref = pytorch_sparse_gqa_reference( + q, + sink, + k_cache, + v_cache, + req_to_token, + seq_lens, + blk, + topk_idx, + ) + o_ref = o_ref.to(o_kernel.dtype) + + assert torch.allclose( + o_kernel.float(), o_ref.float(), rtol=RTOL, atol=ATOL + ), f"max abs diff {(o_kernel.float() - o_ref.float()).abs().max().item():.4e}" + + +@pytest.mark.parametrize( + "bs,nqh,nkh,hd,blk,tk,with_sink,seq_pat,paged", + [ + _case(2, 8, 1, 128, 64, 32, False, "short"), + ], +) +def test_sparse_gqa_topk_exceeds_blocks( + bs, nqh, nkh, hd, blk, tk, with_sink, seq_pat, paged +): + """topk > num_blocks: kernel should handle gracefully (some topk_idx = -1).""" + torch.manual_seed(42) + seq_lens_list = [blk] * bs + + q, sink, k_cache, v_cache, req_to_token, seq_lens, slot_ids, topk_idx = ( + build_inputs( + bs, + nqh, + nkh, + hd, + seq_lens_list, + blk, + tk, + with_sink=with_sink, + paged=paged, + ) + ) + + o_kernel = flash_decode_with_gqa_share_sparse( + q, + sink, + k_cache, + v_cache, + req_to_token, + seq_lens, + slot_ids, + blk, + topk_idx, + ) + o_ref = pytorch_sparse_gqa_reference( + q, + sink, + k_cache, + v_cache, + req_to_token, + seq_lens, + blk, + topk_idx, + ) + o_ref = o_ref.to(o_kernel.dtype) + + assert torch.allclose( + o_kernel.float(), o_ref.float(), rtol=RTOL, atol=ATOL + ), f"max abs diff {(o_kernel.float() - o_ref.float()).abs().max().item():.4e}" + + +@pytest.mark.parametrize( + "bs,nqh,nkh,hd,blk,tk,with_sink,seq_pat,paged", + [ + _case(4, 8, 1, 128, 64, 16, False, "aligned"), + ], +) +def test_sparse_gqa_deterministic(bs, nqh, nkh, hd, blk, tk, with_sink, seq_pat, paged): + """Two calls with same inputs must produce identical outputs.""" + torch.manual_seed(42) + seq_lens_list = make_seq_lens(seq_pat, bs, blk) + + q, sink, k_cache, v_cache, req_to_token, seq_lens, slot_ids, topk_idx = ( + build_inputs( + bs, + nqh, + nkh, + hd, + seq_lens_list, + blk, + tk, + with_sink=with_sink, + paged=paged, + ) + ) + + o1 = flash_decode_with_gqa_share_sparse( + q, + sink, + k_cache, + v_cache, + req_to_token, + seq_lens, + slot_ids, + blk, + topk_idx, + ) + o2 = flash_decode_with_gqa_share_sparse( + q, + sink, + k_cache, + v_cache, + req_to_token, + seq_lens, + slot_ids, + blk, + topk_idx, + ) + + assert torch.equal( + o1, o2 + ), f"non-deterministic: max diff {(o1.float() - o2.float()).abs().max().item():.4e}" + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py index 8f74b354e..6acd20ed0 100644 --- a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py +++ b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py @@ -179,6 +179,7 @@ class FusedMoE(torch.nn.Module): with_bias=False, routing_method_type: Optional[RoutingMethodType] = None, is_gated: bool = True, + gate_up_interleaved: bool = True, ): super().__init__() if params_dtype is None: @@ -265,6 +266,7 @@ class FusedMoE(torch.nn.Module): swiglu_limit=swiglu_limit, is_gated=is_gated, routing_method_type=routing_method_type, + gate_up_interleaved=gate_up_interleaved, ) self.quant_method: Optional[FusedMoEMethodBase] = None diff --git a/python/sglang/srt/layers/moe/moe_runner/base.py b/python/sglang/srt/layers/moe/moe_runner/base.py index 6363a4de8..9a4c2aa0d 100644 --- a/python/sglang/srt/layers/moe/moe_runner/base.py +++ b/python/sglang/srt/layers/moe/moe_runner/base.py @@ -65,6 +65,10 @@ class MoeRunnerConfig: gemm1_alpha: Optional[float] = None gemm1_clamp_limit: Optional[float] = None swiglu_limit: Optional[float] = None + # Whether gate/up weights are stored interleaved (vs split). Only the + # silu+is_gated swiglu path consumes it (interleaved -> swiglu_gpt_oss_*, + # otherwise chunk gate/up then apply alpha/limit). + gate_up_interleaved: bool = True @dataclass diff --git a/python/sglang/srt/layers/moe/moe_runner/triton.py b/python/sglang/srt/layers/moe/moe_runner/triton.py index 96e431d4e..4b9cc7bf6 100644 --- a/python/sglang/srt/layers/moe/moe_runner/triton.py +++ b/python/sglang/srt/layers/moe/moe_runner/triton.py @@ -16,6 +16,7 @@ from sglang.srt.layers.moe.moe_runner.base import ( register_pre_permute, ) from sglang.srt.layers.moe.utils import MoeRunnerBackend +from sglang.srt.utils import is_cuda, is_gfx95_supported, is_hip if TYPE_CHECKING: from sglang.srt.layers.moe.token_dispatcher.standard import ( @@ -55,6 +56,7 @@ class TritonMoeQuantInfo(MoeQuantInfo): w2_weight: torch.Tensor b13: Optional[torch.Tensor] = None b2: Optional[torch.Tensor] = None + use_mxfp8: bool = False use_fp8_w8a8: bool = False use_int8_w8a8: bool = False use_int8_w8a16: bool = False @@ -81,6 +83,40 @@ class TritonRunnerCore(MoeRunnerCore): running_state: dict, hooks: Optional[Any] = None, ) -> TritonRunnerOutput: + if quant_info.use_mxfp8 and is_hip() and is_gfx95_supported(): + from sglang.srt.layers.moe.moe_runner.triton_utils.mxfp8_moe_amd_gfx95 import ( + fused_experts_mxfp8, + ) + + out = fused_experts_mxfp8( + runner_input.hidden_states, + quant_info.w13_weight, + quant_info.w2_weight, + runner_input.topk_weights, + runner_input.topk_ids, + quant_info.w13_scale, + quant_info.w2_scale, + b1=quant_info.b13, + b2=quant_info.b2, + activation=self.config.activation, + is_gated=self.config.is_gated, + no_combine=self.config.no_combine, + inplace=self.config.inplace, + apply_router_weight_on_input=self.config.apply_router_weight_on_input, + routed_scaling_factor=self.config.routed_scaling_factor, + gemm1_alpha=self.config.gemm1_alpha, + gemm1_limit=self.config.gemm1_clamp_limit, + swiglu_limit=self.config.swiglu_limit, + gate_up_interleaved=self.config.gate_up_interleaved, + ) + return TritonRunnerOutput(hidden_states=out) + + if quant_info.use_mxfp8 and is_cuda(): + raise NotImplementedError( + "Triton MoE runner does not support NVIDIA MXFP8; use " + "--moe-runner-backend deep_gemm (or flashinfer_trtllm/cutlass)." + ) + from sglang.srt.layers.moe.moe_runner.triton_utils.fused_moe import ( _fused_moe_kernel_sequence, ) @@ -142,30 +178,66 @@ def fused_experts_none_to_triton( quant_info: TritonMoeQuantInfo, runner_config: MoeRunnerConfig, ) -> StandardCombineInput: - from sglang.srt.layers.moe.moe_runner.triton_utils.fused_moe import fused_experts from sglang.srt.layers.moe.token_dispatcher.standard import StandardCombineInput - output = fused_experts( - hidden_states=dispatch_output.hidden_states, - w1=quant_info.w13_weight, - w2=quant_info.w2_weight, - topk_output=dispatch_output.topk_output, - moe_runner_config=runner_config, - b1=quant_info.b13, - b2=quant_info.b2, - use_fp8_w8a8=quant_info.use_fp8_w8a8, - use_int8_w8a8=quant_info.use_int8_w8a8, - use_int8_w8a16=quant_info.use_int8_w8a16, - use_int4_w4a16=quant_info.use_int4_w4a16, - per_channel_quant=quant_info.per_channel_quant, - w1_scale=quant_info.w13_scale, - w2_scale=quant_info.w2_scale, - w1_zp=quant_info.w13_zp, - w2_zp=quant_info.w2_zp, - a1_scale=quant_info.a13_scale, - a2_scale=quant_info.a2_scale, - block_shape=quant_info.block_shape, - ) + if quant_info.use_mxfp8 and is_hip() and is_gfx95_supported(): + from sglang.srt.layers.moe.moe_runner.triton_utils.mxfp8_moe_amd_gfx95 import ( + fused_experts_mxfp8, + ) + + topk_weights, topk_ids, _ = dispatch_output.topk_output + output = fused_experts_mxfp8( + hidden_states=dispatch_output.hidden_states, + w1=quant_info.w13_weight, + w2=quant_info.w2_weight, + topk_weights=topk_weights, + topk_ids=topk_ids, + w1_scale=quant_info.w13_scale, + w2_scale=quant_info.w2_scale, + b1=quant_info.b13, + b2=quant_info.b2, + activation=runner_config.activation, + is_gated=runner_config.is_gated, + no_combine=runner_config.no_combine, + inplace=runner_config.inplace, + apply_router_weight_on_input=runner_config.apply_router_weight_on_input, + routed_scaling_factor=runner_config.routed_scaling_factor, + gemm1_alpha=runner_config.gemm1_alpha, + gemm1_limit=runner_config.gemm1_clamp_limit, + swiglu_limit=runner_config.swiglu_limit, + gate_up_interleaved=runner_config.gate_up_interleaved, + ) + else: + if quant_info.use_mxfp8 and is_cuda(): + raise NotImplementedError( + "Triton MoE runner does not support NVIDIA MXFP8; use " + "--moe-runner-backend deep_gemm (or flashinfer_trtllm/cutlass)." + ) + from sglang.srt.layers.moe.moe_runner.triton_utils.fused_moe import ( + fused_experts, + ) + + output = fused_experts( + hidden_states=dispatch_output.hidden_states, + w1=quant_info.w13_weight, + w2=quant_info.w2_weight, + topk_output=dispatch_output.topk_output, + moe_runner_config=runner_config, + b1=quant_info.b13, + b2=quant_info.b2, + use_fp8_w8a8=quant_info.use_fp8_w8a8, + use_int8_w8a8=quant_info.use_int8_w8a8, + use_int8_w8a16=quant_info.use_int8_w8a16, + use_int4_w4a16=quant_info.use_int4_w4a16, + per_channel_quant=quant_info.per_channel_quant, + w1_scale=quant_info.w13_scale, + w2_scale=quant_info.w2_scale, + w1_zp=quant_info.w13_zp, + w2_zp=quant_info.w2_zp, + a1_scale=quant_info.a13_scale, + a2_scale=quant_info.a2_scale, + block_shape=quant_info.block_shape, + ) return StandardCombineInput( hidden_states=output, diff --git a/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/minimax_m3_gfx950_mxfp8_compact_moe.json b/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/minimax_m3_gfx950_mxfp8_compact_moe.json new file mode 100644 index 000000000..3f84875e0 --- /dev/null +++ b/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/minimax_m3_gfx950_mxfp8_compact_moe.json @@ -0,0 +1,70 @@ +{ + "model": "MiniMax-M3", + "device": "gfx950", + "experts": 128, + "hidden_size": 6144, + "intermediate_size": 384, + "top_k": 4, + "tokens": { + "1": { + "best_gemm1": {"block_m": 128, "block_n": 256, "block_k": 128, "num_stages": 1, "num_warps": 4}, + "best_gemm2": {"block_m": 128, "block_n": 256, "block_k": 128, "num_stages": 1, "num_warps": 4} + }, + "2": { + "best_gemm1": {"block_m": 64, "block_n": 128, "block_k": 128, "num_stages": 2, "num_warps": 8}, + "best_gemm2": {"block_m": 64, "block_n": 256, "block_k": 128, "num_stages": 2, "num_warps": 4} + }, + "4": { + "best_gemm1": {"block_m": 128, "block_n": 256, "block_k": 128, "num_stages": 1, "num_warps": 4}, + "best_gemm2": {"block_m": 128, "block_n": 128, "block_k": 64, "num_stages": 1, "num_warps": 4} + }, + "8": { + "best_gemm1": {"block_m": 128, "block_n": 128, "block_k": 256, "num_stages": 1, "num_warps": 4}, + "best_gemm2": {"block_m": 128, "block_n": 256, "block_k": 128, "num_stages": 2, "num_warps": 8} + }, + "16": { + "best_gemm1": {"block_m": 64, "block_n": 128, "block_k": 256, "num_stages": 2, "num_warps": 8}, + "best_gemm2": {"block_m": 64, "block_n": 256, "block_k": 128, "num_stages": 2, "num_warps": 4} + }, + "32": { + "best_gemm1": {"block_m": 64, "block_n": 128, "block_k": 128, "num_stages": 2, "num_warps": 8}, + "best_gemm2": {"block_m": 64, "block_n": 128, "block_k": 64, "num_stages": 2, "num_warps": 4} + }, + "64": { + "best_gemm1": {"block_m": 32, "block_n": 256, "block_k": 128, "num_stages": 1, "num_warps": 4}, + "best_gemm2": {"block_m": 32, "block_n": 256, "block_k": 128, "num_stages": 1, "num_warps": 4} + }, + "128": { + "best_gemm1": {"block_m": 64, "block_n": 256, "block_k": 256, "num_stages": 1, "num_warps": 4}, + "best_gemm2": {"block_m": 32, "block_n": 256, "block_k": 128, "num_stages": 1, "num_warps": 8} + }, + "256": { + "best_gemm1": {"block_m": 32, "block_n": 128, "block_k": 256, "num_stages": 1, "num_warps": 4}, + "best_gemm2": {"block_m": 32, "block_n": 256, "block_k": 128, "num_stages": 1, "num_warps": 8} + }, + "512": { + "best_gemm1": {"block_m": 64, "block_n": 128, "block_k": 256, "num_stages": 1, "num_warps": 4}, + "best_gemm2": {"block_m": 32, "block_n": 256, "block_k": 128, "num_stages": 1, "num_warps": 8} + }, + "1024": { + "best_gemm1": {"block_m": 64, "block_n": 128, "block_k": 256, "num_stages": 2, "num_warps": 8}, + "best_gemm2": {"block_m": 64, "block_n": 256, "block_k": 128, "num_stages": 1, "num_warps": 8} + }, + "1536": { + "best_gemm1": {"block_m": 64, "block_n": 128, "block_k": 256, "num_stages": 2, "num_warps": 8}, + "best_gemm2": {"block_m": 64, "block_n": 128, "block_k": 128, "num_stages": 1, "num_warps": 4} + }, + "2048": { + "best_gemm1": {"block_m": 128, "block_n": 128, "block_k": 256, "num_stages": 2, "num_warps": 8}, + "best_gemm2": {"block_m": 128, "block_n": 128, "block_k": 128, "num_stages": 2, "num_warps": 8} + }, + "3072": { + "best_gemm1": {"block_m": 128, "block_n": 256, "block_k": 128, "num_stages": 1, "num_warps": 4}, + "best_gemm2": {"block_m": 128, "block_n": 256, "block_k": 128, "num_stages": 1, "num_warps": 4} + }, + "4096": { + "best_gemm1": {"block_m": 128, "block_n": 256, "block_k": 128, "num_stages": 1, "num_warps": 4}, + "best_gemm2": {"block_m": 128, "block_n": 256, "block_k": 128, "num_stages": 1, "num_warps": 4} + } + } +} diff --git a/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=128,N=384,device_name=AMD_Instinct_MI300X,dtype=fp8_w8a8,block_shape=[128, 128].json b/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=128,N=384,device_name=AMD_Instinct_MI300X,dtype=fp8_w8a8,block_shape=[128, 128].json new file mode 100644 index 000000000..c474ad5f7 --- /dev/null +++ b/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=128,N=384,device_name=AMD_Instinct_MI300X,dtype=fp8_w8a8,block_shape=[128, 128].json @@ -0,0 +1,164 @@ +{ + "1": { + "BLOCK_SIZE_M": 32, + "BLOCK_SIZE_N": 16, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 16, + "num_warps": 2, + "num_stages": 2, + "waves_per_eu": 0 + }, + "2": { + "BLOCK_SIZE_M": 32, + "BLOCK_SIZE_N": 16, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 2, + "num_stages": 2, + "waves_per_eu": 0 + }, + "4": { + "BLOCK_SIZE_M": 32, + "BLOCK_SIZE_N": 16, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 8, + "num_warps": 2, + "num_stages": 2, + "waves_per_eu": 0 + }, + "8": { + "BLOCK_SIZE_M": 32, + "BLOCK_SIZE_N": 32, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 1, + "num_stages": 2, + "waves_per_eu": 0 + }, + "16": { + "BLOCK_SIZE_M": 32, + "BLOCK_SIZE_N": 32, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 1, + "num_stages": 2, + "waves_per_eu": 0 + }, + "24": { + "BLOCK_SIZE_M": 32, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 2, + "num_stages": 2, + "waves_per_eu": 0 + }, + "32": { + "BLOCK_SIZE_M": 32, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 2, + "waves_per_eu": 0 + }, + "48": { + "BLOCK_SIZE_M": 32, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 2, + "num_stages": 2, + "waves_per_eu": 0 + }, + "64": { + "BLOCK_SIZE_M": 32, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 8, + "num_warps": 2, + "num_stages": 2, + "waves_per_eu": 0 + }, + "96": { + "BLOCK_SIZE_M": 32, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 4, + "num_warps": 2, + "num_stages": 2, + "waves_per_eu": 0 + }, + "128": { + "BLOCK_SIZE_M": 32, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 8, + "num_warps": 2, + "num_stages": 2, + "waves_per_eu": 0 + }, + "256": { + "BLOCK_SIZE_M": 32, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 8, + "num_warps": 2, + "num_stages": 2, + "waves_per_eu": 0 + }, + "512": { + "BLOCK_SIZE_M": 32, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 4, + "num_warps": 2, + "num_stages": 2, + "waves_per_eu": 0 + }, + "1024": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 8, + "num_warps": 4, + "num_stages": 2, + "waves_per_eu": 0 + }, + "1536": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 4, + "num_warps": 2, + "num_stages": 2, + "waves_per_eu": 0 + }, + "2048": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 8, + "num_warps": 4, + "num_stages": 2, + "waves_per_eu": 0 + }, + "3072": { + "BLOCK_SIZE_M": 128, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 16, + "num_warps": 8, + "num_stages": 2, + "waves_per_eu": 0 + }, + "4096": { + "BLOCK_SIZE_M": 128, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 8, + "num_warps": 4, + "num_stages": 2, + "waves_per_eu": 0 + } +} diff --git a/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=129,N=384,device_name=NVIDIA_H200.json b/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=129,N=384,device_name=NVIDIA_H200.json new file mode 100644 index 000000000..b5d25b962 --- /dev/null +++ b/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=129,N=384,device_name=NVIDIA_H200.json @@ -0,0 +1,146 @@ +{ + "1": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 32, + "BLOCK_SIZE_K": 256, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 4 + }, + "2": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 32, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 16, + "num_warps": 4, + "num_stages": 5 + }, + "4": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 32, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 64, + "num_warps": 4, + "num_stages": 4 + }, + "8": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 32, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 4 + }, + "16": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 3 + }, + "24": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 2 + }, + "32": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 2 + }, + "48": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 32, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 64, + "num_warps": 4, + "num_stages": 2 + }, + "64": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 16, + "num_warps": 4, + "num_stages": 3 + }, + "96": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 8, + "num_stages": 3 + }, + "128": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 8, + "num_stages": 3 + }, + "256": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 8, + "num_stages": 3 + }, + "512": { + "BLOCK_SIZE_M": 32, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 16, + "num_warps": 4, + "num_stages": 2 + }, + "1024": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 1, + "num_warps": 8, + "num_stages": 3 + }, + "1536": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 1, + "num_warps": 8, + "num_stages": 3 + }, + "2048": { + "BLOCK_SIZE_M": 128, + "BLOCK_SIZE_N": 256, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 1, + "num_warps": 8, + "num_stages": 4 + }, + "3072": { + "BLOCK_SIZE_M": 128, + "BLOCK_SIZE_N": 256, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 16, + "num_warps": 8, + "num_stages": 4 + }, + "4096": { + "BLOCK_SIZE_M": 128, + "BLOCK_SIZE_N": 256, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 1, + "num_warps": 8, + "num_stages": 4 + } +} diff --git a/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe.py b/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe.py index b7245e5ec..b92654f21 100644 --- a/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe.py +++ b/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe.py @@ -122,6 +122,7 @@ def inplace_fused_experts( gemm1_limit: Optional[float] = None, filter_expert: bool = True, swiglu_limit: Optional[float] = None, + gate_up_interleaved: bool = True, ) -> None: fused_experts_impl( hidden_states, @@ -153,6 +154,7 @@ def inplace_fused_experts( gemm1_limit, filter_expert, swiglu_limit=swiglu_limit, + gate_up_interleaved=gate_up_interleaved, ) @@ -186,6 +188,7 @@ def outplace_fused_experts( gemm1_limit: Optional[float] = None, filter_expert: bool = True, swiglu_limit: Optional[float] = None, + gate_up_interleaved: bool = True, ) -> torch.Tensor: return fused_experts_impl( hidden_states, @@ -217,6 +220,7 @@ def outplace_fused_experts( gemm1_limit=gemm1_limit, filter_expert=filter_expert, swiglu_limit=swiglu_limit, + gate_up_interleaved=gate_up_interleaved, ) @@ -276,6 +280,7 @@ def fused_experts( moe_runner_config.gemm1_clamp_limit, filter_expert, swiglu_limit=moe_runner_config.swiglu_limit, + gate_up_interleaved=moe_runner_config.gate_up_interleaved, ) return hidden_states else: @@ -308,6 +313,7 @@ def fused_experts( gemm1_limit=moe_runner_config.gemm1_clamp_limit, filter_expert=filter_expert, swiglu_limit=moe_runner_config.swiglu_limit, + gate_up_interleaved=moe_runner_config.gate_up_interleaved, ) @@ -336,6 +342,14 @@ def swiglu_gpt_oss_sigmoid_alpha(x, gemm1_alpha, gemm1_limit): return gate * torch.sigmoid(gate * gemm1_alpha) * (up + 1) +@torch.compile +def swiglu_no_interleaved_with_alpha_and_limit(x, gemm1_alpha, gemm1_limit): + gate, up = x.chunk(2, dim=-1) + gate = gate.clamp(min=None, max=gemm1_limit) + up = up.clamp(min=-gemm1_limit, max=gemm1_limit) + return gate * torch.sigmoid(gate * gemm1_alpha) * (up + 1) + + @functools.lru_cache() def _down_moe_use_tma(): return support_tensor_descriptor() @@ -441,6 +455,7 @@ def _fused_moe_kernel_sequence( filter_expert: bool, hooks: Optional[Any] = None, swiglu_limit: Optional[float] = None, + gate_up_interleaved: bool = True, ) -> torch.Tensor: """Run the MoE kernel/activation/kernel/combine sequence in a single shot. @@ -538,9 +553,18 @@ def _fused_moe_kernel_sequence( # - swiglu_limit != None: DeepSeek V4 swiglu clamp + silu_and_mul (CUDA/HIP only) if gemm1_alpha is not None: assert gemm1_limit is not None - intermediate_cache2 = swiglu_gpt_oss_sigmoid_alpha( - intermediate_cache1.view(-1, N), gemm1_alpha, gemm1_limit - ) + if gate_up_interleaved: + intermediate_cache2 = swiglu_gpt_oss_sigmoid_alpha( + intermediate_cache1.view(-1, N), + gemm1_alpha, + gemm1_limit, + ) + else: + intermediate_cache2 = swiglu_no_interleaved_with_alpha_and_limit( + intermediate_cache1.view(-1, N), + gemm1_alpha, + gemm1_limit, + ) elif gemm1_limit is not None: intermediate_cache2 = _swiglu_silu_clamp_mul( intermediate_cache1.view(-1, N), gemm1_limit @@ -820,6 +844,7 @@ def fused_experts_impl( gemm1_limit: Optional[float] = None, filter_expert: bool = True, swiglu_limit: Optional[float] = None, + gate_up_interleaved: bool = True, ): padded_size = padding_size if not (use_fp8_w8a8 or use_int8_w8a8) or block_shape is not None or _use_aiter: @@ -895,6 +920,7 @@ def fused_experts_impl( filter_expert=filter_expert, hooks=None, swiglu_limit=swiglu_limit, + gate_up_interleaved=gate_up_interleaved, ) diff --git a/python/sglang/srt/layers/moe/moe_runner/triton_utils/mxfp8_moe_amd_gfx95.py b/python/sglang/srt/layers/moe/moe_runner/triton_utils/mxfp8_moe_amd_gfx95.py new file mode 100644 index 000000000..de900e57c --- /dev/null +++ b/python/sglang/srt/layers/moe/moe_runner/triton_utils/mxfp8_moe_amd_gfx95.py @@ -0,0 +1,433 @@ +"""Native MXFP8 (1x32 block, E8M0 scale) MoE for AMD CDNA4 (gfx950). + +Replaces the prior SGLang MXFP8 MoE family (dense / hybrid / packed / +grouped_gemm1 / grouped_gemm12 / compact / fused_act) with a single grouped +``tl.dot_scaled`` kernel. Instead of an explicit ``argsort`` + ``index_select`` +gather, a materialized intermediate, a separate activation quant, and a +``tl.atomic_add`` combine: + + * tokens are sorted by expert with ``moe_align_block_size``; + * GEMM1 reads the activation by token-id indirection (``a_row = token // top_k``) + so the hidden states are MXFP8-quantized exactly ONCE (not top_k times); + * the SwiGLU-OAI activation is the fused fp32 Triton kernel (split layout); + * GEMM2 applies the top-k weight inside the kernel and writes each route to a + distinct output row (no atomics); the final reduction is a strided sum. +""" + +from __future__ import annotations + +from typing import Optional + +import torch +import triton +import triton.language as tl + +from sglang.srt.environ import envs +from sglang.srt.layers.moe.moe_runner.triton_utils.moe_align_block_size import ( + moe_align_block_size, +) +from sglang.srt.layers.quantization.mxfp8_amd_gfx95 import mxfp8_e4m3_quantize + + +@triton.jit +def _mxfp8_grouped_gemm_kernel( + a_ptr, + a_scale_ptr, + b_ptr, + b_scale_ptr, + c_ptr, + topk_weights_ptr, + sorted_token_ids_ptr, + expert_ids_ptr, + num_tokens_post_padded_ptr, + E, + N, + K, + num_valid_tokens, + top_k, + stride_am, + stride_ak, + stride_asm, + stride_ask, + stride_be, + stride_bn, + stride_bk, + stride_bse, + stride_bsn, + stride_bsk, + stride_cm, + stride_cn, + A_DIV: tl.constexpr, + MUL_WEIGHT: tl.constexpr, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, + BLOCK_K: tl.constexpr, +): + pid_m = tl.program_id(0) + pid_n = tl.program_id(1) + num_post = tl.load(num_tokens_post_padded_ptr) + if pid_m * BLOCK_M >= num_post: + return + + offs_tid = pid_m * BLOCK_M + tl.arange(0, BLOCK_M) + offs_token = tl.load(sorted_token_ids_ptr + offs_tid).to(tl.int64) + token_mask = offs_token < num_valid_tokens + off_e = tl.load(expert_ids_ptr + pid_m).to(tl.int64) + valid_expert = (off_e >= 0) & (off_e < E) + + offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N) + offs_k = tl.arange(0, BLOCK_K) + offs_sk = tl.arange(0, BLOCK_K // 32) + a_row = offs_token // A_DIV + + a_ptrs = a_ptr + a_row[:, None] * stride_am + offs_k[None, :] * stride_ak + as_ptrs = a_scale_ptr + a_row[:, None] * stride_asm + offs_sk[None, :] * stride_ask + b_ptrs = ( + b_ptr + + off_e * stride_be + + offs_n[:, None] * stride_bn + + offs_k[None, :] * stride_bk + ) + bs_ptrs = ( + b_scale_ptr + + off_e * stride_bse + + offs_n[:, None] * stride_bsn + + offs_sk[None, :] * stride_bsk + ) + + acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) + n_mask = offs_n < N + for _ in range(0, tl.cdiv(K, BLOCK_K)): + a = tl.load(a_ptrs, mask=token_mask[:, None], other=0.0) + b = tl.load(b_ptrs, mask=valid_expert & n_mask[:, None], other=0.0) + asc = tl.load(as_ptrs, mask=token_mask[:, None], other=0) + bsc = tl.load(bs_ptrs, mask=valid_expert & n_mask[:, None], other=0) + acc += tl.dot_scaled(a, asc, "e4m3", b.T, bsc, "e4m3") + + a_ptrs += BLOCK_K * stride_ak + b_ptrs += BLOCK_K * stride_bk + as_ptrs += (BLOCK_K // 32) * stride_ask + bs_ptrs += (BLOCK_K // 32) * stride_bsk + + if MUL_WEIGHT: + w = tl.load(topk_weights_ptr + offs_token, mask=token_mask, other=0.0) + acc = acc * w[:, None] + + c_ptrs = c_ptr + offs_token[:, None] * stride_cm + offs_n[None, :] * stride_cn + tl.store( + c_ptrs, + acc.to(c_ptr.dtype.element_ty), + mask=token_mask[:, None] & n_mask[None, :], + ) + + +def _grouped_gemm_mxfp8( + a_q: torch.Tensor, # [M, K] fp8 e4m3 + a_scale: torch.Tensor, # [M, K//32] uint8 (E8M0) + w: torch.Tensor, # [E, N, K] fp8 e4m3 + w_scale: torch.Tensor, # [E, N, K//32] uint8 (E8M0) + sorted_token_ids: torch.Tensor, + expert_ids: torch.Tensor, + num_tokens_post_padded: torch.Tensor, + num_valid_tokens: int, + top_k: int, + block_m: int, + out_dtype: torch.dtype, + a_div: int, + mul_weight_by: Optional[torch.Tensor] = None, +) -> torch.Tensor: + M_routed = num_valid_tokens + E, N, K = w.shape + assert K % 128 == 0, f"MXFP8 native MoE requires K%128==0, got K={K}" + # Keep zero-fill: moe_align_block_size reserves an extra expert bucket for + # filtered routes, which should contribute zeros if present. + out = torch.zeros((M_routed, N), dtype=out_dtype, device=a_q.device) + if a_div == top_k and M_routed <= 32 and K >= 3072: + BLOCK_N = 64 + num_warps = 4 + else: + BLOCK_N = 128 + num_warps = 8 + BLOCK_K = 128 + grid = (triton.cdiv(sorted_token_ids.shape[0], block_m), triton.cdiv(N, BLOCK_N)) + _mxfp8_grouped_gemm_kernel[grid]( + a_q, + a_scale, + w, + w_scale, + out, + mul_weight_by if mul_weight_by is not None else a_q, + sorted_token_ids, + expert_ids, + num_tokens_post_padded, + E, + N, + K, + num_valid_tokens, + top_k, + a_q.stride(0), + a_q.stride(1), + a_scale.stride(0), + a_scale.stride(1), + w.stride(0), + w.stride(1), + w.stride(2), + w_scale.stride(0), + w_scale.stride(1), + w_scale.stride(2), + out.stride(0), + out.stride(1), + A_DIV=a_div, + MUL_WEIGHT=mul_weight_by is not None, + BLOCK_M=block_m, + BLOCK_N=BLOCK_N, + BLOCK_K=BLOCK_K, + num_warps=num_warps, + ) + return out + + +@triton.jit +def _combine_topk_routes_kernel( + route_ptr, + out_ptr, + H, + top_k: tl.constexpr, + stride_rm, + stride_rh, + stride_ot, + stride_oh, + BLOCK_H: tl.constexpr, +): + token_id = tl.program_id(0) + h_block = tl.program_id(1) + offs_h = h_block * BLOCK_H + tl.arange(0, BLOCK_H) + h_mask = offs_h < H + + acc = tl.zeros((BLOCK_H,), dtype=tl.float32) + for route_idx in range(0, top_k): + route_row = token_id * top_k + route_idx + vals = tl.load( + route_ptr + route_row * stride_rm + offs_h * stride_rh, + mask=h_mask, + other=0.0, + ) + acc += vals.to(tl.float32) + + tl.store(out_ptr + token_id * stride_ot + offs_h * stride_oh, acc, mask=h_mask) + + +def _combine_topk_routes( + route_outputs: torch.Tensor, + num_tokens: int, + top_k: int, + hidden_size: int, + out_dtype: torch.dtype, +) -> torch.Tensor: + out = torch.empty( + (num_tokens, hidden_size), dtype=out_dtype, device=route_outputs.device + ) + block_h = 1024 + grid = (num_tokens, triton.cdiv(hidden_size, block_h)) + _combine_topk_routes_kernel[grid]( + route_outputs, + out, + hidden_size, + top_k, + route_outputs.stride(0), + route_outputs.stride(1), + out.stride(0), + out.stride(1), + BLOCK_H=block_h, + num_warps=8, + ) + return out + + +def fused_moe_mxfp8_native( + hidden_states: torch.Tensor, # [T, H] bf16 + w13: torch.Tensor, # [E, 2I, H] fp8 + w13_scale: torch.Tensor, # [E, 2I, H//32] uint8 + w2: torch.Tensor, # [E, H, I] fp8 + w2_scale: torch.Tensor, # [E, H, I//32] uint8 + topk_weights: torch.Tensor, # [T, top_k] + topk_ids: torch.Tensor, # [T, top_k] (local expert ids; -1 for non-local EP) + *, + alpha: float, + beta: float, + limit: Optional[float], + no_combine: bool = False, + expert_map: Optional[torch.Tensor] = None, +) -> torch.Tensor: + # Lazy import: the jit_kernel package pulls in Triton at first use; importing + # at call time avoids any import-time cycle with the moe runner package. + from sglang.jit_kernel.minimax_m3 import swiglu_oai_mxfp8_quant, swiglu_oai_split + + T, H = hidden_states.shape + top_k = topk_ids.shape[1] + M = T * top_k + local_num_experts = w13.shape[0] + + if expert_map is not None: + valid_global = (topk_ids >= 0) & (topk_ids < expert_map.numel()) + topk_ids = expert_map[topk_ids.clamp(0, expert_map.numel() - 1).long()].to( + torch.int32 + ) + topk_ids.masked_fill_( + ~valid_global | (topk_ids < 0) | (topk_ids >= local_num_experts), -1 + ) + else: + topk_ids = topk_ids.to(torch.int32, copy=True) + topk_ids.masked_fill_((topk_ids < 0) | (topk_ids >= local_num_experts), -1) + + block_m = 64 + sorted_ids, expert_ids, num_post = moe_align_block_size( + topk_ids, block_m, local_num_experts + ) + + # GEMM1: x (mxfp8) @ w13^T -> [M, 2I]. The activation is quantized ONCE over + # the T hidden rows; the kernel gathers per route via a_row = token // top_k. + a_q, a_s = mxfp8_e4m3_quantize(hidden_states) + g1 = _grouped_gemm_mxfp8( + a_q, + a_s, + w13, + w13_scale, + sorted_ids, + expert_ids, + num_post, + M, + top_k, + block_m, + hidden_states.dtype, + a_div=top_k, + ) # [M, 2I] + + if envs.SGLANG_MINIMAX_M3_FUSED_SWIGLU_MXFP8.get(): + # SwiGLU-OAI (split layout) + MiniMax MXFP8 quant in one kernel, fp32 all + # the way to the E8M0 scale (no bf16 activation round-trip; matches the + # vLLM/ame fused kernel). Opt-in until full-model accuracy is re-qualified. + act_q, act_s = swiglu_oai_mxfp8_quant(g1, alpha=alpha, beta=beta, limit=limit) + else: + # Default accuracy path: keep the historical two-kernel boundary. + act = swiglu_oai_split( + g1, alpha=alpha, beta=beta, limit=limit, out_dtype=hidden_states.dtype + ) + act_q, act_s = mxfp8_e4m3_quantize(act) + + if no_combine: + # Per-route outputs, unweighted, no reduction: [T, top_k, H]. + g2 = _grouped_gemm_mxfp8( + act_q, + act_s, + w2, + w2_scale, + sorted_ids, + expert_ids, + num_post, + M, + top_k, + block_m, + hidden_states.dtype, + a_div=1, + ) + return g2.view(T, top_k, H) + + # GEMM2: act (mxfp8) @ w2^T -> [M, H], weighted by topk_weights, then reduce. + g2 = _grouped_gemm_mxfp8( + act_q, + act_s, + w2, + w2_scale, + sorted_ids, + expert_ids, + num_post, + M, + top_k, + block_m, + torch.float32, + a_div=1, + mul_weight_by=topk_weights.reshape(-1).to(torch.float32), + ) # [M, H] == [T*top_k, H] + + if envs.SGLANG_MINIMAX_M3_FUSED_MOE_COMBINE.get(): + return _combine_topk_routes(g2, T, top_k, H, hidden_states.dtype) + return g2.view(T, top_k, H).sum(dim=1).to(hidden_states.dtype) + + +def fused_experts_mxfp8( + hidden_states: torch.Tensor, + w1: torch.Tensor, + w2: torch.Tensor, + topk_weights: torch.Tensor, + topk_ids: torch.Tensor, + w1_scale: torch.Tensor, + w2_scale: torch.Tensor, + *, + b1: Optional[torch.Tensor] = None, + b2: Optional[torch.Tensor] = None, + activation: str = "silu", + is_gated: bool = True, + no_combine: bool = False, + inplace: bool = False, + apply_router_weight_on_input: bool = False, + routed_scaling_factor: Optional[float] = None, + gemm1_alpha: Optional[float] = None, + gemm1_limit: Optional[float] = None, + swiglu_limit: Optional[float] = None, + gate_up_interleaved: bool = True, + expert_map: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """Native MXFP8 MoE entry (CDNA4 ``dot_scaled``). + + Keeps the SGLang ``fused_experts_mxfp8`` call contract but routes to the + single grouped-GEMM kernel. Only the MiniMax-M3 SwiGLU-OAI (split, + uninterleaved, gated silu with ``gemm1_alpha``/``gemm1_limit``) configuration + is supported -- the unsupported cases below never occur on the M3 path. + """ + if not (activation == "silu" and is_gated): + raise NotImplementedError( + f"native MXFP8 MoE only supports gated swiglu-oai, got " + f"{activation=} {is_gated=}." + ) + if b1 is not None or b2 is not None: + raise NotImplementedError("native MXFP8 MoE does not support expert bias.") + if apply_router_weight_on_input: + raise NotImplementedError( + "native MXFP8 MoE does not support apply_router_weight_on_input." + ) + if gate_up_interleaved: + raise NotImplementedError( + "native MXFP8 MoE expects uninterleaved (split) gate/up layout." + ) + + # SwiGLU-OAI default activation alpha (gpt-oss); M3 may override via gemm1_alpha. + alpha = 1.702 if gemm1_alpha is None else float(gemm1_alpha) + beta = 1.0 + limit = None if gemm1_limit is None else float(gemm1_limit) + # NOTE: routed_scaling_factor is intentionally NOT re-applied here. For M3 + # (sigmoid routing with apply_routed_scaling_factor_on_output=True) it is + # already folded into topk_weights by the topk kernel; re-applying would + # double-count it. This matches the prior SGLang MXFP8 behaviour. + + out = fused_moe_mxfp8_native( + hidden_states, + w1, + w1_scale, + w2, + w2_scale, + topk_weights, + topk_ids, + alpha=alpha, + beta=beta, + limit=limit, + no_combine=no_combine, + expert_map=expert_map, + ) + + if no_combine: + return out + if inplace: + hidden_states.copy_(out) + return hidden_states + return out diff --git a/python/sglang/srt/layers/quantization/mxfp8_amd_gfx95.py b/python/sglang/srt/layers/quantization/mxfp8_amd_gfx95.py new file mode 100644 index 000000000..9c90ce8ed --- /dev/null +++ b/python/sglang/srt/layers/quantization/mxfp8_amd_gfx95.py @@ -0,0 +1,333 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Native MXFP8 (1x32 block, E8M0 scale) ops for AMD CDNA4 (gfx950). + + * per-token MXFP8 activation quant (single fused Triton pass) + * dense GEMM via Triton ``tl.dot_scaled`` (consumes FP8 E4M3 weights + E8M0 + block scales directly, no dequant-to-BF16), lowering to the CDNA4 native MX + matrix-core ops; ``K % 128 != 0`` falls back to dequant + ``F.linear``. + +Replaces the FlyDSL ``v_mfma_scale_f32_32x32x64`` dense path with a single +Triton ``dot_scaled`` GEMM: no load-time weight reformat (fp8 + E8M0 are +consumed as-is) and the activation is MXFP8-quantized in one fused pass. +""" + +from __future__ import annotations + +from typing import Optional, Tuple + +import torch +import torch.nn.functional as F +import triton +import triton.language as tl + +# MXFP8 constants (OCP microscaling: 1x32 block, E8M0 shared scale). +MXFP8_VALUE_DTYPE = torch.float8_e4m3fn +MXFP8_SCALE_DTYPE = torch.uint8 +MXFP8_BLOCK_SIZE = 32 +MXFP8_E4M3_MAX = 448.0 # max representable magnitude of float8_e4m3fn + + +# --------------------------------------------------------------------------- # +# MXFP8 quantization (per-32-block E8M0 scale + FP8-E4M3 values) +# --------------------------------------------------------------------------- # +def _mxfp8_e4m3_quantize_torch(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: + """Naive (reference) MXFP8 quantization. + + For each block of 32 elements along the last dim, compute a shared E8M0 + scale and quantize each element to float8_e4m3fn. The E8M0 exponent is + rounded *up* -- ``ceil(log2(amax / e4m3_max)) + 127`` -- so the block amax + stays inside the e4m3 range (no clipping) and the full dynamic range is + used, matching ``triton_kernels`` ``downcast_to_mxfp`` (ROUND_UP) and the + SGLang fp8 quant kernels. Returns ``(values [same shape, fp8], scales + [..., K//32] u8)``. + """ + assert x.shape[-1] % MXFP8_BLOCK_SIZE == 0 + orig_shape = x.shape + num_blocks = x.shape[-1] // MXFP8_BLOCK_SIZE + + x_fp32 = x.to(torch.float32) + x_blocked = x_fp32.view(*orig_shape[:-1], num_blocks, MXFP8_BLOCK_SIZE) + + amax = x_blocked.abs().amax(dim=-1) + amax = amax.clamp(min=torch.finfo(torch.float32).tiny) + scale_biased = (torch.ceil(torch.log2(amax / MXFP8_E4M3_MAX)) + 127.0).clamp(0, 254) + scales_uint8 = scale_biased.to(torch.uint8) + + descale = torch.exp2(scale_biased - 127.0) + x_scaled = (x_blocked / descale.unsqueeze(-1)).clamp( + -MXFP8_E4M3_MAX, MXFP8_E4M3_MAX + ) + x_fp8 = x_scaled.view(orig_shape).to(MXFP8_VALUE_DTYPE) + + scales_uint8 = scales_uint8.view(*orig_shape[:-1], num_blocks) + return x_fp8, scales_uint8 + + +@triton.jit +def _mxfp8_quant_kernel( + x_ptr, + xq_ptr, + s_ptr, + M, + K, + sxm, + sxk, + sqm, + sqk, + ssm, + ssk, + BLOCK_M: tl.constexpr, +): + """Per-32-block E8M0 scale + FP8-E4M3 quant, one program per ``[BLOCK_M, 32]``.""" + pid_m = tl.program_id(0) + pid_b = tl.program_id(1) # which 32-element block along K + offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M) + offs_k = pid_b * 32 + tl.arange(0, 32) + m_mask = offs_m < M + x = tl.load( + x_ptr + offs_m[:, None] * sxm + offs_k[None, :] * sxk, + mask=m_mask[:, None], + other=0.0, + ).to(tl.float32) + amax = tl.maximum(tl.max(tl.abs(x), axis=1), 1e-30) # [BLOCK_M] + # Round the E8M0 exponent up (ceil(log2(amax / e4m3_max))) so the block amax + # stays inside the e4m3 range and the full dynamic range is used. + sb = tl.ceil(tl.log2(amax / 448.0)) + 127.0 + sb = tl.minimum(tl.maximum(sb, 0.0), 254.0) + descale = tl.exp2(sb - 127.0) + xq = tl.clamp(x / descale[:, None], -448.0, 448.0).to(xq_ptr.dtype.element_ty) + tl.store( + xq_ptr + offs_m[:, None] * sqm + offs_k[None, :] * sqk, + xq, + mask=m_mask[:, None], + ) + tl.store(s_ptr + offs_m * ssm + pid_b * ssk, sb.to(tl.uint8), mask=m_mask) + + +def _mxfp8_e4m3_quantize_triton( + x: torch.Tensor, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Fused 2D MXFP8 quant (row-major [M, K//32] UE8M0 scales).""" + M, K = x.shape + x = x.contiguous() + xq = torch.empty((M, K), dtype=MXFP8_VALUE_DTYPE, device=x.device) + scales = torch.empty( + (M, K // MXFP8_BLOCK_SIZE), dtype=MXFP8_SCALE_DTYPE, device=x.device + ) + BLOCK_M = 64 + grid = (triton.cdiv(M, BLOCK_M), K // MXFP8_BLOCK_SIZE) + _mxfp8_quant_kernel[grid]( + x, + xq, + scales, + M, + K, + x.stride(0), + x.stride(1), + xq.stride(0), + xq.stride(1), + scales.stride(0), + scales.stride(1), + BLOCK_M=BLOCK_M, + ) + return xq, scales + + +def mxfp8_e4m3_quantize(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: + """Per-token MXFP8 quant -> (fp8 values, [.., K//32] uint8 UE8M0 scales). + + Uses the single fused Triton kernel for the common 2D, ``K % 32 == 0`` case + (activations); falls back to the torch reference otherwise. + """ + if x.ndim == 2 and x.shape[-1] % MXFP8_BLOCK_SIZE == 0 and x.is_cuda: + return _mxfp8_e4m3_quantize_triton(x.contiguous()) + return _mxfp8_e4m3_quantize_torch(x) + + +def dequant_mxfp8_to_bf16(x: torch.Tensor, scales: torch.Tensor) -> torch.Tensor: + """Dequantize an MXFP8 tensor (fp8 values + UE8M0 scales) to BF16.""" + x_float = x.to(torch.float32) + num_blocks = x.shape[-1] // MXFP8_BLOCK_SIZE + x_blocked = x_float.view(*x.shape[:-1], num_blocks, MXFP8_BLOCK_SIZE) + descale = torch.exp2(scales.to(torch.float32) - 127.0) + dequantized = (x_blocked * descale.unsqueeze(-1)).view(*x.shape) + return dequantized.to(torch.bfloat16) + + +# --------------------------------------------------------------------------- # +# Dense MXFP8 linear via Triton tl.dot_scaled (CDNA4 native microscaling) +# --------------------------------------------------------------------------- # +@triton.jit +def _mxfp8_linear_kernel( + x_ptr, + xs_ptr, + w_ptr, + ws_ptr, + out_ptr, + M, + N, + K, + stride_xm, + stride_xk, + stride_xsm, + stride_xsk, + stride_wn, + stride_wk, + stride_wsn, + stride_wsk, + stride_om, + stride_on, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, + BLOCK_K: tl.constexpr, +): + pid_m = tl.program_id(0) + pid_n = tl.program_id(1) + offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M) + offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N) + offs_k = tl.arange(0, BLOCK_K) + offs_sk = tl.arange(0, BLOCK_K // 32) + m_mask = offs_m < M + n_mask = offs_n < N + + x_ptrs = x_ptr + offs_m[:, None] * stride_xm + offs_k[None, :] * stride_xk + xs_ptrs = xs_ptr + offs_m[:, None] * stride_xsm + offs_sk[None, :] * stride_xsk + w_ptrs = w_ptr + offs_n[:, None] * stride_wn + offs_k[None, :] * stride_wk + ws_ptrs = ws_ptr + offs_n[:, None] * stride_wsn + offs_sk[None, :] * stride_wsk + + acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) + for _ in range(0, tl.cdiv(K, BLOCK_K)): + x = tl.load(x_ptrs, mask=m_mask[:, None], other=0.0) + w = tl.load(w_ptrs, mask=n_mask[:, None], other=0.0) + xs = tl.load(xs_ptrs, mask=m_mask[:, None], other=0) + ws = tl.load(ws_ptrs, mask=n_mask[:, None], other=0) + acc += tl.dot_scaled(x, xs, "e4m3", w.T, ws, "e4m3") + x_ptrs += BLOCK_K * stride_xk + w_ptrs += BLOCK_K * stride_wk + xs_ptrs += (BLOCK_K // 32) * stride_xsk + ws_ptrs += (BLOCK_K // 32) * stride_wsk + + o_ptrs = out_ptr + offs_m[:, None] * stride_om + offs_n[None, :] * stride_on + tl.store( + o_ptrs, + acc.to(out_ptr.dtype.element_ty), + mask=m_mask[:, None] & n_mask[None, :], + ) + + +def _run_mxfp8_linear_kernel( + x_q: torch.Tensor, # [M, K] fp8 e4m3 + x_scale: torch.Tensor, # [M, K//32] uint8 (E8M0) + w: torch.Tensor, # [N, K] fp8 e4m3 + w_scale: torch.Tensor, # [N, K//32] uint8 (E8M0) + out_dtype: torch.dtype, +) -> torch.Tensor: + M, K = x_q.shape + N = w.shape[0] + out = torch.empty((M, N), dtype=out_dtype, device=x_q.device) + BLOCK_M, BLOCK_K = 64, 128 + if M <= 512 and (K >= 4096 or (N == 6144 and K in (2048, 3072))): + BLOCK_N, num_warps = 64, 4 + else: + BLOCK_N, num_warps = 128, 8 + grid = (triton.cdiv(M, BLOCK_M), triton.cdiv(N, BLOCK_N)) + _mxfp8_linear_kernel[grid]( + x_q, + x_scale, + w, + w_scale, + out, + M, + N, + K, + x_q.stride(0), + x_q.stride(1), + x_scale.stride(0), + x_scale.stride(1), + w.stride(0), + w.stride(1), + w_scale.stride(0), + w_scale.stride(1), + out.stride(0), + out.stride(1), + BLOCK_M=BLOCK_M, + BLOCK_N=BLOCK_N, + BLOCK_K=BLOCK_K, + num_warps=num_warps, + ) + return out + + +def _mxfp8_dot_scaled_linear( + x: torch.Tensor, # [M, K] bf16/fp16 + w: torch.Tensor, # [N, K] fp8 e4m3 + w_scale: torch.Tensor, # [N, K//32] uint8 (E8M0) +) -> torch.Tensor: + """bf16/fp16 input -> per-token MXFP8 quant -> dot_scaled GEMM.""" + x_q, x_scale = mxfp8_e4m3_quantize(x) + return _run_mxfp8_linear_kernel(x_q, x_scale, w, w_scale, x.dtype) + + +def dot_scaled_mxfp8_blockscaled_linear( + input: torch.Tensor, + weight: torch.Tensor, + weight_scale: torch.Tensor, + input_scale: Optional[torch.Tensor] = None, + bias: Optional[torch.Tensor] = None, + output_dtype: Optional[torch.dtype] = None, +) -> torch.Tensor: + """Native dense MXFP8 linear (CDNA4 ``tl.dot_scaled``). + + Consumes FP8 E4M3 ``weight`` + canonical 2D UE8M0 ``weight_scale`` [N, K//32] + directly. Activations are MXFP8-quantized per token inside the kernel path. + Drop-in for the SGLang ``w8a8_mxfp8_linear`` callable signature. + """ + assert weight.dtype == torch.float8_e4m3fn, "MXFP8 weight must be FP8 E4M3." + assert weight_scale.dtype == torch.uint8, "MXFP8 weight_scale must be UE8M0 uint8." + assert weight_scale.dim() == 2, ( + "dot_scaled MXFP8 linear expects canonical 2D [N, K//32] weight scales, " + f"got {weight_scale.dim()}D." + ) + + input_2d = input.view(-1, input.shape[-1]).contiguous() + output_shape = [*input.shape[:-1], weight.shape[0]] + if output_dtype is None: + output_dtype = ( + input_2d.dtype + if input_2d.dtype in (torch.float16, torch.bfloat16, torch.float32) + else torch.bfloat16 + ) + + m, k = input_2d.shape + n, k_w = weight.shape + assert k == k_w, f"{k=} does not match {k_w=}" + + if k % 128 == 0: + if input_scale is None: + # Quantize the bf16/fp16 activations per token inside the path. + x_q, x_scale = mxfp8_e4m3_quantize(input_2d) + kernel_out_dtype = input_2d.dtype + else: + # Activations already MXFP8-quantized by a fused upstream op. + assert ( + input_2d.dtype == MXFP8_VALUE_DTYPE + ), "pre-quantized input must be FP8 E4M3 when input_scale is given." + assert input_scale.dtype == torch.uint8 and input_scale.shape == ( + m, + k // 32, + ), "input_scale must be UE8M0 uint8 [M, K//32]." + x_q, x_scale = input_2d, input_scale + kernel_out_dtype = output_dtype + out = _run_mxfp8_linear_kernel( + x_q, x_scale, weight, weight_scale, kernel_out_dtype + ) + else: + # dot_scaled tiling needs K % 128 == 0; dequantize fallback otherwise. + w_bf16 = dequant_mxfp8_to_bf16(weight, weight_scale) + if input_scale is not None: + input_2d = dequant_mxfp8_to_bf16(input_2d, input_scale) + out = F.linear(input_2d.to(w_bf16.dtype), w_bf16).to(output_dtype) + + if bias is not None: + out = out + bias + return out.to(output_dtype).view(*output_shape) diff --git a/test/registered/jit/benchmark/minimax/bench_minimax_decode_topk.py b/test/registered/jit/benchmark/minimax/bench_minimax_decode_topk.py new file mode 100644 index 000000000..7c3722fc8 --- /dev/null +++ b/test/registered/jit/benchmark/minimax/bench_minimax_decode_topk.py @@ -0,0 +1,121 @@ +"""Benchmark: MiniMax-M3 single-stage radix-select decode topk (JIT CUDA) vs the +2-stage split-K Triton baseline (_topk_index_partial_kernel + _topk_index_merge_kernel). + +Both consume the decode score tensor [num_heads, batch, max_seqblock] and produce +topk_idx [num_heads, batch, topk]. The JIT kernel is one launch with no +intermediate buffers; the baseline is two launches with split-K partials. +""" + +import torch +import triton + +from sglang.jit_kernel.benchmark import marker +from sglang.jit_kernel.minimax_decode_topk import minimax_decode_topk +from sglang.srt.layers.attention.minimax_sparse_ops.decode.flash_with_topk_idx import ( + _topk_index_merge_kernel, + _topk_index_partial_kernel, +) +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=8, suite="base-b-kernel-benchmark-1-gpu-large") + +BLOCK_SIZE = 128 +TOPK = 16 +NUM_HEADS = 1 # per-rank index heads at TP>=4 + + +def _triton_2stage(score, seq_lens): + num_q_heads, batch_size, max_seqblock = score.shape + TOPK_TARGET_GRID = 64 + MAX_NUM_TOPK_CHUNKS = 16 + t = max( + 1, + min(MAX_NUM_TOPK_CHUNKS, TOPK_TARGET_GRID // max(1, batch_size * num_q_heads)), + ) + nchunks = 1 << (t.bit_length() - 1) + bt = triton.next_power_of_2(TOPK) + chunk_blocks = (max_seqblock + nchunks - 1) // nchunks + out = torch.empty( + (num_q_heads, batch_size, TOPK), device=score.device, dtype=torch.int32 + ) + tsp = torch.empty( + nchunks, num_q_heads, batch_size, bt, dtype=torch.float32, device=score.device + ) + tip = torch.empty( + nchunks, num_q_heads, batch_size, bt, dtype=torch.int32, device=score.device + ) + _topk_index_partial_kernel[(batch_size, num_q_heads, nchunks)]( + score, + tsp, + tip, + seq_lens, + BLOCK_SIZE, + TOPK, + chunk_blocks, + score.stride(0), + score.stride(1), + score.stride(2), + tsp.stride(0), + tsp.stride(1), + tsp.stride(2), + tsp.stride(3), + tip.stride(0), + tip.stride(1), + tip.stride(2), + tip.stride(3), + ) + _topk_index_merge_kernel[(batch_size, num_q_heads)]( + tsp, + tip, + out, + seq_lens, + BLOCK_SIZE, + TOPK, + tsp.stride(0), + tsp.stride(1), + tsp.stride(2), + tsp.stride(3), + tip.stride(0), + tip.stride(1), + tip.stride(2), + tip.stride(3), + out.stride(0), + out.stride(1), + out.stride(2), + NUM_TOPK_CHUNKS=nchunks, + ) + return out + + +def _jit(score, seq_lens): + return minimax_decode_topk(score, seq_lens, BLOCK_SIZE, TOPK) + + +FN_MAP = {"jit": _jit, "triton_2stage": _triton_2stage} + + +@marker.parametrize("ctx", [4096, 32768, 131072, 524288], [4096, 524288]) +@marker.parametrize("batch", [1, 4, 16, 64, 256], [1, 64]) +@marker.benchmark("impl", ["jit", "triton_2stage"]) +def benchmark(ctx: int, batch: int, impl: str): + max_seqblock = (524288 + BLOCK_SIZE - 1) // BLOCK_SIZE + nb = min((ctx + BLOCK_SIZE - 1) // BLOCK_SIZE, max_seqblock) + score = torch.full( + (NUM_HEADS, batch, max_seqblock), + float("-inf"), + dtype=torch.float32, + device="cuda", + ) + score[:, :, :nb] = torch.randn(NUM_HEADS, batch, nb, device="cuda") * 5.0 + score[:, :, nb - 1] = 1e29 # forced local block + seq_lens = torch.full((batch,), ctx, device="cuda", dtype=torch.int32) + return marker.do_bench( + FN_MAP[impl], + input_args=(score, seq_lens), + graph_clone_args=(0, 1), # both read-only inputs + memory_args=(score,), + ) + + +if __name__ == "__main__": + benchmark.run() diff --git a/test/registered/jit/benchmark/minimax/bench_minimax_qknorm_rope.py b/test/registered/jit/benchmark/minimax/bench_minimax_qknorm_rope.py new file mode 100644 index 000000000..89b1ad666 --- /dev/null +++ b/test/registered/jit/benchmark/minimax/bench_minimax_qknorm_rope.py @@ -0,0 +1,146 @@ +"""Benchmark: fused MiniMax-M3 Gemma-RMSNorm + partial NeoX RoPE (1 in-place +launch) vs the unfused path (GemmaRMSNorm(q) + GemmaRMSNorm(k) + rotary_emb, +3 launches + intermediates). Main attention branch, per-rank TP8 shape (nq=8, nk=1). +""" + +import torch + +from sglang.jit_kernel.benchmark import marker +from sglang.jit_kernel.minimax_qknorm_rope import ( + minimax_qknorm_rope, + minimax_qknorm_rope_grouped, +) +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=6, suite="base-b-kernel-benchmark-1-gpu-large") + +HEAD_DIM, ROTARY_DIM, BASE, EPS, MAXPOS = 128, 64, 5_000_000, 1e-6, 131072 +NQ, NK = 64, 4 + + +def _cache(): + inv_freq = 1.0 / ( + BASE + ** ( + torch.arange(0, ROTARY_DIM, 2, dtype=torch.float, device="cuda") + / ROTARY_DIM + ) + ) + t = torch.arange(MAXPOS, dtype=torch.float, device="cuda") + freqs = torch.outer(t, inv_freq) + return torch.cat([freqs.cos(), freqs.sin()], dim=-1).contiguous() + + +def _unfused(qkv, wq, wk, cache, positions): + # GemmaRMSNorm (1+w) on q,k head-wise + partial neox rope, in plain torch + # (representative of the separate norm + rope launches). + T = qkv.shape[0] + q, k, v = qkv.split([NQ * HEAD_DIM, NK * HEAD_DIM, NK * HEAD_DIM], dim=-1) + + def norm(x, w, nh): + x = x.reshape(T, nh, HEAD_DIM).float() + var = x.pow(2).mean(-1, keepdim=True) + return (x * torch.rsqrt(var + EPS) * (1.0 + w.float())).to(torch.bfloat16) + + qn, kn = norm(q, wq, NQ), norm(k, wk, NK) + cs = cache.index_select(0, positions).float() + cos, sin = cs[:, None, :32], cs[:, None, 32:] + + def rope(x): + x1, x2 = x[..., :32].float(), x[..., 32:64].float() + o1 = x1 * cos - x2 * sin + o2 = x2 * cos + x1 * sin + return torch.cat([o1, o2, x[..., 64:].float()], dim=-1).to(torch.bfloat16) + + return rope(qn), rope(kn) + + +def _fused(qkv, wq, wk, cache, positions): + minimax_qknorm_rope(qkv, wq, wk, cache, positions, NQ, NK, NK, EPS) + return qkv + + +FN_MAP = {"fused": _fused, "unfused_torch": _unfused} + + +@marker.parametrize("T", [1, 16, 64, 256, 1024, 8192], [64, 1024]) +@marker.benchmark("impl", ["fused", "unfused_torch"]) +def benchmark(T: int, impl: str): + cache = _cache() + wq = (torch.randn(HEAD_DIM, device="cuda") * 0.1).to(torch.bfloat16) + wk = (torch.randn(HEAD_DIM, device="cuda") * 0.1).to(torch.bfloat16) + qkv = torch.randn(T, (NQ + 2 * NK) * HEAD_DIM, dtype=torch.bfloat16, device="cuda") + positions = torch.randint(0, MAXPOS, (T,), device="cuda", dtype=torch.int64) + return marker.do_bench( + FN_MAP[impl], + input_args=(qkv, wq, wk, cache, positions), + graph_clone_args=(0,), + memory_args=None, + ) + + +# --- Combined main + index single launch (the fused qkv+index_qkv GEMM path) --- +# Per-rank TP8 sparse shape: main q=8/k=1/v=1 + index idx_q=1/idx_k=1 (value +# disabled). One grouped launch (q, k, idx_q, idx_k) vs two separate launches. +C_NQ, C_NKV, C_NIQ = 8, 1, 1 +C_OFF_Q = 0 +C_OFF_K = C_NQ +C_OFF_V = C_NQ + C_NKV +C_OFF_IQ = C_NQ + 2 * C_NKV +C_OFF_IK = C_OFF_IQ + C_NIQ +C_TOTAL_HEADS = C_OFF_IK + 1 + + +def _combined_one(args): + qkv, wq, wk, wiq, wik, cache, positions = args + minimax_qknorm_rope_grouped( + qkv, + [ + (wq, C_OFF_Q, C_NQ), + (wk, C_OFF_K, C_NKV), + (wiq, C_OFF_IQ, C_NIQ), + (wik, C_OFF_IK, 1), + ], + cache, + positions, + EPS, + ) + return qkv + + +def _combined_two(args): + # Two launches over the same buffer: main (q,k) then index (idx_q, idx_k). + qkv, wq, wk, wiq, wik, cache, positions = args + minimax_qknorm_rope_grouped( + qkv, [(wq, C_OFF_Q, C_NQ), (wk, C_OFF_K, C_NKV)], cache, positions, EPS + ) + minimax_qknorm_rope_grouped( + qkv, [(wiq, C_OFF_IQ, C_NIQ), (wik, C_OFF_IK, 1)], cache, positions, EPS + ) + return qkv + + +C_FN_MAP = {"combined_one_launch": _combined_one, "two_launches": _combined_two} + + +@marker.parametrize("T", [1, 16, 64, 256, 1024, 8192], [64, 1024]) +@marker.benchmark("impl", ["combined_one_launch", "two_launches"]) +def benchmark_combined(T: int, impl: str): + cache = _cache() + ws = [ + (torch.randn(HEAD_DIM, device="cuda") * 0.1).to(torch.bfloat16) + for _ in range(4) + ] + qkv = torch.randn(T, C_TOTAL_HEADS * HEAD_DIM, dtype=torch.bfloat16, device="cuda") + positions = torch.randint(0, MAXPOS, (T,), device="cuda", dtype=torch.int64) + return marker.do_bench( + C_FN_MAP[impl], + input_args=((qkv, *ws, cache, positions),), + graph_clone_args=(0,), + memory_args=None, + ) + + +if __name__ == "__main__": + benchmark.run() + benchmark_combined.run() diff --git a/test/registered/jit/benchmark/minimax/bench_minimax_store_kv_index.py b/test/registered/jit/benchmark/minimax/bench_minimax_store_kv_index.py new file mode 100644 index 000000000..367a5b9c2 --- /dev/null +++ b/test/registered/jit/benchmark/minimax/bench_minimax_store_kv_index.py @@ -0,0 +1,67 @@ +"""Benchmark: fused MiniMax-M3 KV + index cache store (1 launch) vs the separate +per-buffer index_put_ stores (main K, main V, index K, optional index V).""" + +import torch + +from sglang.jit_kernel.benchmark import marker +from sglang.jit_kernel.minimax_store_kv_index import store_kv_index +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=6, suite="base-b-kernel-benchmark-1-gpu-large") + +HEAD_DIM = 128 +NUM_KV_HEADS = 1 +HAS_V = False +N = 1 << 20 +DTYPE = torch.bfloat16 + + +def _separate(k, v, kc, vc, ik, ikc, loc, **_): + kc[loc] = k + vc[loc] = v + ikc[loc] = ik + + +def _fused(k, v, kc, vc, ik, ikc, loc, *, num_kv_heads, head_bytes): + store_kv_index( + k, + v, + kc, + vc, + ik, + ikc, + None, + None, + loc, + num_kv_heads=num_kv_heads, + head_bytes=head_bytes, + ) + + +FN_MAP = {"fused": _fused, "separate": _separate} + + +@marker.parametrize("T", [16, 64, 256, 1024, 4096, 16384], [256, 4096]) +@marker.benchmark("impl", ["fused", "separate"]) +def benchmark(T: int, impl: str): + k = torch.randn(T, NUM_KV_HEADS * HEAD_DIM, dtype=DTYPE, device="cuda") + v = torch.randn_like(k) + ik = torch.randn(T, HEAD_DIM, dtype=DTYPE, device="cuda") + kc = torch.zeros(N, NUM_KV_HEADS * HEAD_DIM, dtype=DTYPE, device="cuda") + vc = torch.zeros_like(kc) + ikc = torch.zeros(N, HEAD_DIM, dtype=DTYPE, device="cuda") + loc = torch.randperm(N, device="cuda")[:T] + extra_kwargs = dict(num_kv_heads=NUM_KV_HEADS, head_bytes=HEAD_DIM * DTYPE.itemsize) + return marker.do_bench( + FN_MAP[impl], + input_args=(k, v, kc, vc, ik, ikc, loc), + input_kwargs=extra_kwargs if impl == "fused" else {}, + # Read inputs cloned per iter; caches are write targets (kept hot). + graph_clone_args=(0, 1, 4, 6), + memory_args=(k, v, ik, loc), + memory_output=(k, v, ik), + ) + + +if __name__ == "__main__": + benchmark.run() diff --git a/test/registered/jit/minimax/test_minimax_decode_topk.py b/test/registered/jit/minimax/test_minimax_decode_topk.py new file mode 100644 index 000000000..21e8ed3fa --- /dev/null +++ b/test/registered/jit/minimax/test_minimax_decode_topk.py @@ -0,0 +1,144 @@ +"""Correctness tests for the MiniMax-M3 single-stage radix-select decode topk. + +The kernel selects, per (head, batch) row, the indices of the ``topk`` largest +block scores among the row's first ``num_blocks = ceil(seq_len / block_size)`` +entries, front-packing valid block ids and ``-1``-padding the tail. This mirrors +the consumer ``_gqa_share_sparse_decode_kernel`` contract. +""" + +import pytest +import torch + +from sglang.jit_kernel.minimax_decode_topk import minimax_decode_topk +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=40, suite="base-b-kernel-unit-1-gpu-large") +register_cuda_ci(est_time=40, suite="base-b-kernel-unit-1-gpu-b200") + + +def _ref(score, seq_lens, block_size, topk): + H, B, S = score.shape + out = torch.full((H, B, topk), -1, dtype=torch.int32, device=score.device) + for h in range(H): + for b in range(B): + sl = int(seq_lens[b]) + nb = min((sl + block_size - 1) // block_size, S) + if nb <= topk: + for i in range(nb): + out[h, b, i] = i + continue + keff = min(topk, nb) + _, idx = torch.topk(score[h, b, :nb], keff) + out[h, b, :keff] = idx.to(torch.int32) + return out + + +def _selected_scores_sorted(score, out): + """Per-row sorted-desc multiset of the scores the kernel selected (tie-robust).""" + H, B, _ = score.shape + rows = [] + for h in range(H): + for b in range(B): + sel = out[h, b] + sel = sel[sel >= 0].long() + assert len(sel.unique()) == len(sel), f"duplicate idx h{h} b{b}: {sel}" + rows.append(torch.sort(score[h, b, sel], descending=True).values) + return rows + + +def _check_contract(out, seq_lens, block_size, topk, S): + H, B, _ = out.shape + for h in range(H): + for b in range(B): + o = out[h, b] + nvalid = int((o >= 0).sum()) + # valid entries are front-packed, -1 after + assert torch.all(o[:nvalid] >= 0) and torch.all(o[nvalid:] == -1), o + nb = min((int(seq_lens[b]) + block_size - 1) // block_size, S) + assert nvalid == min(topk, nb) + assert torch.all(o[:nvalid] < nb) + + +@pytest.mark.parametrize("dtype_sl", [torch.int32, torch.int64]) +@pytest.mark.parametrize("H", [1, 2]) +@pytest.mark.parametrize("B", [1, 5, 32]) +@pytest.mark.parametrize("topk", [16, 32, 64]) +@pytest.mark.parametrize("max_ctx", [4096, 131072, 524288]) +def test_decode_topk_distinct(dtype_sl, H, B, topk, max_ctx): + torch.manual_seed(1234) + block_size = 128 + S = (max_ctx + block_size - 1) // block_size + # distinct scores per row -> exact index-set match against torch.topk + score = torch.empty(H, B, S, dtype=torch.float32, device="cuda") + for h in range(H): + for b in range(B): + score[h, b] = torch.randperm(S, device="cuda").float() + torch.rand( + 1, device="cuda" + ) + seq_lens = torch.randint(1, max_ctx + 1, (B,), device="cuda", dtype=dtype_sl) + + out = minimax_decode_topk(score, seq_lens, block_size, topk) + ref = _ref(score, seq_lens, block_size, topk) + _check_contract(out, seq_lens, block_size, topk, S) + # exact index-set equality (distinct scores) + for h in range(H): + for b in range(B): + assert set(out[h, b][out[h, b] >= 0].tolist()) == set( + ref[h, b][ref[h, b] >= 0].tolist() + ) + + +@pytest.mark.parametrize("kind", ["ties", "negative", "neg_inf_padding", "all_equal"]) +def test_decode_topk_adversarial(kind): + torch.manual_seed(7) + block_size = 128 + H, B, S, topk = 1, 6, 1024, 16 + if kind == "ties": + score = torch.randint(0, 4, (H, B, S), device="cuda").float() + elif kind == "negative": + score = -torch.rand(H, B, S, device="cuda") * 1000 - 1.0 + elif kind == "neg_inf_padding": + score = torch.randn(H, B, S, device="cuda") + score[:, :, ::7] = float("-inf") # scattered -inf in valid range + else: # all_equal + score = torch.full((H, B, S), 3.14, dtype=torch.float32, device="cuda") + seq_lens = torch.randint(1, S * block_size, (B,), device="cuda", dtype=torch.int32) + + out = minimax_decode_topk(score, seq_lens, block_size, topk) + ref = _ref(score, seq_lens, block_size, topk) + _check_contract(out, seq_lens, block_size, topk, S) + # tie-robust: the multiset of selected scores must match torch.topk's + for a, b in zip( + _selected_scores_sorted(score, out), _selected_scores_sorted(score, ref) + ): + torch.testing.assert_close(a, b, rtol=0, atol=0) + + +@pytest.mark.parametrize("seq_len", [1, 128, 129, 2048, 2049]) +def test_decode_topk_small_num_blocks(seq_len): + # num_blocks around / below topk -> naive identity path and boundary. + block_size = 128 + H, B, S, topk = 1, 1, 64, 16 + score = torch.randn(H, B, S, dtype=torch.float32, device="cuda") + seq_lens = torch.tensor([seq_len], device="cuda", dtype=torch.int32) + out = minimax_decode_topk(score, seq_lens, block_size, topk) + _check_contract(out, seq_lens, block_size, topk, S) + nb = min((seq_len + block_size - 1) // block_size, S) + if nb <= topk: + assert out[0, 0, :nb].tolist() == list(range(nb)) + + +def test_decode_topk_out_param(): + block_size = 128 + H, B, S, topk = 1, 4, 1024, 16 + score = torch.randn(H, B, S, dtype=torch.float32, device="cuda") + seq_lens = torch.full((B,), 100000, device="cuda", dtype=torch.int32) + out = torch.empty((H, B, topk), dtype=torch.int32, device="cuda") + res = minimax_decode_topk(score, seq_lens, block_size, topk, out=out) + assert res is out + + +if __name__ == "__main__": + import sys + + sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/test/registered/jit/minimax/test_minimax_decode_topk_page_table.py b/test/registered/jit/minimax/test_minimax_decode_topk_page_table.py new file mode 100644 index 000000000..c4d55cb13 --- /dev/null +++ b/test/registered/jit/minimax/test_minimax_decode_topk_page_table.py @@ -0,0 +1,178 @@ +"""Fused decode top-k + page-table transform. + +`minimax_decode_topk_page_table` selects the top-k blocks (same as the block-id +`minimax_decode_topk`) and emits the per-query paged page table consumed by the +dense backend (trtllm_mha), instead of block ids. This checks the fused output +end-to-end: trtllm decode over the emitted page table matches the custom +`_gqa_share_sparse_decode_kernel` fed the block-id selection from the same score. +Only the TP>=4 case (num_kv_heads == 1) is covered. +""" + +import random +import sys + +import pytest +import torch + +flashinfer = pytest.importorskip("flashinfer") + +from sglang.jit_kernel.minimax_decode_topk import ( + minimax_decode_topk, + minimax_decode_topk_page_table, +) +from sglang.srt.layers.attention.minimax_sparse_ops.decode.topk_sparse import ( + flash_decode_with_gqa_share_sparse, +) +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=25, suite="base-b-kernel-unit-1-gpu-b200") + +dev = "cuda" + + +def _effective_kv_from_selection(ti, seq_lens, block): + # Ground-truth effective KV length: sum of valid tokens over the selected + # blocks (only the final block can be partial), per query. + bs = seq_lens.shape[0] + out = torch.zeros(bs, dtype=torch.int32, device=ti.device) + for b in range(bs): + sl = int(seq_lens[b]) + tot = 0 + for c in ti[0, b].tolist(): + if c < 0: + continue + tot += min(block, sl - c * block) + out[b] = tot + return out + + +@pytest.mark.parametrize( + "bs,seq_len", + [ + (1, 5000), + (2, 300), + (3, 160), + (4, 2048), + (8, 4096), + (16, 8000), + (2, 40000), # num_blocks=313 -> medium radix path + (1, 90000), # num_blocks=704 -> large compaction path + (1, 480000), # num_blocks=3750 -> large path near kMaxNumBlocks + ], +) +@pytest.mark.parametrize("nqh", [8, 16]) +def test_fused_page_table_matches_custom(bs, seq_len, nqh): + arch_major, _ = torch.cuda.get_device_capability(torch.cuda.current_device()) + if arch_major < 10: + pytest.skip("trtllm-gen decode is Blackwell (sm100)") + nkv, D, block, topk, ps = 1, 128, 128, 16, 64 + torch.manual_seed(bs * 31 + seq_len + nqh) + random.seed(bs * 31 + seq_len) + + ppr = (seq_len + ps - 1) // ps + npages = bs * ppr + 4 + kf = torch.randn(npages * ps, nkv, D, device=dev, dtype=torch.bfloat16) * 0.5 + vf = torch.randn(npages * ps, nkv, D, device=dev, dtype=torch.bfloat16) * 0.5 + q = torch.randn(bs, nqh, D, device=dev, dtype=torch.bfloat16) * 0.5 + + r2t = torch.zeros(bs, seq_len, dtype=torch.int32, device=dev) + sl = torch.full((bs,), seq_len, dtype=torch.int32, device=dev) + sid = torch.arange(bs, dtype=torch.int64, device=dev) + idxs = torch.arange(seq_len, device=dev) + for b in range(bs): + r2t[b] = ((b * ppr + idxs // ps) * ps + idxs % ps).int() + + nb = (seq_len + block - 1) // block + score = torch.full((1, bs, nb), -float("inf"), device=dev, dtype=torch.float32) + score[0, :, :nb] = torch.randn(bs, nb, device=dev) + score[0, :, nb - 1] = 1e30 # final (local) block always selected + + # block-id selection -> custom kernel (reference) + ti = minimax_decode_topk(score, sl, block, topk) + ref = flash_decode_with_gqa_share_sparse( + q, None, kf, vf, r2t, sl, sid, block, ti, sm_scale=D**-0.5 + ) + + # fused page-table + effective KV length -> trtllm (allocated + returned) + pt, cache = minimax_decode_topk_page_table(score, sl, r2t, sid, block, topk, ps) + # the kernel's effective KV length must match the actual block selection + expect_cache = _effective_kv_from_selection(ti, sl, block) + assert torch.equal(cache, expect_cache), f"{cache} != {expect_cache}" + ws = torch.zeros(128 * 1024 * 1024, dtype=torch.int8, device=dev) + kv = ( + kf.view(npages, ps, nkv, D).permute(0, 2, 1, 3), + vf.view(npages, ps, nkv, D).permute(0, 2, 1, 3), + ) + o = flashinfer.decode.trtllm_batch_decode_with_kv_cache( + query=q, + kv_cache=kv, + workspace_buffer=ws, + block_tables=pt, + seq_lens=cache, + max_seq_len=topk * block, + bmm1_scale=D**-0.5, + bmm2_scale=1.0, + ) + cos = torch.nn.functional.cosine_similarity( + ref.float().flatten(), o.float().flatten(), dim=0 + ).item() + assert cos > 0.999, f"cos={cos}" + + +@pytest.mark.parametrize("seq_len", [300, 5000, 90000]) +@pytest.mark.parametrize("bs", [1, 3]) +@pytest.mark.parametrize("nkv", [2, 4]) +def test_dp_flattened_page_table(nkv, bs, seq_len): + """DP attention (num_kv_heads>1): each kv head selects its own blocks, flattened + into bs*nkv pseudo-requests (row = b*nkv + h). Validate the flattened page table + + effective KV length against the per-head block-id selection, including the + head-minor head-encoded page index (base_page*nkv + h, the index into an HND + cache [num_pages, nkv, ps, D] reshaped to [num_pages*nkv, 1, ps, D]).""" + D, block, topk, ps = 128, 128, 16, 64 + ppb = block // ps + torch.manual_seed(nkv * 131 + bs * 31 + seq_len) + ppr = (seq_len + ps - 1) // ps + max_kv = seq_len # req_to_token width + + r2t = torch.zeros(bs, seq_len, dtype=torch.int32, device=dev) + sl = torch.full((bs,), seq_len, dtype=torch.int32, device=dev) + sid = torch.arange(bs, dtype=torch.int64, device=dev) + idxs = torch.arange(seq_len, device=dev) + for b in range(bs): + r2t[b] = ((b * ppr + idxs // ps) * ps + idxs % ps).int() + + nb = (seq_len + block - 1) // block + score = torch.full((nkv, bs, nb), -float("inf"), device=dev, dtype=torch.float32) + score[:, :, :nb] = torch.randn(nkv, bs, nb, device=dev) + score[:, :, nb - 1] = 1e30 # final (local) block always selected + + # per-head block-id selection is the reference for the flattened page table + ti = minimax_decode_topk(score, sl, block, topk) # [nkv, bs, topk] + pt, cache = minimax_decode_topk_page_table(score, sl, r2t, sid, block, topk, ps) + msp = topk * ppb + assert pt.shape == (bs * nkv, msp) and cache.shape == (bs * nkv,) + + r2t_cpu = r2t.cpu() + for b in range(bs): + for h in range(nkv): + blocks = sorted(c for c in ti[h, b].tolist() if c >= 0) + row = b * nkv + h + # effective KV length = sum of valid tokens over selected blocks + exp_kv = sum(min(block, seq_len - c * block) for c in blocks) + assert ( + int(cache[row]) == exp_kv + ), f"row {row}: {int(cache[row])} != {exp_kv}" + # page table: each block -> ppb pages via req_to_token, head-minor encoded + for e in range(len(blocks) * ppb): + c = blocks[e // ppb] + tok = c * block + (e % ppb) * ps + if tok >= max_kv: + tok = max_kv - 1 + exp = int(r2t_cpu[b, tok]) // ps * nkv + h + assert ( + int(pt[row, e]) == exp + ), f"row {row} e {e}: {int(pt[row,e])} != {exp}" + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/jit/minimax/test_minimax_qknorm_rope.py b/test/registered/jit/minimax/test_minimax_qknorm_rope.py new file mode 100644 index 000000000..91d653a6f --- /dev/null +++ b/test/registered/jit/minimax/test_minimax_qknorm_rope.py @@ -0,0 +1,178 @@ +"""Correctness for the fused MiniMax-M3 Gemma-RMSNorm + partial NeoX RoPE kernel. + +Verifies the in-place fused kernel reproduces GemmaRMSNorm((1+w)) + partial NeoX +RoPE to the bf16 round-off floor, leaves V untouched, and matches sglang's RoPE +convention (cos|sin cache, neox pairs (i, i+rotary_dim/2)). +""" + +import pytest +import torch + +from sglang.jit_kernel.minimax_qknorm_rope import ( + minimax_qknorm_rope, + minimax_qknorm_rope_grouped, +) +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=30, suite="base-b-kernel-unit-1-gpu-large") +register_cuda_ci(est_time=30, suite="base-b-kernel-unit-1-gpu-b200") + +dev = "cuda" +HEAD_DIM, ROTARY_DIM, BASE, EPS = 128, 64, 5_000_000, 1e-6 + + +def _build_cache(max_pos): + inv_freq = 1.0 / ( + BASE + ** (torch.arange(0, ROTARY_DIM, 2, dtype=torch.float, device=dev) / ROTARY_DIM) + ) # [32] + t = torch.arange(max_pos, dtype=torch.float, device=dev) + freqs = torch.outer(t, inv_freq) # [max_pos, 32] + return torch.cat([freqs.cos(), freqs.sin()], dim=-1).contiguous() # [max_pos, 64] + + +def _ref(q, k, wq, wk, cache, positions, nq, nk): + T = q.shape[0] + + def norm(x, w, nh): + x = x.reshape(T, nh, HEAD_DIM).float() + var = x.pow(2).mean(-1, keepdim=True) + return x * torch.rsqrt(var + EPS) * (1.0 + w.float()) + + cs = cache.index_select(0, positions).float() + cos, sin = cs[:, :32], cs[:, 32:] + + def rope(x): + x1, x2 = x[..., :32], x[..., 32:64] + o1 = x1 * cos[:, None, :] - x2 * sin[:, None, :] + o2 = x2 * cos[:, None, :] + x1 * sin[:, None, :] + return torch.cat([o1, o2, x[..., 64:]], dim=-1) + + return rope(norm(q, wq, nq)).reshape(T, -1), rope(norm(k, wk, nk)).reshape(T, -1) + + +@pytest.mark.parametrize("nq,nk", [(8, 1), (64, 8), (16, 2)]) +@pytest.mark.parametrize("T", [1, 7, 64, 1024, 4096]) +@pytest.mark.parametrize("pos_dtype", [torch.int32, torch.int64]) +def test_fused_qknorm_rope(nq, nk, T, pos_dtype): + torch.manual_seed(T * 131 + nq) + max_pos = 8192 + cache = _build_cache(max_pos) + wq = (torch.randn(HEAD_DIM, device=dev) * 0.1).to(torch.bfloat16) + wk = (torch.randn(HEAD_DIM, device=dev) * 0.1).to(torch.bfloat16) + q = torch.randn(T, nq * HEAD_DIM, dtype=torch.bfloat16, device=dev) + k = torch.randn(T, nk * HEAD_DIM, dtype=torch.bfloat16, device=dev) + v = torch.randn(T, nk * HEAD_DIM, dtype=torch.bfloat16, device=dev) + positions = torch.randint(0, max_pos, (T,), device=dev, dtype=pos_dtype) + + qr, kr = _ref(q, k, wq, wk, cache, positions.long(), nq, nk) + + qkv = torch.cat([q, k, v], dim=-1).contiguous() + minimax_qknorm_rope(qkv, wq, wk, cache, positions, nq, nk, nk, EPS) + qf, kf, vf = qkv.split([nq * HEAD_DIM, nk * HEAD_DIM, nk * HEAD_DIM], dim=-1) + + floor = (qr.bfloat16().float() - qr.float()).abs().max().item() + assert (qf.float() - qr.float()).abs().max().item() <= 2 * floor + 1e-3 + assert (kf.float() - kr.float()).abs().max().item() <= 2 * floor + 1e-3 + assert (vf.float() - v.float()).abs().max().item() == 0.0 # V untouched + + +def test_index_branch_shapes(): + # idx_q (nq=num_idx_heads, nk=0) and idx_k (nq=1) in-place calls. + torch.manual_seed(0) + max_pos = 4096 + cache = _build_cache(max_pos) + w = (torch.randn(HEAD_DIM, device=dev) * 0.1).to(torch.bfloat16) + T, num_idx = 33, 4 + positions = torch.randint(0, max_pos, (T,), device=dev, dtype=torch.int64) + for nq in (num_idx, 1): + x = torch.randn(T, nq * HEAD_DIM, dtype=torch.bfloat16, device=dev) + ref, _ = _ref(x, x[:, :HEAD_DIM], w, w, cache, positions, nq, 1) + xc = x.clone() + minimax_qknorm_rope(xc, w, w, cache, positions, nq, 0, 0, EPS) + floor = (ref.bfloat16().float() - ref.float()).abs().max().item() + assert (xc.float() - ref.float()).abs().max().item() <= 2 * floor + 1e-3 + + +def _norm_rope_one(x_heads, w, cache, positions): + # x_heads: [T, nh, HEAD_DIM] fp32; returns same shape, GemmaRMSNorm(1+w)+rope. + var = x_heads.pow(2).mean(-1, keepdim=True) + y = x_heads * torch.rsqrt(var + EPS) * (1.0 + w.float()) + cs = cache.index_select(0, positions).float() + cos, sin = cs[:, None, :32], cs[:, None, 32:] + x1, x2 = y[..., :32], y[..., 32:64] + o1 = x1 * cos - x2 * sin + o2 = x2 * cos + x1 * sin + return torch.cat([o1, o2, y[..., 64:]], dim=-1) + + +@pytest.mark.parametrize("nq,nkv,niq", [(8, 1, 1), (8, 1, 4), (16, 2, 2)]) +@pytest.mark.parametrize("idx_v", [0, 1]) +@pytest.mark.parametrize("T", [1, 7, 64, 1024]) +@pytest.mark.parametrize("pos_dtype", [torch.int32, torch.int64]) +def test_combined_main_index_grouped(nq, nkv, niq, idx_v, T, pos_dtype): + """Combined main(q,k,v) + index(idx_q,idx_k,[idx_v]) layout in one launch. + + Mirrors the fused qkv+index_qkv GEMM output: a uniform [total_heads, 128] + grid where Q / K main heads and index-Q / index-K heads are normed+roped in + one pass and the V / index-V heads are left untouched. + """ + torch.manual_seed(T * 17 + nq * 3 + niq + idx_v) + max_pos = 8192 + cache = _build_cache(max_pos) + positions = torch.randint(0, max_pos, (T,), device=dev, dtype=pos_dtype) + + # head layout: q(nq) k(nkv) v(nkv) idx_q(niq) idx_k(1) [idx_v(1)] + off_q = 0 + off_k = nq + off_v = nq + nkv + off_iq = nq + 2 * nkv + off_ik = off_iq + niq + total_heads = off_ik + 1 + idx_v + + wq = (torch.randn(HEAD_DIM, device=dev) * 0.1).to(torch.bfloat16) + wk = (torch.randn(HEAD_DIM, device=dev) * 0.1).to(torch.bfloat16) + wiq = (torch.randn(HEAD_DIM, device=dev) * 0.1).to(torch.bfloat16) + wik = (torch.randn(HEAD_DIM, device=dev) * 0.1).to(torch.bfloat16) + + x = torch.randn(T, total_heads, HEAD_DIM, dtype=torch.bfloat16, device=dev) + ref = x.float().clone() + ref[:, off_q:off_k] = _norm_rope_one( + ref[:, off_q:off_k], wq, cache, positions.long() + ) + ref[:, off_k:off_v] = _norm_rope_one( + ref[:, off_k:off_v], wk, cache, positions.long() + ) + ref[:, off_iq:off_ik] = _norm_rope_one( + ref[:, off_iq:off_ik], wiq, cache, positions.long() + ) + ref[:, off_ik : off_ik + 1] = _norm_rope_one( + ref[:, off_ik : off_ik + 1], wik, cache, positions.long() + ) + + qkv = x.reshape(T, total_heads * HEAD_DIM).contiguous() + minimax_qknorm_rope_grouped( + qkv, + [(wq, off_q, nq), (wk, off_k, nkv), (wiq, off_iq, niq), (wik, off_ik, 1)], + cache, + positions, + EPS, + ) + out = qkv.reshape(T, total_heads, HEAD_DIM) + + floor = (ref.bfloat16().float() - ref).abs().max().item() + assert (out.float() - ref).abs().max().item() <= 2 * floor + 1e-3 + # V and index-V heads untouched (bit-exact). + assert ( + out[:, off_v:off_iq].float() - x[:, off_v:off_iq].float() + ).abs().max() == 0.0 + if idx_v: + assert ( + out[:, off_ik + 1 :].float() - x[:, off_ik + 1 :].float() + ).abs().max() == 0.0 + + +if __name__ == "__main__": + import sys + + sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/test/registered/jit/minimax/test_minimax_store_kv_index.py b/test/registered/jit/minimax/test_minimax_store_kv_index.py new file mode 100644 index 000000000..8f74cbeac --- /dev/null +++ b/test/registered/jit/minimax/test_minimax_store_kv_index.py @@ -0,0 +1,78 @@ +"""Correctness for the fused MiniMax-M3 KV + index cache store kernel. + +Verifies the single fused launch writes the main K/V, the index K, and the +optional index V into their pools at out_cache_loc rows exactly as the separate +index_put_ stores would, for both value modes and int32/int64 indices. +""" + +import pytest +import torch + +from sglang.jit_kernel.minimax_store_kv_index import store_kv_index +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=30, suite="base-b-kernel-unit-1-gpu-large") +register_cuda_ci(est_time=30, suite="base-b-kernel-unit-1-gpu-b200") + +dev = "cuda" +HEAD_DIM = 128 + + +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +@pytest.mark.parametrize("num_kv_heads", [1, 4, 8]) +@pytest.mark.parametrize("has_v", [False, True]) +@pytest.mark.parametrize("idx_dtype", [torch.int32, torch.int64]) +@pytest.mark.parametrize("T", [1, 7, 128, 513]) +def test_store_kv_index(dtype, num_kv_heads, has_v, idx_dtype, T): + torch.manual_seed(T * 17 + num_kv_heads) + head_bytes = HEAD_DIM * dtype.itemsize + N = 4096 + + def rnd(*shape): + return torch.randn(*shape, dtype=dtype, device=dev) + + k = rnd(T, num_kv_heads * HEAD_DIM) + v = rnd(T, num_kv_heads * HEAD_DIM) + idx_k = rnd(T, HEAD_DIM) + idx_v = rnd(T, HEAD_DIM) if has_v else None + + k_cache = torch.zeros(N, num_kv_heads * HEAD_DIM, dtype=dtype, device=dev) + v_cache = torch.zeros_like(k_cache) + idx_k_cache = torch.zeros(N, HEAD_DIM, dtype=dtype, device=dev) + idx_v_cache = torch.zeros_like(idx_k_cache) if has_v else None + + loc = torch.randperm(N, device=dev)[:T].to(idx_dtype) + + store_kv_index( + k, + v, + k_cache, + v_cache, + idx_k, + idx_k_cache, + idx_v, + idx_v_cache, + loc, + num_kv_heads=num_kv_heads, + head_bytes=head_bytes, + ) + + ll = loc.long() + k_ref = torch.zeros_like(k_cache) + v_ref = torch.zeros_like(v_cache) + ik_ref = torch.zeros_like(idx_k_cache) + k_ref[ll], v_ref[ll], ik_ref[ll] = k, v, idx_k + + assert torch.equal(k_cache, k_ref) + assert torch.equal(v_cache, v_ref) + assert torch.equal(idx_k_cache, ik_ref) + if has_v: + iv_ref = torch.zeros_like(idx_v_cache) + iv_ref[ll] = idx_v + assert torch.equal(idx_v_cache, iv_ref) + + +if __name__ == "__main__": + import sys + + sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/test/registered/jit/test_minimax_m3_qk_norm_rope.py b/test/registered/jit/test_minimax_m3_qk_norm_rope.py new file mode 100644 index 000000000..2026cc313 --- /dev/null +++ b/test/registered/jit/test_minimax_m3_qk_norm_rope.py @@ -0,0 +1,334 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Reference tests for MiniMax-M3 fused Q/K Gemma RMSNorm + RoPE.""" + +import pytest +import torch + +from sglang.srt.utils import is_hip + +if not is_hip(): + pytest.skip( + "MiniMax-M3 fused Q/K norm + RoPE kernel is ROCm-only.", + allow_module_level=True, + ) +if not torch.cuda.is_available(): + pytest.skip("Requires a GPU.", allow_module_level=True) + +from sglang.jit_kernel.minimax_m3.qk_norm_rope import ( # noqa: E402 + qk_gemma_rmsnorm_rope, + sparse_qk_index_gemma_rmsnorm_rope, + sparse_qk_index_gemma_rmsnorm_rope_cache, +) +from sglang.test.ci.ci_register import register_amd_ci # noqa: E402 + +# ROCm-only fused kernel; runs in the AMD jit-kernel unit suite. +register_amd_ci(est_time=30, suite="jit-kernel-unit-test-amd") + +DEVICE = "cuda" +EPS = 1e-6 + + +def _gemma_norm_by_head(x: torch.Tensor, weight: torch.Tensor, head_dim: int): + orig_shape = x.shape + orig_dtype = x.dtype + xh = x.view(x.shape[0], -1, head_dim).float() + var = xh.pow(2).mean(dim=-1, keepdim=True) + out = xh * torch.rsqrt(var + EPS) * (1.0 + weight.float()) + return out.to(orig_dtype).reshape(orig_shape) + + +def _apply_rope_ref( + x: torch.Tensor, + positions: torch.Tensor, + cos_sin_cache: torch.Tensor, + head_dim: int, + rotary_dim: int, + is_neox_style: bool, +): + orig_shape = x.shape + xh = x.view(x.shape[0], -1, head_dim) + x_rot = xh[..., :rotary_dim].float() + x_pass = xh[..., rotary_dim:] + cos_sin = cos_sin_cache.index_select(0, positions) + cos, sin = cos_sin.chunk(2, dim=-1) + cos = cos[:, None, :].float() + sin = sin[:, None, :].float() + + if is_neox_style: + x1, x2 = x_rot.chunk(2, dim=-1) + y_rot = torch.cat((x1 * cos - x2 * sin, x2 * cos + x1 * sin), dim=-1) + else: + x1 = x_rot[..., ::2] + x2 = x_rot[..., 1::2] + y_rot = torch.stack((x1 * cos - x2 * sin, x2 * cos + x1 * sin), dim=-1) + y_rot = y_rot.flatten(-2) + + return torch.cat((y_rot.to(x.dtype), x_pass), dim=-1).reshape(orig_shape) + + +def _reference( + q, + k, + q_weight, + k_weight, + positions, + cos_sin_cache, + head_dim, + rotary_dim, + is_neox_style, +): + q_norm = _gemma_norm_by_head(q, q_weight, head_dim) + k_norm = _gemma_norm_by_head(k, k_weight, head_dim) + q_ref = _apply_rope_ref( + q_norm, positions, cos_sin_cache, head_dim, rotary_dim, is_neox_style + ) + k_ref = _apply_rope_ref( + k_norm, positions, cos_sin_cache, head_dim, rotary_dim, is_neox_style + ) + return q_ref, k_ref + + +def _sparse_reference( + q, + k, + idx_q, + idx_k, + q_weight, + k_weight, + idx_q_weight, + idx_k_weight, + positions, + cos_sin_cache, + head_dim, + rotary_dim, + is_neox_style, +): + q_ref, k_ref = _reference( + q, + k, + q_weight, + k_weight, + positions, + cos_sin_cache, + head_dim, + rotary_dim, + is_neox_style, + ) + idx_q_norm = _gemma_norm_by_head(idx_q, idx_q_weight, head_dim) + idx_k_norm = _gemma_norm_by_head(idx_k, idx_k_weight, head_dim) + idx_q_ref = _apply_rope_ref( + idx_q_norm, positions, cos_sin_cache, head_dim, rotary_dim, is_neox_style + ) + idx_k_ref = _apply_rope_ref( + idx_k_norm, positions, cos_sin_cache, head_dim, rotary_dim, is_neox_style + ) + return q_ref, k_ref, idx_q_ref, idx_k_ref + + +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +@pytest.mark.parametrize("is_neox_style", [True, False]) +@pytest.mark.parametrize( + "num_tokens,q_heads,k_heads,head_dim,rotary_dim", + [(1, 16, 1, 128, 64), (17, 16, 1, 128, 64), (64, 4, 1, 128, 64)], +) +@torch.inference_mode() +def test_qk_gemma_rmsnorm_rope_matches_reference( + dtype, is_neox_style, num_tokens, q_heads, k_heads, head_dim, rotary_dim +): + torch.manual_seed(0) + q_dim = q_heads * head_dim + k_dim = k_heads * head_dim + padding_dim = 37 + qkv = torch.randn( + num_tokens, q_dim + k_dim + padding_dim, device=DEVICE, dtype=dtype + ) + q, k, _ = qkv.split([q_dim, k_dim, padding_dim], dim=-1) + if num_tokens > 1: + assert not q.is_contiguous() + assert not k.is_contiguous() + + q_weight = torch.randn(head_dim, device=DEVICE, dtype=torch.float32) + k_weight = torch.randn(head_dim, device=DEVICE, dtype=torch.float32) + positions = torch.randint(0, 512, (num_tokens,), device=DEVICE, dtype=torch.long) + cos_sin_cache = torch.randn(512, rotary_dim, device=DEVICE, dtype=dtype) + + got_q, got_k = qk_gemma_rmsnorm_rope( + q, + k, + q_weight, + k_weight, + positions, + cos_sin_cache, + EPS, + head_dim, + rotary_dim, + is_neox_style, + ) + ref_q, ref_k = _reference( + q, + k, + q_weight, + k_weight, + positions, + cos_sin_cache, + head_dim, + rotary_dim, + is_neox_style, + ) + + torch.testing.assert_close(got_q, ref_q, atol=3e-2, rtol=3e-2) + torch.testing.assert_close(got_k, ref_k, atol=3e-2, rtol=3e-2) + + +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +@pytest.mark.parametrize("is_neox_style", [True, False]) +@pytest.mark.parametrize( + "num_tokens,q_heads,k_heads,idx_q_heads,head_dim,rotary_dim", + [(1, 16, 1, 16, 128, 64), (19, 16, 1, 16, 128, 64)], +) +@torch.inference_mode() +def test_sparse_qk_index_gemma_rmsnorm_rope_matches_reference( + dtype, + is_neox_style, + num_tokens, + q_heads, + k_heads, + idx_q_heads, + head_dim, + rotary_dim, +): + torch.manual_seed(1) + q = torch.randn(num_tokens, q_heads * head_dim, device=DEVICE, dtype=dtype) + k = torch.randn(num_tokens, k_heads * head_dim, device=DEVICE, dtype=dtype) + idx_q = torch.randn(num_tokens, idx_q_heads * head_dim, device=DEVICE, dtype=dtype) + idx_k = torch.randn(num_tokens, head_dim, device=DEVICE, dtype=dtype) + q_weight = torch.randn(head_dim, device=DEVICE, dtype=torch.float32) + k_weight = torch.randn(head_dim, device=DEVICE, dtype=torch.float32) + idx_q_weight = torch.randn(head_dim, device=DEVICE, dtype=torch.float32) + idx_k_weight = torch.randn(head_dim, device=DEVICE, dtype=torch.float32) + positions = torch.randint(0, 512, (num_tokens,), device=DEVICE, dtype=torch.long) + cos_sin_cache = torch.randn(512, rotary_dim, device=DEVICE, dtype=dtype) + + got = sparse_qk_index_gemma_rmsnorm_rope( + q, + k, + idx_q, + idx_k, + q_weight, + k_weight, + idx_q_weight, + idx_k_weight, + positions, + cos_sin_cache, + EPS, + head_dim, + rotary_dim, + is_neox_style, + ) + ref = _sparse_reference( + q, + k, + idx_q, + idx_k, + q_weight, + k_weight, + idx_q_weight, + idx_k_weight, + positions, + cos_sin_cache, + head_dim, + rotary_dim, + is_neox_style, + ) + for got_tensor, ref_tensor in zip(got, ref): + torch.testing.assert_close(got_tensor, ref_tensor, atol=3e-2, rtol=3e-2) + + +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +@pytest.mark.parametrize("is_neox_style", [True, False]) +@torch.inference_mode() +def test_sparse_qk_index_gemma_rmsnorm_rope_cache_matches_reference( + dtype, is_neox_style +): + torch.manual_seed(2) + num_tokens, q_heads, k_heads, idx_q_heads = 11, 16, 1, 16 + head_dim, rotary_dim = 128, 64 + q = torch.randn(num_tokens, q_heads * head_dim, device=DEVICE, dtype=dtype) + k = torch.randn(num_tokens, k_heads * head_dim, device=DEVICE, dtype=dtype) + v = torch.randn(num_tokens, k_heads * head_dim, device=DEVICE, dtype=dtype) + idx_q = torch.randn(num_tokens, idx_q_heads * head_dim, device=DEVICE, dtype=dtype) + idx_k = torch.randn(num_tokens, head_dim, device=DEVICE, dtype=dtype) + q_weight = torch.randn(head_dim, device=DEVICE, dtype=torch.float32) + k_weight = torch.randn(head_dim, device=DEVICE, dtype=torch.float32) + idx_q_weight = torch.randn(head_dim, device=DEVICE, dtype=torch.float32) + idx_k_weight = torch.randn(head_dim, device=DEVICE, dtype=torch.float32) + positions = torch.randint(0, 512, (num_tokens,), device=DEVICE, dtype=torch.long) + cos_sin_cache = torch.randn(512, rotary_dim, device=DEVICE, dtype=dtype) + out_cache_loc = torch.randperm(64, device=DEVICE, dtype=torch.int64)[:num_tokens] + k_cache = torch.empty(64, k_heads, head_dim, device=DEVICE, dtype=dtype) + v_cache = torch.empty(64, k_heads, head_dim, device=DEVICE, dtype=dtype) + idx_k_cache = torch.empty(64, 1, head_dim, device=DEVICE, dtype=dtype) + + got = sparse_qk_index_gemma_rmsnorm_rope_cache( + q, + k, + v, + idx_q, + idx_k, + k_cache, + v_cache, + idx_k_cache, + out_cache_loc, + q_weight, + k_weight, + idx_q_weight, + idx_k_weight, + positions, + cos_sin_cache, + EPS, + head_dim, + rotary_dim, + is_neox_style, + ) + ref = _sparse_reference( + q, + k, + idx_q, + idx_k, + q_weight, + k_weight, + idx_q_weight, + idx_k_weight, + positions, + cos_sin_cache, + head_dim, + rotary_dim, + is_neox_style, + ) + for got_tensor, ref_tensor in zip(got, ref): + torch.testing.assert_close(got_tensor, ref_tensor, atol=3e-2, rtol=3e-2) + + torch.testing.assert_close( + k_cache.index_select(0, out_cache_loc), + ref[1].view(num_tokens, k_heads, head_dim), + atol=3e-2, + rtol=3e-2, + ) + torch.testing.assert_close( + v_cache.index_select(0, out_cache_loc), + v.view(num_tokens, k_heads, head_dim), + atol=0, + rtol=0, + ) + torch.testing.assert_close( + idx_k_cache.index_select(0, out_cache_loc), + ref[3].view(num_tokens, 1, head_dim), + atol=3e-2, + rtol=3e-2, + ) + + +if __name__ == "__main__": + import sys + + sys.exit(pytest.main([__file__, "-v", "-s"]))