diff --git a/python/sglang/jit_kernel/csrc/gemm/dsv3_router_gemm.cuh b/python/sglang/jit_kernel/csrc/gemm/dsv3_router_gemm.cuh new file mode 100644 index 000000000..1fcc5b63b --- /dev/null +++ b/python/sglang/jit_kernel/csrc/gemm/dsv3_router_gemm.cuh @@ -0,0 +1,184 @@ +/* + * Adapted from + * https://github.com/NVIDIA/TensorRT-LLM/blob/main/cpp/tensorrt_llm/kernels/dsv3MinLatencyKernels/dsv3RouterGemm.cu + * https://github.com/NVIDIA/TensorRT-LLM/blob/main/cpp/tensorrt_llm/thop/dsv3RouterGemmOp.cpp + * + * Copyright (c) 2019-2023, NVIDIA CORPORATION. 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 +#include +#include + +#include + +#include + +namespace { + +using namespace device; + +static constexpr int kDefaultNumExperts = 256; +static constexpr int kKimiK2NumExperts = 384; +static constexpr int kDefaultHiddenDim = 7168; + +// kOutFloat: true = float32 output, false = bfloat16 output +template < + typename T, + typename OutT, + int kBlockSize, + int VPT, + int kNumTokens, + int kNumExperts, + int kHiddenDim, + bool kUsePDL> +__global__ __launch_bounds__(kBlockSize, 1) void router_gemm_kernel(OutT* out, T const* mat_a, T const* mat_b) { + constexpr int kWarpSize = 32; + constexpr int kNumWarps = kBlockSize / kWarpSize; + constexpr int kElemsPerKIter = VPT * kBlockSize; + static_assert(kHiddenDim % kElemsPerKIter == 0, "hidden_dim must be divisible by one K iteration"); + constexpr int kIters = kHiddenDim / kElemsPerKIter; + // Padding to avoid shared memory bank conflicts when kNumTokens > 8 + constexpr int kSmReductionPad = (kNumTokens > 8) ? 1 : 0; + static_assert(kSmReductionPad == 0 || kSmReductionPad == 1, "kSmReductionPad only supports 0 or 1"); + + int const n_idx = blockIdx.x; + int const tid = threadIdx.x; + int const warp_id = tid / kWarpSize; + int const lane_id = tid % kWarpSize; + + float acc[kNumTokens] = {}; + __shared__ float sm_reduction[kNumTokens][kNumWarps + kSmReductionPad]; + + T const* b_col = mat_b + n_idx * kHiddenDim; + + PDLWaitPrimary(); + + int k_base = tid * VPT; +#pragma unroll + for (int ki = 0; ki < kIters; ++ki, k_base += kElemsPerKIter) { + AlignedVector b_vec; + b_vec.load(b_col + k_base); +#pragma unroll + for (int m_idx = 0; m_idx < kNumTokens; ++m_idx) { + AlignedVector a_vec; + a_vec.load(mat_a + m_idx * kHiddenDim + k_base); +#pragma unroll + for (int k = 0; k < VPT; ++k) { + acc[m_idx] += cast(a_vec[k]) * cast(b_vec[k]); + } + } + } + +#pragma unroll + for (int m_idx = 0; m_idx < kNumTokens; ++m_idx) { + float sum = warp::reduce_sum(acc[m_idx]); + if (lane_id == 0) { + sm_reduction[m_idx][warp_id] = sum; + } + } + + __syncthreads(); + + if (warp_id == 0 && lane_id < kNumTokens) { + float final_sum = 0.0f; +#pragma unroll + for (int w = 0; w < kNumWarps; ++w) { + final_sum += sm_reduction[lane_id][w]; + } + out[lane_id * kNumExperts + n_idx] = cast(final_sum); + } + + PDLTriggerSecondary(); +} + +template +void invokeRouterGemm(OutT* output, T const* mat_a, T const* mat_b, DLDevice device) { + constexpr int VPT = 16 / sizeof(T); + constexpr int kBlockSize = 128; + constexpr auto kernel = router_gemm_kernel; + host::LaunchKernel(kNumExperts, kBlockSize, device).enable_pdl(kUsePDL)(kernel, output, mat_a, mat_b); +} + +// Dispatch runtime num_tokens to compile-time template parameter [kBegin, kEnd] +template +struct RouterGemmDispatcher { + static void run(int num_tokens, OutT* output, bf16_t const* mat_a, bf16_t const* mat_b, DLDevice device) { + if (num_tokens == kBegin) { + invokeRouterGemm(output, mat_a, mat_b, device); + } else { + RouterGemmDispatcher::run( + num_tokens, output, mat_a, mat_b, device); + } + } +}; + +// Base case: kBegin == kEnd +template +struct RouterGemmDispatcher { + static void run(int num_tokens, OutT* output, bf16_t const* mat_a, bf16_t const* mat_b, DLDevice device) { + if (num_tokens == kEnd) { + invokeRouterGemm(output, mat_a, mat_b, device); + } else { + host::panic({}, "dsv3_router_gemm: num_tokens must be between 1 and 16, got ", num_tokens); + } + } +}; + +// kNumExperts: compile-time 256 or 384 +// kHiddenDim: compile-time hidden dim, any multiple of one K iteration (1024) +// kUsePDL: compile-time bool (true on SM90+) +// kOutFloat: compile-time bool (true = float32 output, false = bfloat16 output) +template +struct DSV3RouterGemmKernel { + static_assert( + kNumExperts == kDefaultNumExperts || kNumExperts == kKimiK2NumExperts, + "required num_experts == 256 or num_experts == 384"); + + using OutT = std::conditional_t; + + static void + run(const tvm::ffi::TensorView mat_a, const tvm::ffi::TensorView mat_b, const tvm::ffi::TensorView output) { + using namespace host; + + auto M = SymbolicSize{"num_tokens"}; + auto K = SymbolicSize{"hidden_dim"}; + auto N = SymbolicSize{"num_experts"}; + auto device = SymbolicDevice{}; + K.set_value(kHiddenDim); + N.set_value(kNumExperts); + device.set_options(); + + TensorMatcher({M, K}).with_dtype().with_device(device).verify(mat_a); + TensorMatcher({N, K}).with_dtype().with_device(device).verify(mat_b); + TensorMatcher({M, N}).with_dtype().with_device(device).verify(output); + + const int num_tokens = static_cast(M.unwrap()); + + RouterGemmDispatcher<1, 16, OutT, kNumExperts, kHiddenDim, kUsePDL>::run( + num_tokens, + static_cast(output.data_ptr()), + static_cast(mat_a.data_ptr()), + static_cast(mat_b.data_ptr()), + device.unwrap()); + } +}; + +} // namespace diff --git a/python/sglang/jit_kernel/dsv3_router_gemm.py b/python/sglang/jit_kernel/dsv3_router_gemm.py new file mode 100644 index 000000000..f33d883dd --- /dev/null +++ b/python/sglang/jit_kernel/dsv3_router_gemm.py @@ -0,0 +1,92 @@ +""" +JIT kernel for DeepSeek V3 router GEMM. + +Replaces the AOT sgl_kernel.dsv3_router_gemm for SM90+ (Hopper) GPUs. +Supports num_experts in {256, 384}, hidden_dim a multiple of 1024, num_tokens 1-16. +""" + +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, +) +from sglang.kernel_api_logging import debug_kernel_api +from sglang.srt.utils.custom_op import register_custom_op + +if TYPE_CHECKING: + from tvm_ffi.module import Module + + +@cache_once +def _jit_dsv3_router_gemm_module( + num_experts: int, + hidden_dim: int, + use_pdl: bool, + out_float: bool, +) -> Module: + args = make_cpp_args(num_experts, hidden_dim, use_pdl, out_float) + return load_jit( + "dsv3_router_gemm", + *args, + cuda_files=["gemm/dsv3_router_gemm.cuh"], + cuda_wrappers=[ + ("dsv3_router_gemm", f"DSV3RouterGemmKernel<{args}>::run"), + ], + ) + + +@register_custom_op( + op_name="dsv3_router_gemm", + mutates_args=["output"], +) +def _dsv3_router_gemm_custom_op( + hidden_states: torch.Tensor, + router_weights: torch.Tensor, + output: torch.Tensor, +) -> None: + num_experts = router_weights.shape[0] + hidden_dim = hidden_states.shape[1] + out_float = output.dtype == torch.float32 + module = _jit_dsv3_router_gemm_module( + num_experts, hidden_dim, is_arch_support_pdl(), out_float + ) + module.dsv3_router_gemm(hidden_states, router_weights, output) + return None + + +@debug_kernel_api +def dsv3_router_gemm( + hidden_states: torch.Tensor, + router_weights: torch.Tensor, + out_dtype: torch.dtype = torch.bfloat16, + output: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """ + DeepSeek V3 router GEMM kernel (JIT variant). + + Args: + hidden_states: Input tensor of shape [num_tokens, hidden_dim], bfloat16. + hidden_dim must be a multiple of 1024 and num_tokens in [1, 16]. + router_weights: Weight tensor of shape [num_experts, hidden_dim], bfloat16. + out_dtype: Output dtype, either torch.bfloat16 or torch.float32. + output: Optional pre-allocated output tensor. + + Returns: + Output tensor of shape [num_tokens, num_experts]. + """ + if output is None: + output = torch.empty( + hidden_states.shape[0], + router_weights.shape[0], + device=hidden_states.device, + dtype=out_dtype, + ) + _dsv3_router_gemm_custom_op(hidden_states, router_weights, output) + return output diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 9ccb0797d..dfab7fd5b 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -201,8 +201,11 @@ if _use_aiter: pass if _is_cuda: - from flashinfer.gemm import mm_M1_16_K7168_N256 as _raw_dsv3_router_gemm - from sgl_kernel import dsv3_fused_a_gemm, dsv3_router_gemm + from sgl_kernel import dsv3_fused_a_gemm + + from sglang.jit_kernel.dsv3_router_gemm import ( + dsv3_router_gemm as _jit_dsv3_router_gemm, + ) elif _is_npu: from sglang.srt.hardware_backend.npu.modules.deepseek_v2_attention_mla_npu import ( forward_dsa_core_npu, @@ -213,7 +216,7 @@ elif _is_npu: forward_mla_prepare_npu, ) elif _is_musa: - from sgl_kernel import dsv3_fused_a_gemm, dsv3_router_gemm + from sgl_kernel import dsv3_fused_a_gemm else: pass @@ -494,28 +497,16 @@ class MoEGate(nn.Module): ): logits = F.linear(hidden_states, self.weight, None) else: - # NOTE: For some unknown reason, router_gemm seems degrade accept length. if ( _is_cuda and hidden_states.shape[0] <= 16 - and hidden_states.shape[1] == 7168 + and hidden_states.shape[1] % 1024 == 0 and (self.weight.shape[0] == 256 or self.weight.shape[0] == 384) and _device_sm >= 90 ): - if _device_sm in [100, 103] and self.weight.shape[0] == 256: - # TODO: will check the dtype to be bf16 - # router gemm output float32 - logits = torch.empty( - hidden_states.shape[0], - self.weight.shape[0], - device=hidden_states.device, - dtype=torch.float32, - ) - flashinfer_dsv3_router_gemm(logits, hidden_states, self.weight) - else: - logits = dsv3_router_gemm( - hidden_states, self.weight, out_dtype=torch.float32 - ) + logits = _jit_dsv3_router_gemm( + hidden_states, self.weight, out_dtype=torch.float32 + ) elif _use_aiter: logits = aiter_dsv3_router_gemm(hidden_states, self.weight) @@ -2946,24 +2937,6 @@ class DeepseekV32ForCausalLM(DeepseekV2ForCausalLM): pass -@register_custom_op( - op_name="flashinfer_dsv3_router_gemm", - mutates_args=[], - fake_impl=lambda logits, hidden_states, weight: None, -) -def flashinfer_dsv3_router_gemm( - logits: torch.Tensor, - hidden_states: torch.Tensor, - weight: torch.Tensor, -) -> None: - _raw_dsv3_router_gemm( - hidden_states, - weight.t(), - logits, - launch_with_pdl=True, - ) - - @register_custom_op(out_shape="hidden_states") def dsv2_flashinfer_moe_dual_stream_graph( hidden_states: torch.Tensor, diff --git a/sgl-kernel/CMakeLists.txt b/sgl-kernel/CMakeLists.txt index 99015ff18..e3409d2e1 100644 --- a/sgl-kernel/CMakeLists.txt +++ b/sgl-kernel/CMakeLists.txt @@ -264,9 +264,6 @@ set(SOURCES "csrc/gemm/awq_kernel.cu" "csrc/gemm/bmm_fp8.cu" "csrc/gemm/dsv3_fused_a_gemm.cu" - "csrc/gemm/dsv3_router_gemm_bf16_out.cu" - "csrc/gemm/dsv3_router_gemm_entry.cu" - "csrc/gemm/dsv3_router_gemm_float_out.cu" "csrc/gemm/fp8_blockwise_gemm_kernel.cu" "csrc/gemm/fp8_gemm_kernel.cu" "csrc/gemm/int8_gemm_kernel.cu" diff --git a/sgl-kernel/benchmark/bench_dsv3_router_gemm.py b/sgl-kernel/benchmark/bench_dsv3_router_gemm.py deleted file mode 100644 index 85ad88529..000000000 --- a/sgl-kernel/benchmark/bench_dsv3_router_gemm.py +++ /dev/null @@ -1,151 +0,0 @@ -import argparse - -import torch -import torch.nn.functional as F -import triton -import triton.testing -from sgl_kernel import dsv3_router_gemm - -from sglang.utils import is_in_ci - -IS_CI = is_in_ci() - -# CI environment uses simplified parameters -if IS_CI: - num_tokens_vals = [1] # Only test 1 value in CI - line_vals = ["sgl-kernel-256"] # Only test one implementation in CI -else: - num_tokens_vals = [i + 1 for i in range(16)] # Test 1-16 in full mode - line_vals = ["torch-256", "sgl-kernel-256", "torch-384", "sgl-kernel-384"] - - -@triton.testing.perf_report( - triton.testing.Benchmark( - x_names=["num_tokens"], - x_vals=num_tokens_vals, - x_log=False, - line_arg="impl", - line_vals=line_vals, - line_names=( - [ - "torch-256", - "dsv3_router_gemm-256", - "torch-384", - "dsv3_router_gemm-384", - ] - if not IS_CI - else ["dsv3_router_gemm-256"] - ), - styles=( - [("blue", "-"), ("orange", "-"), ("green", "-"), ("red", "-")] - if not IS_CI - else [("orange", "-")] - ), - ylabel="TFLOPs", - plot_name="input-bf16-output-bf16 dsv3 router gemm throughput", - args={}, - ) -) -def benchmark_bf16_output(num_tokens, impl): - # M: num_tokens, K: hidden_dim, N: num_experts - M, K = num_tokens, 7168 - - if impl == "torch-256" or impl == "sgl-kernel-256": - N = 256 - elif impl == "torch-384" or impl == "sgl-kernel-384": - N = 384 - else: - raise ValueError(f"Unknown impl: {impl}") - - mat_a = torch.randn((M, K), dtype=torch.bfloat16, device="cuda").contiguous() - mat_b = torch.randn((N, K), dtype=torch.bfloat16, device="cuda").contiguous() - - quantiles = [0.5, 0.2, 0.8] - - if impl == "torch-256" or impl == "torch-384": - - def runner(): - F.linear(mat_a, mat_b) - - elif impl == "sgl-kernel-256" or impl == "sgl-kernel-384": - - def runner(): - dsv3_router_gemm(mat_a, mat_b, out_dtype=torch.bfloat16) - - ms, min_ms, max_ms = triton.testing.do_bench_cudagraph(runner, quantiles=quantiles) - - def tflops(t_ms): - flops = 2 * M * K * N - return flops / (t_ms * 1e-3) / 1e12 - - return tflops(ms), tflops(max_ms), tflops(min_ms) - - -@triton.testing.perf_report( - triton.testing.Benchmark( - x_names=["num_tokens"], - x_vals=num_tokens_vals, - x_log=False, - line_arg="impl", - line_vals=line_vals, - line_names=( - [ - "torch-256", - "dsv3_router_gemm-256", - "torch-384", - "dsv3_router_gemm-384", - ] - if not IS_CI - else ["dsv3_router_gemm-256"] - ), - styles=( - [("blue", "-"), ("orange", "-"), ("green", "-"), ("red", "-")] - if not IS_CI - else [("orange", "-")] - ), - ylabel="TFLOPs", - plot_name="input-bf16-output-fp32 dsv3 router gemm throughput", - args={}, - ) -) -def benchmark_float_output(num_tokens, impl): - # M: num_tokens, K: hidden_dim, N: num_experts - M, K = num_tokens, 7168 - - if impl == "torch-256" or impl == "sgl-kernel-256": - N = 256 - elif impl == "torch-384" or impl == "sgl-kernel-384": - N = 384 - else: - raise ValueError(f"Unknown impl: {impl}") - - mat_a = torch.randn((M, K), dtype=torch.bfloat16, device="cuda").contiguous() - mat_b = torch.randn((N, K), dtype=torch.bfloat16, device="cuda").contiguous() - - quantiles = [0.5, 0.2, 0.8] - - if impl == "torch-256" or impl == "torch-384": - - def runner(): - F.linear(mat_a, mat_b).to(torch.float32) - - elif impl == "sgl-kernel-256" or impl == "sgl-kernel-384": - - def runner(): - dsv3_router_gemm(mat_a, mat_b, out_dtype=torch.float32) - - ms, min_ms, max_ms = triton.testing.do_bench_cudagraph(runner, quantiles=quantiles) - - def tflops(t_ms): - flops = 2 * M * K * N - return flops / (t_ms * 1e-3) / 1e12 - - return tflops(ms), tflops(max_ms), tflops(min_ms) - - -if __name__ == "__main__": - parser = argparse.ArgumentParser() - args = parser.parse_args() - - benchmark_bf16_output.run(print_data=True) - benchmark_float_output.run(print_data=True) diff --git a/sgl-kernel/csrc/common_extension.cc b/sgl-kernel/csrc/common_extension.cc index 8aff53976..f6261dbc7 100644 --- a/sgl-kernel/csrc/common_extension.cc +++ b/sgl-kernel/csrc/common_extension.cc @@ -135,9 +135,6 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) { m.def("dsv3_fused_a_gemm(Tensor! output, Tensor mat_a, Tensor mat_b) -> ()"); m.impl("dsv3_fused_a_gemm", torch::kCUDA, &dsv3_fused_a_gemm); - m.def("dsv3_router_gemm(Tensor! output, Tensor mat_a, Tensor mat_b) -> ()"); - m.impl("dsv3_router_gemm", torch::kCUDA, &dsv3_router_gemm); - /* * From csrc/gemm/gptq */ diff --git a/sgl-kernel/include/sgl_kernel_ops.h b/sgl-kernel/include/sgl_kernel_ops.h index 61c39b6b3..2110aa573 100644 --- a/sgl-kernel/include/sgl_kernel_ops.h +++ b/sgl-kernel/include/sgl_kernel_ops.h @@ -251,7 +251,6 @@ void bmm_fp8( at::Tensor B_scale, at::Tensor workspace_buffer, int64_t cublas_handle); -void dsv3_router_gemm(torch::Tensor& output, const torch::Tensor& mat_a, const torch::Tensor& mat_b); void dsv3_fused_a_gemm(torch::Tensor& output, torch::Tensor const& mat_a, torch::Tensor const& mat_b); torch::Tensor gptq_gemm( diff --git a/sgl-kernel/python/sgl_kernel/__init__.py b/sgl-kernel/python/sgl_kernel/__init__.py index ce6e9d049..e179d889c 100644 --- a/sgl-kernel/python/sgl_kernel/__init__.py +++ b/sgl-kernel/python/sgl_kernel/__init__.py @@ -57,7 +57,6 @@ else: awq_dequantize, bmm_fp8, dsv3_fused_a_gemm, - dsv3_router_gemm, fp8_blockwise_scaled_mm, fp8_scaled_mm, gptq_gemm, diff --git a/sgl-kernel/python/sgl_kernel/gemm.py b/sgl-kernel/python/sgl_kernel/gemm.py index a6e65cd6b..1d68cdf94 100644 --- a/sgl-kernel/python/sgl_kernel/gemm.py +++ b/sgl-kernel/python/sgl_kernel/gemm.py @@ -191,25 +191,6 @@ def qserve_w4a8_per_group_gemm( return out_feats -def dsv3_router_gemm( - hidden_states: torch.Tensor, - router_weights: torch.Tensor, - out_dtype: torch.dtype = torch.bfloat16, -) -> torch.Tensor: - output = torch.empty( - hidden_states.shape[0], - router_weights.shape[0], - device=hidden_states.device, - dtype=out_dtype, - ) - torch.ops.sgl_kernel.dsv3_router_gemm( - output, - hidden_states, - router_weights, - ) - return output - - def shuffle_rows(input_tensor, dst2src_map, output_tensor_shape): output_tensor = torch.empty( output_tensor_shape, diff --git a/sgl-kernel/tests/test_dsv3_router_gemm.py b/sgl-kernel/tests/test_dsv3_router_gemm.py deleted file mode 100644 index fa9a830aa..000000000 --- a/sgl-kernel/tests/test_dsv3_router_gemm.py +++ /dev/null @@ -1,37 +0,0 @@ -import sys - -import pytest -import torch -import torch.nn.functional as F -from sgl_kernel import dsv3_router_gemm - - -@pytest.mark.parametrize("num_tokens", [i + 1 for i in range(16)]) -@pytest.mark.parametrize("num_experts", [256, 384]) -def test_dsv3_router_gemm(num_tokens, num_experts): - hidden_dim = 7168 - - mat_a = torch.randn( - (num_tokens, hidden_dim), dtype=torch.bfloat16, device="cuda" - ).contiguous() - mat_b = torch.randn( - (num_experts, hidden_dim), dtype=torch.bfloat16, device="cuda" - ).contiguous() - - bf16_ref = F.linear(mat_a, mat_b) - float_ref = bf16_ref.to(torch.float32) - - bf16_output = dsv3_router_gemm(mat_a, mat_b, out_dtype=torch.bfloat16) - float_output = dsv3_router_gemm(mat_a, mat_b, out_dtype=torch.float32) - - assert torch.allclose( - bf16_output, bf16_ref, rtol=1e-2, atol=1e-3 - ), "Router GEMM output in bf16 dtype mismatch with torch.nn.functional.linear reference" - - assert torch.allclose( - float_output, float_ref, rtol=1e-2, atol=1e-3 - ), "Router GEMM output in float32 dtype mismatch with torch.nn.functional.linear reference" - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__])) diff --git a/test/registered/jit/benchmark/bench_dsv3_router_gemm.py b/test/registered/jit/benchmark/bench_dsv3_router_gemm.py new file mode 100644 index 000000000..b1e60de7c --- /dev/null +++ b/test/registered/jit/benchmark/bench_dsv3_router_gemm.py @@ -0,0 +1,64 @@ +"""Benchmark for DeepSeek V3 router GEMM (JIT kernel vs sgl_kernel AOT vs torch). + +Run on a Hopper (SM90+) GPU: + python -m sglang.jit_kernel.benchmark.bench_dsv3_router_gemm +""" + +import torch +import torch.nn.functional as F + +try: + from sgl_kernel import dsv3_router_gemm as sgl_kernel_dsv3_router_gemm +except ImportError: + sgl_kernel_dsv3_router_gemm = None + +from sglang.jit_kernel.benchmark import marker +from sglang.jit_kernel.benchmark.utils import create_random +from sglang.jit_kernel.dsv3_router_gemm import dsv3_router_gemm +from sglang.jit_kernel.utils import get_jit_cuda_arch, is_hip_runtime +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=5, suite="base-b-kernel-benchmark-1-gpu-large") + +# sgl_kernel AOT kernel is specialized for hidden_dim=7168 only. +SGL_KERNEL_HIDDEN_DIM = 7168 + + +def _torch(mat_a, mat_b, out_dtype): + return F.linear(mat_a, mat_b).to(out_dtype) + + +FN_MAP = { + "jit": dsv3_router_gemm, + "sgl_kernel": sgl_kernel_dsv3_router_gemm, + "torch": _torch, +} + + +@marker.parametrize("num_experts", [256, 384], [256]) +@marker.parametrize("hidden_dim", [6144, 7168], [7168]) +@marker.parametrize("num_tokens", list(range(1, 17)), [1, 8, 16]) +@marker.parametrize("out_dtype", [torch.bfloat16, torch.float32]) +@marker.benchmark("provider", ["jit", "sgl_kernel", "torch"]) +def benchmark(num_experts, hidden_dim, num_tokens, out_dtype, provider): + if provider == "sgl_kernel": + if sgl_kernel_dsv3_router_gemm is None: + marker.skip("sgl_kernel dsv3_router_gemm not available in this build") + if hidden_dim != SGL_KERNEL_HIDDEN_DIM: + marker.skip("sgl_kernel AOT only supports hidden_dim=7168") + mat_a = create_random(num_tokens, hidden_dim) + mat_b = create_random(num_experts, hidden_dim) + return marker.do_bench( + FN_MAP[provider], + input_args=(mat_a, mat_b), + input_kwargs={"out_dtype": out_dtype}, + ) + + +if __name__ == "__main__": + if is_hip_runtime() or get_jit_cuda_arch().major < 9: + print( + "dsv3_router_gemm JIT kernel requires SM90+ (Hopper). Skipping benchmark." + ) + else: + benchmark.run() diff --git a/test/registered/jit/test_dsv3_router_gemm.py b/test/registered/jit/test_dsv3_router_gemm.py new file mode 100644 index 000000000..c9638b7f8 --- /dev/null +++ b/test/registered/jit/test_dsv3_router_gemm.py @@ -0,0 +1,45 @@ +"""Tests for JIT dsv3_router_gemm kernel.""" + +import sys + +import pytest +import torch + +from sglang.jit_kernel.dsv3_router_gemm import dsv3_router_gemm +from sglang.jit_kernel.utils import get_jit_cuda_arch, is_hip_runtime +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=37, suite="base-b-kernel-unit-1-gpu-large") +register_cuda_ci(est_time=148, suite="nightly-kernel-1-gpu", nightly=True) + +HIDDEN_DIMS = [1024, 4096, 5120, 6144, 7168] +ATOL = 1e-2 +RTOL = 1e-2 + + +def _ref(hidden_states, router_weights, out_dtype): + return (hidden_states.float() @ router_weights.float().T).to(out_dtype) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") +@pytest.mark.parametrize("num_experts", [256, 384]) +@pytest.mark.parametrize("hidden_dim", HIDDEN_DIMS) +@pytest.mark.parametrize("num_tokens", list(range(1, 17))) +@pytest.mark.parametrize("out_dtype", [torch.bfloat16, torch.float32]) +def test_dsv3_router_gemm(num_experts, hidden_dim, num_tokens, out_dtype): + if is_hip_runtime() or get_jit_cuda_arch().major < 9: + pytest.skip("SM90+ required") + + mat_a = torch.randn(num_tokens, hidden_dim, dtype=torch.bfloat16, device="cuda") + mat_b = torch.randn(num_experts, hidden_dim, dtype=torch.bfloat16, device="cuda") + + ref = _ref(mat_a, mat_b, out_dtype) + out = dsv3_router_gemm(mat_a, mat_b, out_dtype=out_dtype) + + assert out.shape == (num_tokens, num_experts) + assert out.dtype == out_dtype + torch.testing.assert_close(out.float(), ref.float(), atol=ATOL, rtol=RTOL) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v"]))