Add FP4 Indexer for DeepSeek V4 on SM120 (#27059)
This commit is contained in:
@@ -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