diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend.py b/python/sglang/srt/layers/attention/deepseek_v4_backend.py index 659ae1438..51b7a091d 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend.py @@ -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, ) diff --git a/python/sglang/srt/layers/attention/dsv4/metadata.py b/python/sglang/srt/layers/attention/dsv4/metadata.py index 2b798ab28..d245ddce3 100644 --- a/python/sglang/srt/layers/attention/dsv4/metadata.py +++ b/python/sglang/srt/layers/attention/dsv4/metadata.py @@ -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, ) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 0e408cf18..715d79604 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -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 diff --git a/test/registered/kernels/benchmark/attention/bench_dsv4_fp4_indexer.py b/test/registered/kernels/benchmark/attention/bench_dsv4_fp4_indexer.py index 1e11568a1..957633540 100644 --- a/test/registered/kernels/benchmark/attention/bench_dsv4_fp4_indexer.py +++ b/test/registered/kernels/benchmark/attention/bench_dsv4_fp4_indexer.py @@ -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.") diff --git a/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py b/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py index 3de07a93c..d13cd02ee 100644 --- a/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py +++ b/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py @@ -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)