From db30a63a1395114b92c544fa6a857c0a1f77bdfc Mon Sep 17 00:00:00 2001 From: Kurt Shuster Date: Wed, 8 Apr 2026 14:45:13 -0400 Subject: [PATCH] [sgl-kernel] support > 1024 experts in moe_align_block_size kernel (#21610) --- .../jit_kernel/csrc/moe/moe_align_kernel.cu | 580 ++++++++++++++++++ python/sglang/jit_kernel/moe_align.py | 46 ++ .../tests/test_moe_align_block_size.py | 349 +++++++++++ 3 files changed, 975 insertions(+) create mode 100644 python/sglang/jit_kernel/csrc/moe/moe_align_kernel.cu create mode 100644 python/sglang/jit_kernel/moe_align.py create mode 100644 python/sglang/jit_kernel/tests/test_moe_align_block_size.py diff --git a/python/sglang/jit_kernel/csrc/moe/moe_align_kernel.cu b/python/sglang/jit_kernel/csrc/moe/moe_align_kernel.cu new file mode 100644 index 000000000..ae1fd0dcd --- /dev/null +++ b/python/sglang/jit_kernel/csrc/moe/moe_align_kernel.cu @@ -0,0 +1,580 @@ +/* Copyright 2025 SGLang Team. All Rights Reserved. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +==============================================================================*/ + +#include +#include + +#include + +#include + +#include + +#ifndef WARP_SIZE +#define WARP_SIZE 32 +#endif + +#define CEILDIV(x, y) (((x) + (y) - 1) / (y)) + +#define VEC_SIZE 4 +using Vec = int4; + +inline uint32_t next_pow2(uint32_t x) noexcept { + --x; + x |= x >> 1; + x |= x >> 2; + x |= x >> 4; + x |= x >> 8; + x |= x >> 16; + return x + 1; +} + +namespace moe { + +__device__ __forceinline__ int warp_exclusive_scan(int v, unsigned mask = 0xffffffffu) { + int original = v; +#pragma unroll + for (int offset = 1; offset < WARP_SIZE; offset <<= 1) { + int n = __shfl_up_sync(mask, v, offset); + if ((threadIdx.x & (WARP_SIZE - 1)) >= offset) v += n; + } + return v - original; +} + +template +__global__ void count_and_sort_expert_tokens_kernel( + const scalar_t* __restrict__ topk_ids, + int32_t* __restrict__ sorted_token_ids, + int32_t* __restrict__ cumsum_buffer, + size_t numel) { + const size_t tid = blockIdx.x * blockDim.x + threadIdx.x; + const size_t stride = blockDim.x * gridDim.x; + + for (size_t i = tid; i < numel; i += stride) { + int32_t expert_id = topk_ids[i] + 1; + int32_t rank_post_pad = atomicAdd(&cumsum_buffer[expert_id], 1); + sorted_token_ids[rank_post_pad] = i; + } +} + +template +__global__ void moe_align_block_size_kernel( + const scalar_t* __restrict__ topk_ids, + int32_t* __restrict__ sorted_token_ids, + int32_t* __restrict__ expert_ids, + int32_t* __restrict__ total_tokens_post_pad, + int32_t num_experts, + int32_t block_size, + size_t numel, + int32_t* __restrict__ cumsum, + bool pad_sorted_token_ids, + const int32_t scan_size, + int32_t max_num_tokens_padded) { + // Use a separate thread block to populate sorted_token_ids + if (blockIdx.x == 1) { + if (pad_sorted_token_ids) { + Vec fill_vec; + fill_vec.x = fill_vec.y = fill_vec.z = fill_vec.w = numel; + int32_t total_vecs = (max_num_tokens_padded + VEC_SIZE - 1) / VEC_SIZE; + Vec* out_ptr = reinterpret_cast(sorted_token_ids); + for (int32_t i = threadIdx.x; i < total_vecs; i += blockDim.x) { + out_ptr[i] = fill_vec; + } + } + return; + } + + extern __shared__ int32_t smem[]; + int32_t* shared_counts = smem; // [num_experts] + int32_t* prefix = shared_counts + num_experts; // [num_experts + 1] + int32_t* scan_buf = prefix + num_experts + 1; // [scan_size] + __shared__ int32_t s_total_tokens_post_pad; + + const size_t tid = threadIdx.x; + const size_t stride = blockDim.x; + + if (tid < num_experts) { + shared_counts[tid] = 0; + } + + __syncthreads(); + + for (size_t i = tid; i < numel; i += stride) { + int expert_id = topk_ids[i] + 1; + atomicAdd(&shared_counts[expert_id], 1); + } + + __syncthreads(); + + int32_t padded_count = 0; + if (tid < num_experts) { + int32_t count = shared_counts[tid]; + padded_count = (count + block_size - 1) / block_size * block_size; + scan_buf[tid] = padded_count; + } + +#ifndef __CUDA_ARCH__ // HIP + + if (tid >= num_experts && tid < scan_size) { + scan_buf[tid] = 0; + } + + __syncthreads(); + + // Blelloch scan + int offset = 1; +#pragma unroll + for (int d = scan_size >> 1; d > 0; d >>= 1) { + if (tid < d) { + int ai = offset * (2 * tid + 1) - 1; + int bi = offset * (2 * tid + 2) - 1; + scan_buf[bi] += scan_buf[ai]; + } + offset <<= 1; + __syncthreads(); + } + + // down-sweep + if (tid == 0) { + prefix[num_experts] = scan_buf[scan_size - 1]; + scan_buf[scan_size - 1] = 0; + } + __syncthreads(); + +#pragma unroll + for (int d = 1; d < scan_size; d <<= 1) { + offset >>= 1; + if (tid < d) { + int ai = offset * (2 * tid + 1) - 1; + int bi = offset * (2 * tid + 2) - 1; + if (bi < scan_size) { + int temp = scan_buf[ai]; + scan_buf[ai] = scan_buf[bi]; + scan_buf[bi] += temp; + } + } + __syncthreads(); + } + + if (tid < num_experts) { + prefix[tid] = scan_buf[tid]; + } + + if (tid == 0) { + s_total_tokens_post_pad = prefix[num_experts]; + *total_tokens_post_pad = s_total_tokens_post_pad; + } + __syncthreads(); + +#else // CUDA + + // Intra warp prefix sum + int32_t* warp_sums = scan_buf + scan_size; // [<= 32] + const int warp_id = tid / WARP_SIZE; + const int lane_id = tid & (WARP_SIZE - 1); + const int num_warps_for_scan = (scan_size + WARP_SIZE - 1) / WARP_SIZE; + const int warp_sum = warp_exclusive_scan(padded_count) + padded_count; + if (lane_id == WARP_SIZE - 1) warp_sums[warp_id] = warp_sum; + __syncthreads(); + + // warp0 accumulate all the block's prefix sum + if (tid < WARP_SIZE) { + int val = (tid < num_warps_for_scan) ? warp_sums[tid] : 0; + int incl = warp_exclusive_scan(val) + val; + warp_sums[tid] = incl; + } + __syncthreads(); + + // Every thread obtains the whole block's sum + if (tid == 0) { + prefix[num_experts] = warp_sums[num_warps_for_scan - 1]; + s_total_tokens_post_pad = prefix[num_experts]; + *total_tokens_post_pad = s_total_tokens_post_pad; + } + __syncthreads(); + + // Fill 0 to scan_buf extended area (tid >= num_expert) + if (tid >= num_experts && tid < scan_size) scan_buf[tid] = 0; + __syncthreads(); + + // Perform 2 level exclusive-prefix-sum to scan_buf + int v = (tid < scan_size) ? scan_buf[tid] : 0; + int pre = warp_exclusive_scan(v); + if (lane_id == WARP_SIZE - 1) warp_sums[warp_id] = pre + v; + __syncthreads(); + + if (warp_id == 0) { + int val = (lane_id < num_warps_for_scan) ? warp_sums[lane_id] : 0; + warp_sums[lane_id] = warp_exclusive_scan(val); + } + __syncthreads(); + + int offset = warp_sums[warp_id]; + if (tid < scan_size) scan_buf[tid] = pre + offset; + __syncthreads(); + + // Write prefix[0..num_experts - 1] and cumsum + if (tid < num_experts) prefix[tid] = scan_buf[tid]; +#endif + + if (tid <= num_experts) { + cumsum[tid] = prefix[tid]; + } + // fill expert_ids + const int32_t num_blocks = s_total_tokens_post_pad / block_size; + for (int32_t i = tid; i < num_blocks; i += stride) { + int32_t block_start = i * block_size; + int left = 0, right = num_experts; + while (left < right) { + int mid = (left + right) >> 1; + if (prefix[mid] <= block_start) { + left = mid + 1; + } else { + right = mid; + } + } + expert_ids[i] = left - 2; + } +} + +template +__global__ void moe_align_block_size_small_batch_expert_kernel( + const scalar_t* __restrict__ topk_ids, + int32_t* __restrict__ sorted_token_ids, + int32_t* __restrict__ expert_ids, + int32_t* __restrict__ total_tokens_post_pad, + int32_t num_experts, + int32_t block_size, + size_t numel, + bool pad_sorted_token_ids, + int32_t max_num_tokens_padded) { + // Adapted from + // https://github.com/vllm-project/vllm/pull/29642/files#diff-5647b1413f4ae9aacba904eca8f8a8aee9079321eadff4c10101a2c6962dcc53R226 + // Use an additional group of threads to fill sorted_token_ids. + // Since the kernel will use sorted_token_ids afterward, + // we fill sorted_token_ids within the same threadblock to make + // synchronization easier. + if (threadIdx.x < fill_threads) { + // Initialize sorted_token_ids with numel + if (pad_sorted_token_ids) { + for (int32_t it = threadIdx.x; it < max_num_tokens_padded; it += fill_threads) { + sorted_token_ids[it] = numel; + } + } + // Three __syncthreads() corresponding to the other threads + __syncthreads(); + __syncthreads(); + __syncthreads(); + return; + } + + const size_t tid = threadIdx.x - fill_threads; + const size_t stride = blockDim.x - fill_threads; + + extern __shared__ int32_t shared_mem[]; + int32_t* cumsum = shared_mem; + int32_t* tokens_cnts = (int32_t*)(shared_mem + num_experts + 1); + + for (int i = 0; i < num_experts; ++i) { + tokens_cnts[(tid + 1) * num_experts + i] = 0; + } + + for (size_t i = tid; i < numel; i += stride) { + int32_t expert_id = topk_ids[i] + 1; + ++tokens_cnts[(tid + 1) * num_experts + expert_id]; + } + + __syncthreads(); + + if (tid < num_experts) { + tokens_cnts[tid] = 0; + for (int i = 1; i <= stride; ++i) { + tokens_cnts[i * num_experts + tid] += tokens_cnts[(i - 1) * num_experts + tid]; + } + } + + __syncthreads(); + + if (tid == 0) { + cumsum[0] = 0; + for (int i = 1; i <= num_experts; ++i) { + cumsum[i] = cumsum[i - 1] + CEILDIV(tokens_cnts[stride * num_experts + i - 1], block_size) * block_size; + } + *total_tokens_post_pad = static_cast(cumsum[num_experts]); + } + + __syncthreads(); + + if (tid < num_experts) { + for (int i = cumsum[tid]; i < cumsum[tid + 1]; i += block_size) { + expert_ids[i / block_size] = tid - 1; + } + } + + for (size_t i = tid; i < numel; i += stride) { + int32_t expert_id = topk_ids[i] + 1; + int32_t rank_post_pad = tokens_cnts[tid * num_experts + expert_id] + cumsum[expert_id]; + sorted_token_ids[rank_post_pad] = i; + ++tokens_cnts[tid * num_experts + expert_id]; + } +} + +// v2 kernel: supports >1024 experts via EXPERTS_PER_THREAD templating +// and a two-level warp scan (no cub dependency). Uses the same +1 offset +// convention as the original kernel (topk_ids shifted by +1 so -1 maps to 0). +// Launched with <<<2, 1024>>>: block 1 fills sorted_token_ids in parallel +// with block 0 doing the alignment compute. +// +// With 1024 threads and EXPERTS_PER_THREAD=4, covers at most 4096 expert +// indices. Since num_experts includes the +1 offset bucket, this supports +// up to 4095 real experts. +template +__global__ void moe_align_block_size_kernel_v2( + const scalar_t* __restrict__ topk_ids, + int32_t* __restrict__ sorted_token_ids, + int32_t* __restrict__ expert_ids, + int32_t* __restrict__ total_tokens_post_pad, + int32_t num_experts, + int32_t padded_num_experts, + int32_t block_size, + size_t numel, + int32_t* __restrict__ cumsum, + bool pad_sorted_token_ids, + int32_t max_num_tokens_padded) { + // Use a separate thread block to populate sorted_token_ids + if (blockIdx.x == 1) { + if (pad_sorted_token_ids) { + Vec fill_vec; + fill_vec.x = fill_vec.y = fill_vec.z = fill_vec.w = numel; + int32_t total_vecs = (max_num_tokens_padded + VEC_SIZE - 1) / VEC_SIZE; + Vec* out_ptr = reinterpret_cast(sorted_token_ids); + for (int32_t i = threadIdx.x; i < total_vecs; i += blockDim.x) { + out_ptr[i] = fill_vec; + } + } + return; + } + + extern __shared__ int32_t smem[]; + // Layout: shared_counts[padded_num_experts] | warp_sums[WARP_SIZE] + int32_t* shared_counts = smem; + int32_t* warp_sums = smem + padded_num_experts; + + const size_t tid = threadIdx.x; + const int warp_id = tid / WARP_SIZE; + const int lane_id = tid & (WARP_SIZE - 1); + + // Phase 1: Zero shared counts and count tokens per expert + const int my_start = tid * EXPERTS_PER_THREAD; + for (size_t i = tid; i < padded_num_experts; i += blockDim.x) { + shared_counts[i] = 0; + } + + __syncthreads(); + + for (size_t i = tid; i < numel; i += blockDim.x) { + int expert_id = topk_ids[i] + 1; // +1 offset convention + if (expert_id < num_experts) { + atomicAdd(&shared_counts[expert_id], 1); + } + } + + __syncthreads(); + + // Phase 2: Compute padded counts and two-level warp exclusive prefix sum + int32_t local_padded[EXPERTS_PER_THREAD]; + int32_t thread_sum = 0; + for (int i = 0; i < EXPERTS_PER_THREAD; ++i) { + int eid = my_start + i; + if (eid < num_experts) { + local_padded[i] = CEILDIV(shared_counts[eid], block_size) * block_size; + } else { + local_padded[i] = 0; + } + thread_sum += local_padded[i]; + } + + // Level 1: intra-warp exclusive scan on thread_sum + int32_t warp_prefix = warp_exclusive_scan(thread_sum); + int32_t warp_total = warp_prefix + thread_sum; + if (lane_id == WARP_SIZE - 1) warp_sums[warp_id] = warp_total; + __syncthreads(); + + // Level 2: warp 0 scans the per-warp totals + const int num_warps = (blockDim.x + WARP_SIZE - 1) / WARP_SIZE; + if (tid < WARP_SIZE) { + int val = (tid < num_warps) ? warp_sums[tid] : 0; + warp_sums[tid] = warp_exclusive_scan(val); + } + __syncthreads(); + + // Combine: thread_prefix = warp_sums[warp_id] + warp_prefix + int32_t thread_prefix = warp_sums[warp_id] + warp_prefix; + + // Local sequential prefix sum within each thread's expert group + int32_t running = 0; + for (int i = 0; i < EXPERTS_PER_THREAD; ++i) { + int eid = my_start + i; + if (eid <= num_experts) { + cumsum[eid] = thread_prefix + running; + } + running += local_padded[i]; + } + + // Last thread writes total + if (tid == blockDim.x - 1) { + cumsum[num_experts] = thread_prefix + thread_sum; + *total_tokens_post_pad = thread_prefix + thread_sum; + } + + __syncthreads(); + + // Phase 3: Fill expert_ids (eid - 1 to match sgl-kernel convention) + for (int i = 0; i < EXPERTS_PER_THREAD; ++i) { + int eid = my_start + i; + if (eid < num_experts) { + for (int j = cumsum[eid]; j < cumsum[eid + 1]; j += block_size) { + expert_ids[j / block_size] = eid - 1; + } + } + } +} + +} // namespace moe + +namespace { + +template +struct MoeAlignBlockSizeKernel { + static void + run(tvm::ffi::TensorView topk_ids, + int64_t num_experts, + int64_t block_size, + tvm::ffi::TensorView sorted_token_ids, + tvm::ffi::TensorView expert_ids, + tvm::ffi::TensorView num_tokens_post_pad, + tvm::ffi::TensorView cumsum_buffer, + bool pad_sorted_token_ids) { + using namespace host; + + auto device = topk_ids.device(); + const cudaStream_t stream = LaunchKernel::resolve_device(device); + + int threads = 1024; + threads = ((threads + WARP_SIZE - 1) / WARP_SIZE) * WARP_SIZE; + + int64_t max_num_tokens_padded = sorted_token_ids.size(0); + + // num_experts from Python is actual_num_experts + 1 (for EP offset convention). + // The v2 kernel (>1024 experts) uses 1024 threads with EXPERTS_PER_THREAD=4, + // covering at most 4096 expert indices, so num_experts (including the +1 + // offset bucket) must be <= 4096. This means up to 4095 real experts. + RuntimeCheck(num_experts <= 4096, "moe_align_block_size: num_experts must be <= 4096, got ", num_experts); + + const scalar_t* topk_ids_ptr = static_cast(topk_ids.data_ptr()); + int32_t* sorted_token_ids_ptr = static_cast(sorted_token_ids.data_ptr()); + int32_t* expert_ids_ptr = static_cast(expert_ids.data_ptr()); + int32_t* num_tokens_post_pad_ptr = static_cast(num_tokens_post_pad.data_ptr()); + int32_t* cumsum_buffer_ptr = static_cast(cumsum_buffer.data_ptr()); + size_t numel = topk_ids.numel(); + + bool small_batch_expert_mode = (numel < 1024) && (num_experts <= 64); + + if (small_batch_expert_mode) { + const int32_t num_thread = std::max((int32_t)num_experts, (int32_t)WARP_SIZE); + constexpr int32_t fill_threads = 256; + const int32_t shared_mem_size = ((num_thread + 1) * num_experts + (num_experts + 1)) * sizeof(int32_t); + + auto kernel = moe::moe_align_block_size_small_batch_expert_kernel; + LaunchKernel(dim3(1), dim3(fill_threads + num_thread), stream, shared_mem_size)( + kernel, + topk_ids_ptr, + sorted_token_ids_ptr, + expert_ids_ptr, + num_tokens_post_pad_ptr, + (int32_t)num_experts, + (int32_t)block_size, + numel, + pad_sorted_token_ids, + (int32_t)max_num_tokens_padded); + } else if (num_experts <= 1024) { + const size_t scan_size = next_pow2(num_experts); + const size_t shared_mem_size = (num_experts + (num_experts + 1) + scan_size + WARP_SIZE) * sizeof(int32_t); + + auto align_kernel = moe::moe_align_block_size_kernel; + LaunchKernel(dim3(2), dim3(threads), stream, shared_mem_size)( + align_kernel, + topk_ids_ptr, + sorted_token_ids_ptr, + expert_ids_ptr, + num_tokens_post_pad_ptr, + (int32_t)num_experts, + (int32_t)block_size, + numel, + cumsum_buffer_ptr, + pad_sorted_token_ids, + (int32_t)scan_size, + (int32_t)max_num_tokens_padded); + + const int block_threads = std::min(256, threads); + const int num_blocks = (numel + block_threads - 1) / block_threads; + const int max_blocks = 65535; + const int actual_blocks = std::min(num_blocks, max_blocks); + + auto sort_kernel = moe::count_and_sort_expert_tokens_kernel; + LaunchKernel(dim3(actual_blocks), dim3(block_threads), stream)( + sort_kernel, topk_ids_ptr, sorted_token_ids_ptr, cumsum_buffer_ptr, numel); + } else { + // v2 path for >1024 experts: two-level warp scan with EXPERTS_PER_THREAD + int64_t padded_num_experts = ((num_experts + WARP_SIZE - 1) / WARP_SIZE) * WARP_SIZE; + size_t shared_mem_size = (padded_num_experts + WARP_SIZE) * sizeof(int32_t); + + auto launch_v2 = [&](auto ept_tag) { + constexpr int EPT = decltype(ept_tag)::value; + auto v2_kernel = moe::moe_align_block_size_kernel_v2; + LaunchKernel(dim3(2), dim3(threads), stream, shared_mem_size)( + v2_kernel, + topk_ids_ptr, + sorted_token_ids_ptr, + expert_ids_ptr, + num_tokens_post_pad_ptr, + (int32_t)num_experts, + (int32_t)padded_num_experts, + (int32_t)block_size, + numel, + cumsum_buffer_ptr, + pad_sorted_token_ids, + (int32_t)max_num_tokens_padded); + }; + + if (padded_num_experts <= 2048) { + launch_v2(std::integral_constant{}); + } else { + launch_v2(std::integral_constant{}); + } + + const int block_threads = std::min(256, threads); + const int num_blocks = (numel + block_threads - 1) / block_threads; + const int max_blocks = 65535; + const int actual_blocks = std::min(num_blocks, max_blocks); + + auto sort_kernel = moe::count_and_sort_expert_tokens_kernel; + LaunchKernel(dim3(actual_blocks), dim3(block_threads), stream)( + sort_kernel, topk_ids_ptr, sorted_token_ids_ptr, cumsum_buffer_ptr, numel); + } + } +}; + +} // namespace diff --git a/python/sglang/jit_kernel/moe_align.py b/python/sglang/jit_kernel/moe_align.py new file mode 100644 index 000000000..ee136f1a8 --- /dev/null +++ b/python/sglang/jit_kernel/moe_align.py @@ -0,0 +1,46 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING + +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_moe_align_module(dtype: torch.dtype) -> Module: + args = make_cpp_args(dtype) + return load_jit( + "moe_align_block_size", + *args, + cuda_files=["moe/moe_align_kernel.cu"], + cuda_wrappers=[ + ("moe_align_block_size", f"MoeAlignBlockSizeKernel<{args}>::run"), + ], + ) + + +def moe_align_block_size( + topk_ids: torch.Tensor, + num_experts: int, + block_size: int, + sorted_token_ids: torch.Tensor, + expert_ids: torch.Tensor, + num_tokens_post_pad: torch.Tensor, + cumsum_buffer: torch.Tensor, + pad_sorted_token_ids: bool = False, +) -> None: + module = _jit_moe_align_module(topk_ids.dtype) + module.moe_align_block_size( + topk_ids, + num_experts, + block_size, + sorted_token_ids, + expert_ids, + num_tokens_post_pad, + cumsum_buffer, + pad_sorted_token_ids, + ) diff --git a/python/sglang/jit_kernel/tests/test_moe_align_block_size.py b/python/sglang/jit_kernel/tests/test_moe_align_block_size.py new file mode 100644 index 000000000..92905058a --- /dev/null +++ b/python/sglang/jit_kernel/tests/test_moe_align_block_size.py @@ -0,0 +1,349 @@ +import itertools +import sys + +import pytest +import torch +import triton +import triton.language as tl + +from sglang.jit_kernel.moe_align import moe_align_block_size +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=28, suite="stage-b-kernel-unit-1-gpu-large") +register_cuda_ci(est_time=120, suite="nightly-kernel-1-gpu", nightly=True) + + +def ceil_div(a, b): + return (a + b - 1) // b + + +@triton.jit +def moe_align_block_size_stage1( + topk_ids_ptr, + tokens_cnts_ptr, + num_experts: tl.constexpr, + numel: tl.constexpr, + tokens_per_thread: tl.constexpr, +): + pid = tl.program_id(0) + start_idx = pid * tokens_per_thread + off_c = (pid + 1) * num_experts + + for i in range(tokens_per_thread): + if start_idx + i < numel: + idx = tl.load(topk_ids_ptr + start_idx + i) + token_cnt = tl.load(tokens_cnts_ptr + off_c + idx) + tl.store(tokens_cnts_ptr + off_c + idx, token_cnt + 1) + + +@triton.jit +def moe_align_block_size_stage2( + tokens_cnts_ptr, + num_experts: tl.constexpr, +): + pid = tl.program_id(0) + last_cnt = 0 + for i in range(1, num_experts + 1): + token_cnt = tl.load(tokens_cnts_ptr + i * num_experts + pid) + last_cnt = last_cnt + token_cnt + tl.store(tokens_cnts_ptr + i * num_experts + pid, last_cnt) + + +@triton.jit +def moe_align_block_size_stage3( + total_tokens_post_pad_ptr, + tokens_cnts_ptr, + cumsum_ptr, + num_experts: tl.constexpr, + block_size: tl.constexpr, +): + last_cumsum = 0 + off_cnt = num_experts * num_experts + for i in range(1, num_experts + 1): + token_cnt = tl.load(tokens_cnts_ptr + off_cnt + i - 1) + last_cumsum = last_cumsum + tl.cdiv(token_cnt, block_size) * block_size + tl.store(cumsum_ptr + i, last_cumsum) + tl.store(total_tokens_post_pad_ptr, last_cumsum) + + +@triton.jit +def moe_align_block_size_stage4( + topk_ids_ptr, + sorted_token_ids_ptr, + expert_ids_ptr, + tokens_cnts_ptr, + cumsum_ptr, + num_experts: tl.constexpr, + block_size: tl.constexpr, + numel: tl.constexpr, + tokens_per_thread: tl.constexpr, +): + pid = tl.program_id(0) + start_idx = tl.load(cumsum_ptr + pid) + end_idx = tl.load(cumsum_ptr + pid + 1) + + for i in range(start_idx, end_idx, block_size): + tl.store(expert_ids_ptr + i // block_size, pid) + + start_idx = pid * tokens_per_thread + off_t = pid * num_experts + + for i in range(start_idx, tl.minimum(start_idx + tokens_per_thread, numel)): + expert_id = tl.load(topk_ids_ptr + i) + token_cnt = tl.load(tokens_cnts_ptr + off_t + expert_id) + rank_post_pad = token_cnt + tl.load(cumsum_ptr + expert_id) + tl.store(sorted_token_ids_ptr + rank_post_pad, i) + tl.store(tokens_cnts_ptr + off_t + expert_id, token_cnt + 1) + + +def moe_align_block_size_triton( + topk_ids: torch.Tensor, + num_experts: int, + block_size: int, + sorted_token_ids: torch.Tensor, + expert_ids: torch.Tensor, + num_tokens_post_pad: torch.Tensor, +) -> None: + numel = topk_ids.numel() + grid = (num_experts,) + tokens_cnts = torch.zeros( + (num_experts + 1, num_experts), dtype=torch.int32, device=topk_ids.device + ) + cumsum = torch.zeros((num_experts + 1,), dtype=torch.int32, device=topk_ids.device) + tokens_per_thread = ceil_div(numel, num_experts) + + moe_align_block_size_stage1[grid]( + topk_ids, + tokens_cnts, + num_experts, + numel, + tokens_per_thread, + ) + moe_align_block_size_stage2[grid]( + tokens_cnts, + num_experts, + ) + moe_align_block_size_stage3[(1,)]( + num_tokens_post_pad, + tokens_cnts, + cumsum, + num_experts, + block_size, + ) + moe_align_block_size_stage4[grid]( + topk_ids, + sorted_token_ids, + expert_ids, + tokens_cnts, + cumsum, + num_experts, + block_size, + numel, + tokens_per_thread, + ) + + +@pytest.mark.parametrize( + "block_size,num_tokens,topk,num_experts,pad_sorted_token_ids", + list( + itertools.product( + [32, 64, 128, 256], # block_size + [1, 2, 4, 8, 16, 32, 64, 128, 256, 512, 1024, 2048, 4096], # num_tokens + [1, 2, 4, 8, 16, 32, 64], # topk + [64, 160, 256, 257, 260, 264], # num_experts + [True, False], # pad_sorted_token_ids + ) + ), +) +def test_moe_align_block_size_compare_implementations( + block_size, num_tokens, topk, num_experts, pad_sorted_token_ids +): + topk_ids = torch.argsort(torch.rand(num_tokens, num_experts, device="cuda"), dim=1)[ + :, :topk + ] + + max_num_tokens_padded = topk_ids.numel() + (num_experts + 1) * (block_size - 1) + if topk_ids.numel() < num_experts + 1: + max_num_tokens_padded = topk_ids.numel() * block_size + + sorted_ids_cuda = torch.empty( + (max_num_tokens_padded,), dtype=torch.int32, device=topk_ids.device + ) + if not pad_sorted_token_ids: + sorted_ids_cuda.fill_(topk_ids.numel()) + max_num_m_blocks = max_num_tokens_padded // block_size + expert_ids_cuda = torch.zeros( + (max_num_m_blocks,), dtype=torch.int32, device=topk_ids.device + ) + num_tokens_post_pad_cuda = torch.empty( + (1), dtype=torch.int32, device=topk_ids.device + ) + cumsum_buffer = torch.empty( + num_experts + 2, dtype=torch.int32, device=topk_ids.device + ) + + sorted_ids_triton = torch.empty_like(sorted_ids_cuda) + sorted_ids_triton.fill_(topk_ids.numel()) + expert_ids_triton = torch.zeros_like(expert_ids_cuda) + num_tokens_post_pad_triton = torch.empty_like(num_tokens_post_pad_cuda) + + moe_align_block_size( + topk_ids, + num_experts + 1, + block_size, + sorted_ids_cuda, + expert_ids_cuda, + num_tokens_post_pad_cuda, + cumsum_buffer, + pad_sorted_token_ids, + ) + + moe_align_block_size_triton( + topk_ids, + num_experts + 1, + block_size, + sorted_ids_triton, + expert_ids_triton, + num_tokens_post_pad_triton, + ) + + assert torch.allclose(expert_ids_cuda, expert_ids_triton, atol=0, rtol=0), ( + f"Expert IDs mismatch for block_size={block_size}, " + f"num_tokens={num_tokens}, topk={topk}\n" + f"CUDA expert_ids: {expert_ids_cuda}\n" + f"Triton expert_ids: {expert_ids_triton}" + ) + + assert torch.allclose( + num_tokens_post_pad_cuda, num_tokens_post_pad_triton, atol=0, rtol=0 + ), ( + f"Num tokens post pad mismatch for block_size={block_size}, " + f"num_tokens={num_tokens}, topk={topk}\n" + f"CUDA num_tokens_post_pad: {num_tokens_post_pad_cuda}\n" + f"Triton num_tokens_post_pad: {num_tokens_post_pad_triton}" + ) + + # Select an expert to check + expert_idx = expert_ids_cuda.max().item() + + # Get the first and last block id where expert_ids_cuda == expert_idx + matching_indices = torch.where(expert_ids_cuda == expert_idx)[0] + block_sorted_start = matching_indices[0].item() * block_size + block_sorted_end = min( + (matching_indices[-1].item() + 1) * block_size, + num_tokens_post_pad_cuda.item(), + ) + + selected_sorted_ids_cuda = sorted_ids_cuda[ + block_sorted_start:block_sorted_end + ].sort()[0] + selected_sorted_ids_triton = sorted_ids_triton[ + block_sorted_start:block_sorted_end + ].sort()[0] + + assert torch.allclose( + selected_sorted_ids_cuda, + selected_sorted_ids_triton, + atol=0, + rtol=0, + ), ( + f"Sorted IDs mismatch for block_size={block_size}, " + f"num_tokens={num_tokens}, topk={topk}\n" + f"CUDA sorted_ids: {selected_sorted_ids_cuda}\n" + f"Triton sorted_ids: {selected_sorted_ids_triton}" + ) + + +@pytest.mark.parametrize( + "block_size,num_tokens,topk,num_experts", + list( + itertools.product( + [64, 128], # block_size + [1, 8, 32, 256], # num_tokens + [8], # topk + [ + 1025, + 2048, + 4095, + ], # num_experts (>1024 to exercise v2 kernel, max 4095 real experts) + ) + ), +) +def test_moe_align_block_size_v2_large_num_experts( + block_size, num_tokens, topk, num_experts +): + """Test moe_align_block_size v2 kernel for >1024 experts against Triton reference.""" + topk_ids = torch.randint( + 0, num_experts, (num_tokens, topk), dtype=torch.int32, device="cuda" + ) + + max_num_tokens_padded = topk_ids.numel() + (num_experts + 1) * (block_size - 1) + if topk_ids.numel() < num_experts + 1: + max_num_tokens_padded = topk_ids.numel() * block_size + + sorted_ids_cuda = torch.empty( + (max_num_tokens_padded,), dtype=torch.int32, device=topk_ids.device + ) + sorted_ids_cuda.fill_(topk_ids.numel()) + max_num_m_blocks = max_num_tokens_padded // block_size + expert_ids_cuda = torch.zeros( + (max_num_m_blocks,), dtype=torch.int32, device=topk_ids.device + ) + num_tokens_post_pad_cuda = torch.empty( + (1), dtype=torch.int32, device=topk_ids.device + ) + cumsum_buffer = torch.empty( + num_experts + 2, dtype=torch.int32, device=topk_ids.device + ) + + sorted_ids_triton = torch.empty_like(sorted_ids_cuda) + sorted_ids_triton.fill_(topk_ids.numel()) + expert_ids_triton = torch.zeros_like(expert_ids_cuda) + num_tokens_post_pad_triton = torch.empty_like(num_tokens_post_pad_cuda) + + moe_align_block_size( + topk_ids, + num_experts + 1, + block_size, + sorted_ids_cuda, + expert_ids_cuda, + num_tokens_post_pad_cuda, + cumsum_buffer, + True, + ) + + moe_align_block_size_triton( + topk_ids, + num_experts + 1, + block_size, + sorted_ids_triton, + expert_ids_triton, + num_tokens_post_pad_triton, + ) + + assert torch.equal(num_tokens_post_pad_cuda, num_tokens_post_pad_triton), ( + f"Num tokens post pad mismatch: CUDA={num_tokens_post_pad_cuda.item()}, " + f"Triton={num_tokens_post_pad_triton.item()}" + ) + + ntp = num_tokens_post_pad_cuda.item() + num_blocks = ntp // block_size + + assert torch.equal(expert_ids_cuda[:num_blocks], expert_ids_triton[:num_blocks]), ( + f"Expert IDs mismatch for block_size={block_size}, " + f"num_tokens={num_tokens}, topk={topk}, num_experts={num_experts}" + ) + + # Compare sorted_token_ids per expert block (order within block may differ) + for b in range(num_blocks): + s, e = b * block_size, (b + 1) * block_size + block_cuda = sorted_ids_cuda[s:e].sort().values + block_triton = sorted_ids_triton[s:e].sort().values + assert torch.equal(block_cuda, block_triton), ( + f"Block {b} sorted_ids mismatch for num_experts={num_experts}, " + f"num_tokens={num_tokens}" + ) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__]))