Add FP4 Indexer for DeepSeek V4 on SM120 (#27059)
This commit is contained in:
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user