From 83c3158014743a5205623d6fadd03ab8ced2c2bf Mon Sep 17 00:00:00 2001 From: Derek Yu <81697272+DerekY2@users.noreply.github.com> Date: Wed, 1 Apr 2026 21:17:38 -0400 Subject: [PATCH] [CI] Add Llama 3.1 8B Instruct FP4 CI test on SM120 (#20648) --- .../registered/quant/test_nvfp4_gemm_sm120.py | 71 +++++++++++++++++++ 1 file changed, 71 insertions(+) create mode 100644 test/registered/quant/test_nvfp4_gemm_sm120.py diff --git a/test/registered/quant/test_nvfp4_gemm_sm120.py b/test/registered/quant/test_nvfp4_gemm_sm120.py new file mode 100644 index 000000000..95f32942e --- /dev/null +++ b/test/registered/quant/test_nvfp4_gemm_sm120.py @@ -0,0 +1,71 @@ +import unittest +from types import SimpleNamespace +from urllib.parse import urlparse + +from sglang.srt.utils import get_device_sm, kill_process_tree +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.few_shot_gsm8k import run_eval as run_eval_few_shot_gsm8k +from sglang.test.test_utils import ( + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + DEFAULT_URL_FOR_TEST, + popen_launch_server, + try_cached_model, +) + +register_cuda_ci(est_time=90, suite="stage-b-test-small-1-gpu") + +MODEL_PATH = "nvidia/Llama-3.1-8B-Instruct-NVFP4" + + +class FP4GemmSM120Base: + backend = None + + @classmethod + def setUpClass(cls): + 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, + "--disable-piecewise-cuda-graph", + ] + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=other_args, + ) + + @classmethod + def tearDownClass(cls): + if hasattr(cls, "process"): + kill_process_tree(cls.process.pid) + + def test_gsm8k(self): + parsed_url = urlparse(self.base_url) + args = SimpleNamespace( + num_shots=5, + data_path=None, + num_questions=1319, + max_new_tokens=512, + parallel=200, + host=parsed_url.hostname, + port=parsed_url.port, + ) + metrics = run_eval_few_shot_gsm8k(args) + print(f"{metrics=}") + self.assertGreater(metrics["accuracy"], 0.64) + + +@unittest.skipIf(get_device_sm() < 100, "Test requires CUDA SM 100 or higher") +class TestFP4GemmSM120Auto(FP4GemmSM120Base, unittest.TestCase): + backend = "auto" + + +if __name__ == "__main__": + unittest.main()