feat: enable piecewise prefill graph for Kimi K2.5/K2.7 (#30889)
This commit is contained in:
@@ -1752,6 +1752,7 @@ piecewise_cuda_graph_disabled_model_archs = [
|
|||||||
# cleanly (vision encoder runs eagerly outside the graph via general_mm_embed_routine).
|
# cleanly (vision encoder runs eagerly outside the graph via general_mm_embed_routine).
|
||||||
multimodal_piecewise_cuda_graph_supported_model_archs = [
|
multimodal_piecewise_cuda_graph_supported_model_archs = [
|
||||||
"Cohere2VisionForConditionalGeneration",
|
"Cohere2VisionForConditionalGeneration",
|
||||||
|
"KimiK25ForConditionalGeneration",
|
||||||
"MiniMaxM3SparseForCausalLM",
|
"MiniMaxM3SparseForCausalLM",
|
||||||
"MiniMaxM3SparseForConditionalGeneration",
|
"MiniMaxM3SparseForConditionalGeneration",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -3495,6 +3495,21 @@ class ServerArgs:
|
|||||||
"""
|
"""
|
||||||
if (Phase.PREFILL, "backend") in self._cuda_graph_config_locked:
|
if (Phase.PREFILL, "backend") in self._cuda_graph_config_locked:
|
||||||
return
|
return
|
||||||
|
|
||||||
|
# Breakable is the general CUDA default, but it is not compatible with
|
||||||
|
# multimodal prefill. Models on this allowlist have had their decoder
|
||||||
|
# prefill validated under tc_piecewise; the vision encoder remains
|
||||||
|
# eager outside the captured LM forward.
|
||||||
|
if (
|
||||||
|
self.cuda_graph_config.prefill.backend == Backend.BREAKABLE
|
||||||
|
and self.get_model_config().is_multimodal_piecewise_cuda_graph_supported
|
||||||
|
):
|
||||||
|
logger.info(
|
||||||
|
"Using tc_piecewise CUDA graph for validated multimodal "
|
||||||
|
"decoder prefill."
|
||||||
|
)
|
||||||
|
self.cuda_graph_config.prefill.backend = Backend.TC_PIECEWISE
|
||||||
|
|
||||||
if self.cuda_graph_config.prefill.backend == Backend.TC_PIECEWISE:
|
if self.cuda_graph_config.prefill.backend == Backend.TC_PIECEWISE:
|
||||||
self._disable_tc_piecewise_cudagraph_if_incompatible()
|
self._disable_tc_piecewise_cudagraph_if_incompatible()
|
||||||
elif self.cuda_graph_config.prefill.backend == Backend.BREAKABLE:
|
elif self.cuda_graph_config.prefill.backend == Backend.BREAKABLE:
|
||||||
|
|||||||
@@ -0,0 +1,57 @@
|
|||||||
|
"""Regression tests for multimodal piecewise CUDA graph opt-ins."""
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
from sglang.srt.configs.model_config import (
|
||||||
|
is_multimodal_piecewise_cuda_graph_supported,
|
||||||
|
)
|
||||||
|
from sglang.srt.model_executor.cuda_graph_config import (
|
||||||
|
Backend,
|
||||||
|
CudaGraphConfig,
|
||||||
|
PhaseConfig,
|
||||||
|
)
|
||||||
|
from sglang.srt.server_args import ServerArgs
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
|
||||||
|
def test_kimi_k25_lm_prefill_is_opted_in(self):
|
||||||
|
self.assertTrue(
|
||||||
|
is_multimodal_piecewise_cuda_graph_supported(
|
||||||
|
["KimiK25ForConditionalGeneration"]
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_unknown_multimodal_arch_is_not_opted_in(self):
|
||||||
|
self.assertFalse(
|
||||||
|
is_multimodal_piecewise_cuda_graph_supported(
|
||||||
|
["UnknownVisionForConditionalGeneration"]
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_supported_multimodal_model_upgrades_default_to_tc_piecewise(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(
|
||||||
|
ServerArgs, "_disable_tc_piecewise_cudagraph_if_incompatible"
|
||||||
|
) as disable_if_incompatible:
|
||||||
|
args._apply_cuda_graph_compatibility()
|
||||||
|
|
||||||
|
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.TC_PIECEWISE)
|
||||||
|
disable_if_incompatible.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user