[Refactor] Add -fp4-gemm-backend to replace SGLANG_FLASHINFER_FP4_GEMM_BACKEND (#16534)

Co-authored-by: Vincent Zhong <207368749+vincentzed@users.noreply.github.com>
This commit is contained in:
b8zhong
2026-01-18 07:25:46 -08:00
committed by GitHub
parent f3a7c7dcd9
commit 4df74eb576
9 changed files with 144 additions and 18 deletions

View File

@@ -35,7 +35,7 @@ suites = {
TestFile("test_deepseek_v3_fp4_4gpu.py", 1500),
TestFile("test_fp8_blockwise_gemm.py", 280),
TestFile("test_gpt_oss_4gpu.py", 700),
TestFile("test_llama31_fp4.py", 90),
TestFile("test_nvfp4_gemm.py", 360),
],
# "per-commit-8-gpu-b200": [
# TestFile("test_mistral_large3_basic.py", 275), # Moved to nightly - large model

View File

@@ -8,21 +8,27 @@ from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
popen_launch_server,
try_cached_model,
)
MODEL_PATH = "nvidia/Llama-3.1-8B-Instruct-FP4"
MODEL_PATH = "nvidia/Llama-3.1-8B-Instruct-NVFP4"
@unittest.skipIf(get_device_sm() < 100, "Test requires CUDA SM 100 or higher")
class TestLlama31FP4(unittest.TestCase):
class FP4GemmBase:
backend = None
@classmethod
def setUpClass(cls):
cls.model = MODEL_PATH
if cls.backend is None:
raise NotImplementedError("Subclass must set 'backend' attribute")
cls.model = try_cached_model(MODEL_PATH)
cls.base_url = DEFAULT_URL_FOR_TEST
other_args = [
"--trust-remote-code",
"--quantization",
"modelopt_fp4",
"--fp4-gemm-backend",
cls.backend,
]
cls.process = popen_launch_server(
cls.model,
@@ -52,5 +58,25 @@ class TestLlama31FP4(unittest.TestCase):
self.assertGreater(metrics["accuracy"], 0.64)
@unittest.skipIf(get_device_sm() < 100, "Test requires CUDA SM 100 or higher")
class TestFP4GemmAuto(FP4GemmBase, unittest.TestCase):
backend = "auto"
@unittest.skipIf(get_device_sm() < 100, "Test requires CUDA SM 100 or higher")
class TestFP4GemmCutlass(FP4GemmBase, unittest.TestCase):
backend = "cutlass"
@unittest.skipIf(get_device_sm() < 100, "Test requires CUDA SM 100 or higher")
class TestFP4GemmCudnn(FP4GemmBase, unittest.TestCase):
backend = "cudnn"
@unittest.skipIf(get_device_sm() < 100, "Test requires CUDA SM 100 or higher")
class TestFP4GemmTrtllm(FP4GemmBase, unittest.TestCase):
backend = "trtllm"
if __name__ == "__main__":
unittest.main()