From 75427c9ca426326f3b65ee97b1d13129f67ca5d2 Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <1182563586@qq.com> Date: Sat, 23 May 2026 14:30:43 +0800 Subject: [PATCH] Route concat MLA to JIT and remove unused downcast (#25843) --- .../sglang/jit_kernel/benchmark/bench_cast.py | 106 -------------- python/sglang/jit_kernel/cast.py | 52 ------- .../jit_kernel/csrc/elementwise/cast.cuh | 137 ------------------ python/sglang/srt/layers/attention/utils.py | 2 +- .../attention_forward_methods/forward_mha.py | 6 +- python/sglang/srt/models/sarvam_moe.py | 3 +- 6 files changed, 8 insertions(+), 298 deletions(-) delete mode 100644 python/sglang/jit_kernel/benchmark/bench_cast.py delete mode 100644 python/sglang/jit_kernel/cast.py delete mode 100644 python/sglang/jit_kernel/csrc/elementwise/cast.cuh diff --git a/python/sglang/jit_kernel/benchmark/bench_cast.py b/python/sglang/jit_kernel/benchmark/bench_cast.py deleted file mode 100644 index 4be874ce5..000000000 --- a/python/sglang/jit_kernel/benchmark/bench_cast.py +++ /dev/null @@ -1,106 +0,0 @@ -import torch -import triton -import triton.testing - -from sglang.jit_kernel.benchmark.utils import ( - DEFAULT_DEVICE, - get_benchmark_range, - run_benchmark, -) -from sglang.jit_kernel.cast import downcast_fp8 as downcast_fp8_jit -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=10, suite="base-b-kernel-benchmark-1-gpu-large") - -DEVICE = DEFAULT_DEVICE -DTYPE = torch.bfloat16 - - -# ── Config ranges ────────────────────────────────────────────────────────────── - -SL_LIST = get_benchmark_range( - full_range=[4, 16, 64, 256, 512, 1024, 2048], - ci_range=[4, 64], -) - -HEAD_DIM_LIST = get_benchmark_range( - full_range=[(8, 128), (32, 128), (8, 256), (32, 256)], - ci_range=[(8, 128)], -) - -CONFIGS = [(sl, h, d, sl * 2) for sl in SL_LIST for h, d in HEAD_DIM_LIST] - -LINE_VALS = ["jit"] -LINE_NAMES = ["JIT (cast.cuh, 256 threads, 2D grid)"] -STYLES = [("orange", "-")] - - -# ── Perf report ──────────────────────────────────────────────────────────────── - - -@triton.testing.perf_report( - triton.testing.Benchmark( - x_names=["input_sl", "head", "dim", "out_sl"], - x_vals=CONFIGS, - line_arg="provider", - line_vals=LINE_VALS, - line_names=LINE_NAMES, - styles=STYLES, - ylabel="us", - plot_name="downcast-fp8-jit", - args={}, - ) -) -def benchmark(input_sl, head, dim, out_sl, provider): - k = torch.randn(input_sl, head, dim, dtype=DTYPE, device=DEVICE) - v = torch.randn(input_sl, head, dim, dtype=DTYPE, device=DEVICE) - k_out = torch.zeros(out_sl, head, dim, dtype=torch.uint8, device=DEVICE) - v_out = torch.zeros(out_sl, head, dim, dtype=torch.uint8, device=DEVICE) - k_scale = torch.tensor([1.0], dtype=torch.float32, device=DEVICE) - v_scale = torch.tensor([1.0], dtype=torch.float32, device=DEVICE) - loc = torch.arange(input_sl, dtype=torch.int64, device=DEVICE) - - fn = lambda: downcast_fp8_jit(k, v, k_out, v_out, k_scale, v_scale, loc) - - return run_benchmark(fn) - - -# ── Bandwidth analysis ───────────────────────────────────────────────────────── - - -def _report_bandwidth(input_sl, head, dim, dtype): - elem_bytes = torch.finfo(dtype).bits // 8 - total_bytes = input_sl * head * dim * (2 * elem_bytes + 2) - - k = torch.randn(input_sl, head, dim, dtype=dtype, device=DEVICE) - v = torch.randn(input_sl, head, dim, dtype=dtype, device=DEVICE) - k_out = torch.zeros(input_sl * 2, head, dim, dtype=torch.uint8, device=DEVICE) - v_out = torch.zeros(input_sl * 2, head, dim, dtype=torch.uint8, device=DEVICE) - k_scale = torch.tensor([1.0], dtype=torch.float32, device=DEVICE) - v_scale = torch.tensor([1.0], dtype=torch.float32, device=DEVICE) - loc = torch.arange(input_sl, dtype=torch.int64, device=DEVICE) - - jit_fn = lambda: downcast_fp8_jit(k, v, k_out, v_out, k_scale, v_scale, loc) - - jit_ms, _, _ = triton.testing.do_bench(jit_fn, quantiles=[0.5, 0.2, 0.8]) - - def fmt(ms): - return f"{ms*1000:6.2f}us {total_bytes/(ms*1e-3)/1e9:6.0f}GB/s" - - print(f" sl={input_sl:5d} h={head:2d} d={dim:4d}" f" | jit {fmt(jit_ms)}") - - -def report_bandwidth(): - print(f"\n{'='*95}") - print(" JIT (cast.cuh, 256 threads, 2D grid)") - print(f" dtype={DTYPE}, device={DEVICE}") - print(f"{'='*95}") - for sl in [64, 256, 1024, 2048]: - for h, d in [(8, 128), (32, 128), (8, 256), (32, 256)]: - _report_bandwidth(sl, h, d, DTYPE) - print() - - -if __name__ == "__main__": - benchmark.run(print_data=True) - report_bandwidth() diff --git a/python/sglang/jit_kernel/cast.py b/python/sglang/jit_kernel/cast.py deleted file mode 100644 index f0201c4ab..000000000 --- a/python/sglang/jit_kernel/cast.py +++ /dev/null @@ -1,52 +0,0 @@ -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_cast_module(dtype: torch.dtype) -> Module: - args = make_cpp_args(dtype) - return load_jit( - "cast", - *args, - cuda_files=["elementwise/cast.cuh"], - cuda_wrappers=[("downcast_fp8", f"downcast_fp8<{args}>")], - ) - - -def downcast_fp8( - k: torch.Tensor, - v: torch.Tensor, - k_out: torch.Tensor, - v_out: torch.Tensor, - k_scale: torch.Tensor, - v_scale: torch.Tensor, - loc: torch.Tensor, - mult: int = 1, - offset: int = 0, -) -> None: - """Fused downcast of KV cache tensors from bf16/fp16 to fp8 (E4M3). - - Scales each value by the inverse of its per-tensor scale, clamps to the - fp8 representable range [-448, 448], then converts to fp8 storage. - - Args: - k: [input_sl, head, dim] bf16/fp16 CUDA tensor - v: [input_sl, head, dim] bf16/fp16 CUDA tensor - k_out: [out_sl, head, dim] uint8 CUDA tensor (fp8 storage) - v_out: [out_sl, head, dim] uint8 CUDA tensor (fp8 storage) - k_scale: [1] float32 CUDA tensor, scale for k - v_scale: [1] float32 CUDA tensor, scale for v - loc: [input_sl] int64 CUDA tensor, destination sequence indices - mult: stride multiplier for output index (default 1) - offset: offset added to output index (default 0) - """ - module = _jit_cast_module(k.dtype) - module.downcast_fp8(k, v, k_out, v_out, k_scale, v_scale, loc, mult, offset) diff --git a/python/sglang/jit_kernel/csrc/elementwise/cast.cuh b/python/sglang/jit_kernel/csrc/elementwise/cast.cuh deleted file mode 100644 index f537ddc58..000000000 --- a/python/sglang/jit_kernel/csrc/elementwise/cast.cuh +++ /dev/null @@ -1,137 +0,0 @@ -#pragma once - -// Optimized cast kernel: fixed 256 threads, scaled out via 2D grid. -// Each thread handles exactly one float4 (kVecSize fp16/bf16 elements). -// No per-thread loop — pure grid scaling for any head*dim. - -#include -#include - -#include // For dtype_trait fp8 specialization -#include // For LaunchKernel -#include // For AlignedVector - -#include -#include - -#include - -namespace { - -constexpr int kBlockSize = 256; - -template -__global__ void fused_downcast_kernel( - const T* __restrict__ cache_k, - const T* __restrict__ cache_v, - const float* __restrict__ k_scale, - const float* __restrict__ v_scale, - fp8_e4m3_t* __restrict__ output_k, - fp8_e4m3_t* __restrict__ output_v, - const int input_num_tokens, - const int head, - const int dim, - const T max_fp8, - const T min_fp8, - const int64_t mult, - const int64_t offset, - const int64_t* __restrict__ loc) { - using namespace device; - - constexpr int kVecSize = 16 / sizeof(T); - using vec_t = AlignedVector; - using out_vec_t = AlignedVector; - - const int token_idx = blockIdx.x; - const int vec_idx = blockIdx.y * kBlockSize + threadIdx.x; - const int num_vecs = head * dim / kVecSize; - - if (token_idx >= input_num_tokens || vec_idx >= num_vecs) return; - - T k_scale_inv = static_cast(1.f) / cast(k_scale[0]); - T v_scale_inv = static_cast(1.f) / cast(v_scale[0]); - - auto clamp = [&](T val) { return val > max_fp8 ? max_fp8 : (min_fp8 > val ? min_fp8 : val); }; - - const int out_seq_idx = loc[token_idx]; - const T* in_k_base = cache_k + token_idx * head * dim; - const T* in_v_base = cache_v + token_idx * head * dim; - fp8_e4m3_t* out_k_base = output_k + (out_seq_idx * mult + offset) * head * dim; - fp8_e4m3_t* out_v_base = output_v + (out_seq_idx * mult + offset) * head * dim; - - vec_t k_vec, v_vec; - k_vec.load(in_k_base, vec_idx); - v_vec.load(in_v_base, vec_idx); - - out_vec_t out_k, out_v; -#pragma unroll - for (int j = 0; j < kVecSize; j++) { - out_k[j] = cast(clamp(k_vec[j] * k_scale_inv)); - out_v[j] = cast(clamp(v_vec[j] * v_scale_inv)); - } - - out_k.store(out_k_base, vec_idx); - out_v.store(out_v_base, vec_idx); -} - -template -void downcast_fp8( - tvm::ffi::TensorView k, - tvm::ffi::TensorView v, - tvm::ffi::TensorView k_out, - tvm::ffi::TensorView v_out, - tvm::ffi::TensorView k_scale, - tvm::ffi::TensorView v_scale, - tvm::ffi::TensorView loc, - int64_t mult, - int64_t offset) { - using namespace host; - - auto input_num_tokens = SymbolicSize{"input_num_tokens"}; - auto head = SymbolicSize{"head"}; - auto dim = SymbolicSize{"dim"}; - auto output_num_tokens = SymbolicSize{"out_sl"}; - auto device = SymbolicDevice{}; - device.set_options(); - - TensorMatcher({input_num_tokens, head, dim}).with_dtype().with_device(device).verify(k); - TensorMatcher({input_num_tokens, head, dim}).with_dtype().with_device(device).verify(v); - TensorMatcher({output_num_tokens, head, dim}).with_dtype().with_device(device).verify(k_out); - TensorMatcher({output_num_tokens, head, dim}).with_dtype().with_device(device).verify(v_out); - TensorMatcher({1}).with_dtype().with_device(device).verify(k_scale); - TensorMatcher({1}).with_dtype().with_device(device).verify(v_scale); - TensorMatcher({input_num_tokens}).with_dtype().with_device(device).verify(loc); - - const int num_tokens = static_cast(input_num_tokens.unwrap()); - const int h = static_cast(head.unwrap()); - const int d = static_cast(dim.unwrap()); - - constexpr int kVecSize = 16 / sizeof(T); - const int num_vecs = h * d / kVecSize; - const int grid_y = (num_vecs + kBlockSize - 1) / kBlockSize; - - dim3 grid(num_tokens, grid_y); - dim3 block(kBlockSize); - - const T max_fp8 = static_cast(kFP8E4M3Max); - const T min_fp8 = static_cast(-kFP8E4M3Max); - - LaunchKernel(grid, block, device.unwrap())( - fused_downcast_kernel, - static_cast(k.data_ptr()), - static_cast(v.data_ptr()), - static_cast(k_scale.data_ptr()), - static_cast(v_scale.data_ptr()), - static_cast(k_out.data_ptr()), - static_cast(v_out.data_ptr()), - num_tokens, - h, - d, - max_fp8, - min_fp8, - mult, - offset, - static_cast(loc.data_ptr())); -} - -} // namespace diff --git a/python/sglang/srt/layers/attention/utils.py b/python/sglang/srt/layers/attention/utils.py index 277d46054..65328b16b 100644 --- a/python/sglang/srt/layers/attention/utils.py +++ b/python/sglang/srt/layers/attention/utils.py @@ -10,7 +10,7 @@ FLASHMLA_CREATE_KV_BLOCK_SIZE_TRITON = tl.constexpr(_FLASHMLA_CREATE_KV_BLOCK_SI _is_cuda = is_cuda() if _is_cuda: - from sgl_kernel import concat_mla_absorb_q + from sglang.jit_kernel.concat_mla import concat_mla_absorb_q from sglang.jit_kernel.utils import is_arch_support_pdl diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py index 67f4483cd..dbcb3ee0f 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py @@ -32,7 +32,11 @@ if TYPE_CHECKING: from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA if _is_cuda: - from sgl_kernel import concat_mla_k, merge_state_v2 + from sgl_kernel import merge_state_v2 + + from sglang.jit_kernel.concat_mla import concat_mla_k +elif _is_musa: + from sgl_kernel import concat_mla_k if _use_aiter_gfx95: from aiter.ops.triton.fused_fp8_quant import fused_rms_fp8_group_quant diff --git a/python/sglang/srt/models/sarvam_moe.py b/python/sglang/srt/models/sarvam_moe.py index 6ef737727..83683933c 100644 --- a/python/sglang/srt/models/sarvam_moe.py +++ b/python/sglang/srt/models/sarvam_moe.py @@ -79,8 +79,9 @@ _is_cublas_ge_129 = is_nvidia_cublas_version_ge_12_9() if _is_cuda: try: - from sgl_kernel import bmm_fp8, concat_mla_k, merge_state_v2 + from sgl_kernel import bmm_fp8, merge_state_v2 + from sglang.jit_kernel.concat_mla import concat_mla_k from sglang.srt.layers.quantization.fp8_kernel import per_tensor_quant_mla_fp8 _has_fp8_support = True