[Kernel] Deprecate DeepGemm in sgl kernel and apply custom wheel sgl-deep-gemm (#24268)

This commit is contained in:
Baizhou Zhang
2026-05-06 18:59:01 -07:00
committed by GitHub
parent eaf074d50e
commit ecb786c8d7
13 changed files with 87 additions and 19 deletions
+2 -1
View File
@@ -59,7 +59,8 @@ dependencies = [
"sentencepiece",
"setproctitle",
"flash-attn-4>=4.0.0b9",
"sglang-kernel==0.4.2",
"sgl-deep-gemm==0.0.1",
"sglang-kernel==0.4.2.post1",
"soundfile==0.13.1",
"tiktoken",
"timm==1.0.16",
+1 -1
View File
@@ -1170,7 +1170,7 @@ def _set_envs_and_config(server_args: ServerArgs):
if _is_cuda:
assert_pkg_version(
"sglang-kernel",
"0.4.2",
"0.4.2.post1",
"Please reinstall the latest version with `pip install sglang-kernel --force-reinstall`",
)
@@ -456,10 +456,18 @@ class Indexer(MultiPlatformOp):
# Reuse pre-computed schedule metadata if available (from init_forward_metadata),
# otherwise fall back to computing it here.
schedule_metadata = getattr(metadata, "paged_mqa_schedule_metadata", None)
# DeepGEMM release-0426 requires context_lens of shape [batch_size, next_n]
# to match q.shape = [batch_size, next_n, heads, head_dim]. The indexer uses
# next_n=1 with batch_size=N_total via q_fp8.unsqueeze(1) below, so mirror
# that layout here.
if seqlens_32.dim() == 2:
seqlens_32_2d = seqlens_32
else:
seqlens_32_2d = seqlens_32.unsqueeze(-1)
if _is_cuda:
if schedule_metadata is None:
schedule_metadata = deep_gemm.get_paged_mqa_logits_metadata(
seqlens_32, blocksize, self.sm_count
seqlens_32_2d, blocksize, self.sm_count
)
assert len(q_fp8.shape) == 3
@@ -508,7 +516,7 @@ class Indexer(MultiPlatformOp):
q_fp8[:q_offset],
kv_cache_fp8,
weights[:q_offset],
seqlens_32,
seqlens_32_2d,
block_tables,
schedule_metadata,
max_seq_len,
@@ -67,6 +67,17 @@ else:
)
def _to_2d_context_lens(seqlens_32: torch.Tensor, batch_size: int) -> torch.Tensor:
if seqlens_32.dim() == 2:
return seqlens_32
n = seqlens_32.numel()
assert (
n % batch_size == 0
), f"seqlens_32 size {n} is not a multiple of batch_size {batch_size}"
next_n = n // batch_size
return seqlens_32.view(batch_size, next_n)
# Reuse this workspace buffer across all NSA backend instances
global_workspace_buffer = None
@@ -647,8 +658,11 @@ class NativeSparseAttnBackend(
)
else cache_seqlens_int32
)
seqlens_32_2d = _to_2d_context_lens(
seqlens_32, forward_batch.batch_size
)
paged_mqa_schedule_metadata = deep_gemm.get_paged_mqa_logits_metadata(
seqlens_32, 64, deep_gemm.get_num_sms()
seqlens_32_2d, 64, deep_gemm.get_num_sms()
)
except (ImportError, ModuleNotFoundError):
paged_mqa_schedule_metadata = None
@@ -932,8 +946,9 @@ class NativeSparseAttnBackend(
)
else cache_seqlens_int32
)
seqlens_32_2d = _to_2d_context_lens(seqlens_32, bs)
paged_mqa_schedule_metadata = deep_gemm.get_paged_mqa_logits_metadata(
seqlens_32, 64, deep_gemm.get_num_sms()
seqlens_32_2d, 64, deep_gemm.get_num_sms()
)
except (ImportError, ModuleNotFoundError):
paged_mqa_schedule_metadata = None
@@ -1082,8 +1097,9 @@ class NativeSparseAttnBackend(
)
else metadata.cache_seqlens_int32
)
seqlens_32_2d = _to_2d_context_lens(seqlens_32, bs)
new_schedule = deep_gemm.get_paged_mqa_logits_metadata(
seqlens_32, 64, deep_gemm.get_num_sms()
seqlens_32_2d, 64, deep_gemm.get_num_sms()
)
if metadata.paged_mqa_schedule_metadata is None:
metadata.paged_mqa_schedule_metadata = new_schedule
@@ -196,12 +196,18 @@ def _compile_deep_gemm_one_type_all(
kernel_type, max_m=max_m, n=n, k=k, num_groups=num_groups
)
old_compile_mode = deep_gemm.get_compile_mode()
deep_gemm.set_compile_mode(1)
has_compile_mode_api = hasattr(deep_gemm, "get_compile_mode") and hasattr(
deep_gemm, "set_compile_mode"
)
if has_compile_mode_api:
old_compile_mode = deep_gemm.get_compile_mode()
deep_gemm.set_compile_mode(1)
# TODO can use multi thread
for m in tqdm(m_list, desc=f"DeepGEMM warmup"):
executor.execute(m=m)
deep_gemm.set_compile_mode(old_compile_mode)
if has_compile_mode_api:
deep_gemm.set_compile_mode(old_compile_mode)
# clean up input buffers
torch.cuda.current_stream().synchronize()
@@ -300,7 +306,7 @@ class _GroupedContWarmupExecutor(_BaseWarmupExecutor):
(self.lhs_q[:m], self.lhs_s[:m]),
(self.rhs_q, self.rhs_s),
self.out[:m],
m_indices=self.m_indices[:m],
self.m_indices[:m],
)
@@ -18,6 +18,8 @@ def _compute_enable_deep_gemm():
sm_version = get_device_sm()
if (_is_cuda and sm_version < 90) or (_is_musa and sm_version < 31):
return False
if not (_is_cuda or _is_musa):
return False
try:
import deep_gemm # noqa: F401
@@ -1286,6 +1286,19 @@ def transform_scale_ue8m0(sf, mn, use_torch_impl: bool = False):
sf = sf.index_select(-2, torch.arange(mn, device=sf.device) // 128)
sf = get_mn_major_tma_aligned_packed_ue8m0_tensor(sf)
# In sgl-deep-gemm, the C++ deepgemm path returns through DLPack which collapses the stride
# of size-1 trailing dims to 1 (happens when packed_sf_k == 1, i.e.
# K <= block_k * 4). Restore the TMA-aligned stride so the deepgemm
# assertion sf.stride(-1) == get_tma_aligned_size(mn, element_size) holds.
if not use_torch_impl and sf.shape[-1] == 1:
from deep_gemm.utils import get_tma_aligned_size
aligned_mn = get_tma_aligned_size(sf.shape[-2], sf.element_size())
if sf.stride(-1) != aligned_mn:
new_stride = list(sf.stride())
new_stride[-1] = aligned_mn
sf = sf.as_strided(sf.shape, tuple(new_stride))
return sf