From d277a86dea6b927204982b2dd2d968834574740b Mon Sep 17 00:00:00 2001 From: Shangming Cai Date: Sun, 14 Dec 2025 23:54:25 +0800 Subject: [PATCH] [CI] Add disaggregation decode PP test (#15114) --- test/srt/test_disaggregation_pp.py | 79 +++++++++++++++++++++++++++++- 1 file changed, 77 insertions(+), 2 deletions(-) diff --git a/test/srt/test_disaggregation_pp.py b/test/srt/test_disaggregation_pp.py index 0a1fb4e8c..4ac271b3f 100644 --- a/test/srt/test_disaggregation_pp.py +++ b/test/srt/test_disaggregation_pp.py @@ -14,7 +14,7 @@ from sglang.test.test_utils import ( ) -class TestDisaggregationPPAccuracy(PDDisaggregationServerBase): +class TestDisaggregationPrefillPPAccuracy(PDDisaggregationServerBase): @classmethod def setUpClass(cls): super().setUpClass() @@ -56,7 +56,82 @@ class TestDisaggregationPPAccuracy(PDDisaggregationServerBase): "--trust-remote-code", "--disaggregation-mode", "decode", - "--tp", + "--tp-size", + "2", + "--base-gpu-id", + "4", + ] + decode_args += cls.transfer_backend + cls.rdma_devices + cls.process_decode = popen_launch_pd_server( + cls.model, + cls.decode_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=decode_args, + ) + + def test_gsm8k(self): + args = SimpleNamespace( + num_shots=5, + data_path=None, + num_questions=200, + max_new_tokens=512, + parallel=128, + host=f"http://{self.base_host}", + port=int(self.lb_port), + ) + metrics = run_eval(args) + print(f"{metrics=}") + + self.assertGreater(metrics["accuracy"], 0.24) + # Wait a little bit so that the memory check happens. + time.sleep(5) + + +class TestDisaggregationDecodePPAccuracy(PDDisaggregationServerBase): + @classmethod + def setUpClass(cls): + super().setUpClass() + cls.model = try_cached_model(DEFAULT_MODEL_NAME_FOR_TEST) + + # Non blocking start servers + cls.start_prefill() + cls.start_decode() + + # Block until both + cls.wait_server_ready(cls.prefill_url + "/health") + cls.wait_server_ready(cls.decode_url + "/health") + + cls.launch_lb() + + @classmethod + def start_prefill(cls): + prefill_args = [ + "--trust-remote-code", + "--disaggregation-mode", + "prefill", + "--tp-size", + "2", + "--pp-size", + "2", + "--disable-overlap-schedule", + ] + prefill_args += cls.transfer_backend + cls.rdma_devices + cls.process_prefill = popen_launch_pd_server( + cls.model, + cls.prefill_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=prefill_args, + ) + + @classmethod + def start_decode(cls): + decode_args = [ + "--trust-remote-code", + "--disaggregation-mode", + "decode", + "--tp-size", + "2", + "--pp-size", "2", "--base-gpu-id", "4",