diff --git a/test/srt/run_suite.py b/test/srt/run_suite.py index 541d59b10..35d3d2fb7 100644 --- a/test/srt/run_suite.py +++ b/test/srt/run_suite.py @@ -161,7 +161,7 @@ suites = { "per-commit-8-gpu-h20": [ TestFile("quant/test_w4a8_deepseek_v3.py", 520), TestFile("test_disaggregation_different_tp.py", 600), - TestFile("test_disaggregation_pp.py", 140), + TestFile("test_disaggregation_pp.py", 180), TestFile("test_disaggregation_dp_attention.py", 155), ], "per-commit-4-gpu-b200-stage-b": [ diff --git a/test/srt/test_disaggregation_pp.py b/test/srt/test_disaggregation_pp.py index 4ac271b3f..4c99ea704 100644 --- a/test/srt/test_disaggregation_pp.py +++ b/test/srt/test_disaggregation_pp.py @@ -87,6 +87,80 @@ class TestDisaggregationPrefillPPAccuracy(PDDisaggregationServerBase): time.sleep(5) +class TestDisaggregationPrefillPPDynamicChunkAccuracy(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", + "--enable-dynamic-chunking", + ] + 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", + "--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):