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_size=self.page_size,
|
||||||
page_table=core_attn_metadata.page_table,
|
page_table=core_attn_metadata.page_table,
|
||||||
c4_seq_lens=core_attn_metadata.c4_topk_lengths_raw,
|
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,
|
use_prefill_cuda_graph=use_prefill_cuda_graph,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -112,6 +112,7 @@ class PagedIndexerMetadata:
|
|||||||
page_size: int
|
page_size: int
|
||||||
page_table: torch.Tensor
|
page_table: torch.Tensor
|
||||||
c4_seq_lens: torch.Tensor
|
c4_seq_lens: torch.Tensor
|
||||||
|
force_deep_gemm_metadata: bool = False
|
||||||
use_prefill_cuda_graph: bool = False
|
use_prefill_cuda_graph: bool = False
|
||||||
deep_gemm_metadata: Any = field(init=False, repr=False)
|
deep_gemm_metadata: Any = field(init=False, repr=False)
|
||||||
topk_metadata: torch.Tensor = 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()
|
envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.get()
|
||||||
or is_xpu()
|
or is_xpu()
|
||||||
or envs.SGLANG_OPT_USE_AITER_INDEXER.get()
|
or envs.SGLANG_OPT_USE_AITER_INDEXER.get()
|
||||||
):
|
) and not self.force_deep_gemm_metadata:
|
||||||
self.deep_gemm_metadata = None
|
self.deep_gemm_metadata = None
|
||||||
else:
|
else:
|
||||||
import deep_gemm
|
import deep_gemm
|
||||||
|
|
||||||
use_jit_indexer = (
|
use_jit_indexer = not self.force_deep_gemm_metadata and (
|
||||||
envs.SGLANG_OPT_USE_JIT_INDEXER_METADATA.get()
|
envs.SGLANG_OPT_USE_JIT_INDEXER_METADATA.get()
|
||||||
or self.c4_seq_lens.numel() > _LARGE_INDEXER_QUERY_THRESHOLD
|
or self.c4_seq_lens.numel() > _LARGE_INDEXER_QUERY_THRESHOLD
|
||||||
)
|
)
|
||||||
@@ -183,7 +184,11 @@ class PagedIndexerMetadata:
|
|||||||
copy_metadata(
|
copy_metadata(
|
||||||
src=other,
|
src=other,
|
||||||
dst=self,
|
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,
|
copy_fields=copy_fields,
|
||||||
assign_fields=assign_fields,
|
assign_fields=assign_fields,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -7263,9 +7263,11 @@ class ServerArgs:
|
|||||||
"Debug mode for CUDA graph is enabled via breakable CUDA graph. "
|
"Debug mode for CUDA graph is enabled via breakable CUDA graph. "
|
||||||
"All operations will run eagerly through the graph capture/replay path."
|
"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(
|
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."
|
"DeepGEMM FP4 indexer support."
|
||||||
)
|
)
|
||||||
# FP8 W_o GEMM needs DeepGEMM JIT. Enable exactly where the runtime can run
|
# 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.benchmark.bench_utils import run_bench
|
||||||
from sglang.kernels.jit.benchmark.utils import get_benchmark_range
|
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
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
|
||||||
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:
|
try:
|
||||||
@@ -167,8 +170,8 @@ def benchmark(batch: int, seq_len_kv: int, provider: str):
|
|||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
if not is_sm100_supported():
|
if not (is_sm100_supported() or is_sm120_supported()):
|
||||||
print("[skip] DeepSeek V4 FP4 indexer benchmark requires SM100 CUDA.")
|
print("[skip] DeepSeek V4 FP4 indexer benchmark requires SM100 or SM120 CUDA.")
|
||||||
sys.exit(0)
|
sys.exit(0)
|
||||||
if deep_gemm is None or per_token_cast_to_fp4 is None:
|
if deep_gemm is None or per_token_cast_to_fp4 is None:
|
||||||
print("[skip] DeepGEMM is unavailable.")
|
print("[skip] DeepGEMM is unavailable.")
|
||||||
|
|||||||
@@ -7,7 +7,10 @@ import torch
|
|||||||
|
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.attention.dsv4.indexer import FP8_DTYPE, C4IndexerBackendMixin
|
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.model_executor.forward_batch_info import ForwardMode
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
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"
|
_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):
|
class TestDSV4NonPagedIndexer(CustomTestCase):
|
||||||
def _is_eligible(self, **overrides):
|
def _is_eligible(self, **overrides):
|
||||||
backend = SimpleNamespace(hisparse_coordinator=None)
|
backend = SimpleNamespace(hisparse_coordinator=None)
|
||||||
|
|||||||
Reference in New Issue
Block a user