diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index ea1f02b32..e0c95ad2a 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -4358,6 +4358,12 @@ class ServerArgs: if ( self.cuda_graph_config.prefill.backend == Backend.BREAKABLE and self.get_model_config().is_multimodal_piecewise_cuda_graph_supported + # Keep trtllm_mla on the preferred breakable path. Its current + # breakable compatibility rule disables the graph, avoiding the + # tc_piecewise FlashInfer paged-MLA fallback; once breakable gains + # native support, that rule can be relaxed without re-enabling the + # deprecated tc_piecewise path. + and self._resolved_attention_backends()[0] != "trtllm_mla" ): logger.info( "Using tc_piecewise CUDA graph for validated multimodal " diff --git a/test/registered/unit/configs/test_multimodal_piecewise_cuda_graph.py b/test/registered/unit/configs/test_multimodal_piecewise_cuda_graph.py index b871b259d..ff71e1848 100644 --- a/test/registered/unit/configs/test_multimodal_piecewise_cuda_graph.py +++ b/test/registered/unit/configs/test_multimodal_piecewise_cuda_graph.py @@ -11,6 +11,7 @@ from sglang.srt.configs.model_config import ( from sglang.srt.model_executor.cuda_graph_config import ( Backend, CudaGraphConfig, + Phase, PhaseConfig, ) from sglang.srt.model_executor.forward_batch_info import ( @@ -79,14 +80,57 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase): ) args._cuda_graph_config_locked = set() - with patch.object( - ServerArgs, "_disable_tc_piecewise_cudagraph_if_incompatible" - ) as disable_if_incompatible: + with ( + patch.object( + ServerArgs, "_disable_tc_piecewise_cudagraph_if_incompatible" + ) as disable_if_incompatible, + patch.object( + args, "_resolved_attention_backends", return_value=("fa3", "fa3") + ), + ): args._apply_cuda_graph_compatibility() self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.TC_PIECEWISE) disable_if_incompatible.assert_called_once() + def test_trtllm_mla_stays_on_breakable_and_is_disabled_by_compatibility(self): + args = ServerArgs(model_path="dummy") + args.model_config = SimpleNamespace( + is_multimodal_piecewise_cuda_graph_supported=True + ) + args.cuda_graph_config = CudaGraphConfig( + prefill=PhaseConfig(backend=Backend.BREAKABLE) + ) + args._cuda_graph_config_locked = set() + + with ( + patch.object( + args, + "_resolved_attention_backends", + return_value=("trtllm_mla", "trtllm_mla"), + ), + patch.object(args, "use_mla_backend", return_value=True), + ): + args._apply_cuda_graph_compatibility() + + self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.DISABLED) + + def test_explicit_tc_piecewise_overrides_trtllm_mla_default(self): + args = ServerArgs(model_path="dummy") + args.cuda_graph_config = CudaGraphConfig( + prefill=PhaseConfig(backend=Backend.TC_PIECEWISE) + ) + args._cuda_graph_config_locked = {(Phase.PREFILL, "backend")} + + with patch.object( + args, + "_resolved_attention_backends", + return_value=("trtllm_mla", "trtllm_mla"), + ): + args._apply_cuda_graph_compatibility() + + self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.TC_PIECEWISE) + def test_multimodal_inputs_keep_tc_piecewise_prefill_enabled(self): runner = self._make_prefill_runner(Backend.TC_PIECEWISE)