Add FP4 Indexer for DeepSeek V4 on SM120 (#27059)

This commit is contained in:
Jinyan Chen
2026-07-24 11:37:23 -07:00
committed by GitHub
parent 2428f56145
commit 1e69765bae
5 changed files with 76 additions and 10 deletions
@@ -650,6 +650,11 @@ class DeepseekV4AttnBackend(
page_size=self.page_size,
page_table=core_attn_metadata.page_table,
c4_seq_lens=core_attn_metadata.c4_topk_lengths_raw,
# The SM120 FP4 kernel schedules split_kv=128, while the generic
# JIT metadata planner encodes split_kv=256.
force_deep_gemm_metadata=(
self.enable_deepseek_v4_fp4_indexer and _is_sm120
),
use_prefill_cuda_graph=use_prefill_cuda_graph,
)
@@ -112,6 +112,7 @@ class PagedIndexerMetadata:
page_size: int
page_table: torch.Tensor
c4_seq_lens: torch.Tensor
force_deep_gemm_metadata: bool = False
use_prefill_cuda_graph: bool = False
deep_gemm_metadata: Any = field(init=False, repr=False)
topk_metadata: torch.Tensor = field(init=False, repr=False)
@@ -124,12 +125,12 @@ class PagedIndexerMetadata:
envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.get()
or is_xpu()
or envs.SGLANG_OPT_USE_AITER_INDEXER.get()
):
) and not self.force_deep_gemm_metadata:
self.deep_gemm_metadata = None
else:
import deep_gemm
use_jit_indexer = (
use_jit_indexer = not self.force_deep_gemm_metadata and (
envs.SGLANG_OPT_USE_JIT_INDEXER_METADATA.get()
or self.c4_seq_lens.numel() > _LARGE_INDEXER_QUERY_THRESHOLD
)
@@ -183,7 +184,11 @@ class PagedIndexerMetadata:
copy_metadata(
src=other,
dst=self,
check_eq_fields=["page_size", "use_prefill_cuda_graph"],
check_eq_fields=[
"page_size",
"force_deep_gemm_metadata",
"use_prefill_cuda_graph",
],
copy_fields=copy_fields,
assign_fields=assign_fields,
)
+4 -2
View File
@@ -7263,9 +7263,11 @@ class ServerArgs:
"Debug mode for CUDA graph is enabled via breakable CUDA graph. "
"All operations will run eagerly through the graph capture/replay path."
)
if self.enable_deepseek_v4_fp4_indexer and not is_sm100_supported():
if self.enable_deepseek_v4_fp4_indexer and not (
is_sm100_supported() or is_sm120_supported()
):
raise ValueError(
"--enable-deepseek-v4-fp4-indexer requires SM100 GPUs with "
"--enable-deepseek-v4-fp4-indexer requires SM100 or SM120 GPUs with "
"DeepGEMM FP4 indexer support."
)
# FP8 W_o GEMM needs DeepGEMM JIT. Enable exactly where the runtime can run
@@ -7,11 +7,14 @@ import triton
from sglang.benchmark.bench_utils import run_bench
from sglang.kernels.jit.benchmark.utils import get_benchmark_range
from sglang.srt.utils import is_sm100_supported
from sglang.srt.utils import is_sm100_supported, is_sm120_supported
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
est_time=5, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
est_time=5,
stage="base-b-kernel-benchmark",
runner_config="1-gpu-large",
disabled="Temporarily disabled until DeepGEMM is ready",
)
try:
@@ -167,8 +170,8 @@ def benchmark(batch: int, seq_len_kv: int, provider: str):
if __name__ == "__main__":
if not is_sm100_supported():
print("[skip] DeepSeek V4 FP4 indexer benchmark requires SM100 CUDA.")
if not (is_sm100_supported() or is_sm120_supported()):
print("[skip] DeepSeek V4 FP4 indexer benchmark requires SM100 or SM120 CUDA.")
sys.exit(0)
if deep_gemm is None or per_token_cast_to_fp4 is None:
print("[skip] DeepGEMM is unavailable.")
@@ -7,7 +7,10 @@ import torch
from sglang.srt.environ import envs
from sglang.srt.layers.attention.dsv4.indexer import FP8_DTYPE, C4IndexerBackendMixin
from sglang.srt.layers.attention.dsv4.metadata import NonPagedIndexerPlan
from sglang.srt.layers.attention.dsv4.metadata import (
NonPagedIndexerPlan,
PagedIndexerMetadata,
)
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.runtime_context import get_parallel
from sglang.test.ci.ci_register import register_cpu_ci
@@ -18,6 +21,54 @@ register_cpu_ci(est_time=2, suite="base-a-test-cpu")
_INDEXER = "sglang.srt.layers.attention.dsv4.indexer"
class TestDSV4PagedIndexerMetadata(CustomTestCase):
def test_sm120_fp4_forces_deep_gemm_metadata(self):
expected = torch.tensor([[0, 0], [1, 0]], dtype=torch.int32)
deep_gemm = SimpleNamespace(
get_num_sms=MagicMock(return_value=1),
get_paged_mqa_logits_metadata=MagicMock(return_value=expected),
)
with (
patch.dict(sys.modules, {"deep_gemm": deep_gemm}),
envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.override(True),
envs.SGLANG_OPT_USE_AITER_INDEXER.override(False),
envs.SGLANG_OPT_USE_JIT_INDEXER_METADATA.override(True),
envs.SGLANG_OPT_USE_TOPK_V2.override(False),
patch(
"sglang.kernels.ops.attention.dsv4.get_paged_mqa_logits_metadata"
) as jit_metadata,
):
metadata = PagedIndexerMetadata(
page_size=256,
page_table=torch.zeros((1, 1), dtype=torch.int32),
c4_seq_lens=torch.tensor([65], dtype=torch.int32),
force_deep_gemm_metadata=True,
)
self.assertIs(metadata.deep_gemm_metadata, expected)
deep_gemm.get_num_sms.assert_called_once_with()
deep_gemm.get_paged_mqa_logits_metadata.assert_called_once()
args = deep_gemm.get_paged_mqa_logits_metadata.call_args.args
torch.testing.assert_close(args[0], torch.tensor([[65]], dtype=torch.int32))
self.assertEqual(args[1:], (64, 1))
jit_metadata.assert_not_called()
def test_sm120_fp8_torch_fallback_keeps_metadata_none(self):
with (
envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.override(True),
envs.SGLANG_OPT_USE_AITER_INDEXER.override(False),
envs.SGLANG_OPT_USE_TOPK_V2.override(False),
):
metadata = PagedIndexerMetadata(
page_size=256,
page_table=torch.zeros((1, 1), dtype=torch.int32),
c4_seq_lens=torch.tensor([65], dtype=torch.int32),
)
self.assertIsNone(metadata.deep_gemm_metadata)
class TestDSV4NonPagedIndexer(CustomTestCase):
def _is_eligible(self, **overrides):
backend = SimpleNamespace(hisparse_coordinator=None)