fix: avoid piecewise prefill graph for trtllm_mla (#32785)
This commit is contained in:
@@ -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 "
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user