[fix] fix DeepGEMM blackwell input quant & ut & fix style and log (#7247)

This commit is contained in:
JieXin Liang
2025-06-16 11:45:54 -07:00
committed by GitHub
parent e30ef368ab
commit 5ca07eed90
10 changed files with 285 additions and 31 deletions
+2 -2
View File
@@ -1201,7 +1201,7 @@ class DeepEPMoE(EPMoE):
gateup_output,
masked_m,
expected_m,
recipe=(1, 128, 128) if deep_gemm_wrapper.DEEPGEMM_V202506 else None,
recipe=(1, 128, 128) if deep_gemm_wrapper.DEEPGEMM_BLACKWELL else None,
)
dispose_tensor(hidden_states_fp8[0])
@@ -1256,7 +1256,7 @@ class DeepEPMoE(EPMoE):
down_output,
masked_m,
expected_m,
recipe=(1, 128, 128) if deep_gemm_wrapper.DEEPGEMM_V202506 else None,
recipe=(1, 128, 128) if deep_gemm_wrapper.DEEPGEMM_BLACKWELL else None,
)
return down_output
@@ -553,9 +553,9 @@ class _DeepEPDispatcherImplLowLatency(_DeepEPDispatcherImplBase):
async_finish=not self.return_recv_hook,
return_recv_hook=self.return_recv_hook,
round_scale=deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM
and deep_gemm_wrapper.DEEPGEMM_V202506,
and deep_gemm_wrapper.DEEPGEMM_BLACKWELL,
use_ue8m0=deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM
and deep_gemm_wrapper.DEEPGEMM_V202506,
and deep_gemm_wrapper.DEEPGEMM_BLACKWELL,
)
)
return packed_recv_hidden, packed_recv_count, event, hook
@@ -8,7 +8,7 @@ from typing import Callable, Dict, List, Optional, Tuple
from tqdm.contrib.concurrent import thread_map
from sglang.srt.layers.quantization.deep_gemm_wrapper.configurer import (
DEEPGEMM_V202506,
DEEPGEMM_BLACKWELL,
ENABLE_JIT_DEEPGEMM,
)
from sglang.srt.server_args import ServerArgs
@@ -16,13 +16,11 @@ from sglang.srt.utils import get_bool_env_var, get_int_env_var
logger = logging.getLogger(__name__)
try:
if ENABLE_JIT_DEEPGEMM and not DEEPGEMM_BLACKWELL:
from deep_gemm import get_num_sms
from deep_gemm.jit import build
from deep_gemm.jit_kernels.gemm import get_best_configs
from deep_gemm.jit_kernels.runtime import FP8GemmRuntime, GemmType
except ImportError:
pass
_BUILTIN_M_LIST = list(range(1, 1024 * 16 + 1))
@@ -313,7 +311,8 @@ def _log_jit_build(M: int, N: int, K: int, kernel_type: DeepGemmKernelType):
ret = origin_func(self, *args, **kwargs)
if ret is None:
kernel_helper = _KERNEL_HELPER_DICT[kernel_type]
_compile_warning_2()
if not DEEPGEMM_BLACKWELL:
_compile_warning_2()
logger.warning(
f"DeepGEMM JIT Compiling for <{kernel_helper.name}> M={M}, N={N}, K={K}. Please wait."
)
@@ -329,10 +328,8 @@ def deep_gemm_execution_hook(
m: int, n: int, k: int, num_groups: int, kernel_type: DeepGemmKernelType
):
# not supported yet
if DEEPGEMM_V202506:
yield
return
if not DEEPGEMM_BLACKWELL:
_maybe_compile_deep_gemm_one_type_all(kernel_type, n, k, num_groups)
_maybe_compile_deep_gemm_one_type_all(kernel_type, n, k, num_groups)
with _log_jit_build(m, n, k, kernel_type):
yield
@@ -6,16 +6,16 @@ logger = logging.getLogger(__name__)
def _compute_enable_deep_gemm():
sm_version = get_device_sm()
if sm_version < 90:
return False
try:
import deep_gemm
except ImportError:
logger.warning("Failed to import deep_gemm, disable ENABLE_JIT_DEEPGEMM.")
return False
sm_version = get_device_sm()
if sm_version < 90:
return False
return get_bool_env_var("SGL_ENABLE_JIT_DEEPGEMM", default="true")
@@ -25,8 +25,8 @@ try:
from deep_gemm import fp8_gemm_nt
# They have not given a name to this breaking change
DEEPGEMM_V202506 = True
DEEPGEMM_BLACKWELL = True
except ImportError:
DEEPGEMM_V202506 = False
DEEPGEMM_BLACKWELL = False
DEEPGEMM_SCALE_UE8M0 = DEEPGEMM_V202506
DEEPGEMM_SCALE_UE8M0 = DEEPGEMM_BLACKWELL
@@ -6,8 +6,8 @@ import torch
from sglang.srt.layers.quantization.deep_gemm_wrapper import compile_utils
from sglang.srt.layers.quantization.deep_gemm_wrapper.configurer import (
DEEPGEMM_BLACKWELL,
DEEPGEMM_SCALE_UE8M0,
DEEPGEMM_V202506,
ENABLE_JIT_DEEPGEMM,
)
from sglang.srt.server_args import ServerArgs
@@ -17,7 +17,7 @@ logger = logging.getLogger(__name__)
if ENABLE_JIT_DEEPGEMM:
import deep_gemm
if DEEPGEMM_V202506:
if DEEPGEMM_BLACKWELL:
from deep_gemm import fp8_gemm_nt as _gemm_nt_f8f8bf16_raw
from deep_gemm import (
fp8_m_grouped_gemm_nt_masked as _grouped_gemm_nt_f8f8bf16_masked_raw,
@@ -57,7 +57,7 @@ def grouped_gemm_nt_f8f8bf16_masked(
out,
masked_m,
expected_m,
**({"recipe": recipe} if DEEPGEMM_V202506 else {})
**({"recipe": recipe} if DEEPGEMM_BLACKWELL else {})
)
@@ -290,11 +290,12 @@ def sglang_per_token_group_quant_fp8(
x_s_mn, x_s_k = x_q_mn, x_q_k // 128
aligned_mn = align(x_s_mn, 4)
aligned_k = align(x_s_k, 4)
x_s = torch.empty(
# TODO(FIXME): Fix cuda kernel and recover here to empty.
x_s = torch.zeros(
(aligned_k // 4, aligned_mn),
device=x.device,
dtype=torch.int,
).permute(-1, -2)[:x_s_mn, :]
).transpose(0, 1)[:x_s_mn, :]
elif column_major_scales:
if scale_tma_aligned:
# TODO extract "align" function
@@ -768,7 +769,7 @@ def prepare_block_fp8_matmul_inputs(
if As.dtype == torch.float:
assert triton.cdiv(A.shape[-1], block_k) == As.shape[-1]
elif Bs.dtype == torch.int:
elif As.dtype == torch.int:
assert (
triton.cdiv(triton.cdiv(A.shape[-1], block_k), 4) == As.shape[-1]
), f"{A.shape=} {As.shape=} {block_size=}"
@@ -241,9 +241,10 @@ def deepgemm_w8a8_block_fp8_linear_with_fallback(
scale_ue8m0=deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0,
)
if get_bool_env_var("SGLANG_W8A8_DEEPGEMM_SANITY_CHECK_UE8M0"):
_check_ue8m0("x_scale", x_scale)
_check_ue8m0("weight_scale", weight_scale)
# NOTE(alcanderian): Useless when scale is packed to int32
# if get_bool_env_var("SGLANG_W8A8_DEEPGEMM_SANITY_CHECK_UE8M0"):
# _check_ue8m0("x_scale", x_scale)
# _check_ue8m0("weight_scale", ws)
output = w8a8_block_fp8_matmul_deepgemm(
q_input, weight, x_scale, weight_scale, block_size, output_dtype=output_dtype
+4 -2
View File
@@ -1829,8 +1829,10 @@ class DeepseekV2ForCausalLM(nn.Module):
and weight_block_size[1] == 128
and model_dtype == torch.bfloat16
):
if deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM and get_bool_env_var(
"SGL_USE_DEEPGEMM_BMM", "false"
if (
deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM
and not deep_gemm_wrapper.DEEPGEMM_BLACKWELL
and get_bool_env_var("SGL_USE_DEEPGEMM_BMM", "false")
):
block_scale = weight_scale
use_deep_gemm_bmm = True