fix: avoid piecewise prefill graph for trtllm_mla (#32785)
This commit is contained in:
@@ -4358,6 +4358,12 @@ class ServerArgs:
|
|||||||
if (
|
if (
|
||||||
self.cuda_graph_config.prefill.backend == Backend.BREAKABLE
|
self.cuda_graph_config.prefill.backend == Backend.BREAKABLE
|
||||||
and self.get_model_config().is_multimodal_piecewise_cuda_graph_supported
|
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(
|
logger.info(
|
||||||
"Using tc_piecewise CUDA graph for validated multimodal "
|
"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 (
|
from sglang.srt.model_executor.cuda_graph_config import (
|
||||||
Backend,
|
Backend,
|
||||||
CudaGraphConfig,
|
CudaGraphConfig,
|
||||||
|
Phase,
|
||||||
PhaseConfig,
|
PhaseConfig,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_executor.forward_batch_info import (
|
from sglang.srt.model_executor.forward_batch_info import (
|
||||||
@@ -79,14 +80,57 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
|
|||||||
)
|
)
|
||||||
args._cuda_graph_config_locked = set()
|
args._cuda_graph_config_locked = set()
|
||||||
|
|
||||||
with patch.object(
|
with (
|
||||||
ServerArgs, "_disable_tc_piecewise_cudagraph_if_incompatible"
|
patch.object(
|
||||||
) as disable_if_incompatible:
|
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()
|
args._apply_cuda_graph_compatibility()
|
||||||
|
|
||||||
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.TC_PIECEWISE)
|
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.TC_PIECEWISE)
|
||||||
disable_if_incompatible.assert_called_once()
|
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):
|
def test_multimodal_inputs_keep_tc_piecewise_prefill_enabled(self):
|
||||||
runner = self._make_prefill_runner(Backend.TC_PIECEWISE)
|
runner = self._make_prefill_runner(Backend.TC_PIECEWISE)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user