From 143b57b8056a63c09804039f4f45ab78df2a3a0c Mon Sep 17 00:00:00 2001 From: fjybiocs <87828006+fjybiocs@users.noreply.github.com> Date: Sat, 29 Nov 2025 12:09:26 +0800 Subject: [PATCH] enable piecewise cuda graph for prefill server (#13377) Co-authored-by: serverance.fu --- python/sglang/srt/server_args.py | 8 +- ...est_disaggregation_piecewise_cuda_graph.py | 86 +++++++++++++++++++ 2 files changed, 92 insertions(+), 2 deletions(-) create mode 100644 test/srt/test_disaggregation_piecewise_cuda_graph.py diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 07a31229a..13cd1533a 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -1825,8 +1825,12 @@ class ServerArgs: self.disaggregation_prefill_pp = self.pp_size self.validate_disagg_tp_size(self.tp_size, self.disaggregation_decode_tp) - self.disable_cuda_graph = True - logger.warning("Cuda graph is disabled for prefill server") + + if not self.enable_piecewise_cuda_graph: + self.disable_cuda_graph = True + logger.warning( + "Cuda graph is disabled for prefill server when piecewise cuda graph is not enabled." + ) def _handle_tokenizer_batching(self): if self.enable_tokenizer_batch_encode and self.enable_dynamic_batch_tokenizer: diff --git a/test/srt/test_disaggregation_piecewise_cuda_graph.py b/test/srt/test_disaggregation_piecewise_cuda_graph.py new file mode 100644 index 000000000..cc75cb9d3 --- /dev/null +++ b/test/srt/test_disaggregation_piecewise_cuda_graph.py @@ -0,0 +1,86 @@ +import unittest +from types import SimpleNamespace + +from sglang.test.few_shot_gsm8k import run_eval as run_eval_few_shot_gsm8k +from sglang.test.test_disaggregation_utils import TestDisaggregationBase +from sglang.test.test_utils import ( + DEFAULT_MODEL_NAME_FOR_TEST, + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + popen_launch_pd_server, +) + + +class TestDisaggregationPiecewiseCudaGraph(TestDisaggregationBase): + """Test piecewise CUDA graph support in disaggregation prefill server""" + + @classmethod + def setUpClass(cls): + super().setUpClass() + cls.model = DEFAULT_MODEL_NAME_FOR_TEST + + # Start servers + cls.start_prefill() + cls.start_decode() + + # Wait for both to be ready + 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", + "1", + "--enable-piecewise-cuda-graph", + ] + 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", + "1", + "--base-gpu-id", + "1", + ] + 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_accuracy(self): + """Verify that piecewise cuda graph works correctly in prefill server""" + 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_few_shot_gsm8k(args) + print(f"GSM8K accuracy with piecewise cuda graph: {metrics['accuracy']:.3f}") + + self.assertGreater(metrics["accuracy"], 0.62) + + +if __name__ == "__main__": + unittest.main()