From ce7141ef98767f756ccc949f5716f3633afad01e Mon Sep 17 00:00:00 2001 From: BingjiaWang Date: Thu, 21 May 2026 06:24:10 +0800 Subject: [PATCH] add git gemm warpper for dispatch_bf16_fp32_backend (#25860) --- python/sglang/jit_kernel/deepseek_v4.py | 15 +++++++-------- 1 file changed, 7 insertions(+), 8 deletions(-) diff --git a/python/sglang/jit_kernel/deepseek_v4.py b/python/sglang/jit_kernel/deepseek_v4.py index 9f64d4ae2..1216dcc04 100644 --- a/python/sglang/jit_kernel/deepseek_v4.py +++ b/python/sglang/jit_kernel/deepseek_v4.py @@ -21,9 +21,13 @@ _use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip if _use_aiter: from aiter.tuned_gemm import tgemm +from sglang.srt.layers import deep_gemm_wrapper + if TYPE_CHECKING: from tvm_ffi.module import Module +_linear_bf16_fp32_algo = envs.SGLANG_OPT_BF16_FP32_GEMM_ALGO.get() + def make_name(name: str) -> str: return f"dpsk_v4_{name}" @@ -1013,10 +1017,7 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { def linear_bf16_fp32(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor: - from sglang.srt.environ import envs - - algo = envs.SGLANG_OPT_BF16_FP32_GEMM_ALGO.get() - return _dispatch_bf16_fp32_backend(x, y, algo=algo) + return _dispatch_bf16_fp32_backend(x, y, algo=_linear_bf16_fp32_algo) def _dispatch_bf16_fp32_backend( @@ -1026,10 +1027,8 @@ def _dispatch_bf16_fp32_backend( module = _jit_torch_cublas_bf16_fp32() return module.linear_bf16_fp32(x, y) elif algo == "deep_gemm": - import deep_gemm - - z = x.new_empty(x.size(0), y.size(0), dtype=torch.float32) - deep_gemm.bf16_gemm_nt(x, y, z) + z = torch.empty(x.size(0), y.size(0), dtype=torch.float32, device=x.device) + deep_gemm_wrapper.gemm_nt_bf16bf16f32(x, y, z) return z elif _use_aiter: return tgemm.mm(x, y, otype=torch.float32)