From c9947b087bf9d3d16b5198234ba4c39b68bb79e9 Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <1182563586@qq.com> Date: Wed, 29 Jul 2026 06:47:40 +0800 Subject: [PATCH] Enable multimodal prefill BCG for VL and audio models (#30872) --- .../srt/layers/rotary_embedding/mrope.py | 1 + .../runner/prefill_cuda_graph_runner.py | 46 +++++++++----- .../test_prefill_cuda_graph_runner_helpers.py | 63 +++++++++++++++++++ 3 files changed, 96 insertions(+), 14 deletions(-) create mode 100644 test/registered/unit/model_executor/test_prefill_cuda_graph_runner_helpers.py diff --git a/python/sglang/srt/layers/rotary_embedding/mrope.py b/python/sglang/srt/layers/rotary_embedding/mrope.py index 4fd88dcc1..68346e505 100644 --- a/python/sglang/srt/layers/rotary_embedding/mrope.py +++ b/python/sglang/srt/layers/rotary_embedding/mrope.py @@ -243,6 +243,7 @@ class MRotaryEmbedding(RotaryEmbedding): fused_set_kv_buffer_arg=None, ) -> Tuple[torch.Tensor, torch.Tensor]: assert positions.ndim == 1 or positions.ndim == 2 + self._match_cos_sin_cache_dtype(query) if positions.ndim == 2 and self.mrope_section: return self.forward_triton(positions, query, key) return self.forward_native(positions, query, key, fused_set_kv_buffer_arg) diff --git a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py index 42f083590..84860b028 100644 --- a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py @@ -101,6 +101,7 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo from sglang.srt.model_executor.runner_utils.buffers import ( PrefillInputBuffers, ) +from sglang.srt.model_loader.utils import resolve_language_model from sglang.srt.runtime_context import get_parallel from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim from sglang.srt.utils import ( @@ -135,6 +136,29 @@ def _ceil_div(a: int, b: int) -> int: return -(-a // b) +def _resolve_transformer_layer_model(model: torch.nn.Module) -> torch.nn.Module: + """Find the module that owns decoder layers behind language/model wrappers.""" + try: + layer_model = resolve_language_model(model) + except AttributeError: + layer_model = getattr(model, "language_model", model) + + seen = set() + while not hasattr(layer_model, "layers") and hasattr(layer_model, "model"): + obj_id = id(layer_model) + if obj_id in seen: + break + seen.add(obj_id) + layer_model = layer_model.model + + if not hasattr(layer_model, "layers"): + raise RuntimeError( + f"could not resolve inner layer_model on {type(model).__name__}; " + f"resolved to {type(layer_model).__name__} without layers." + ) + return layer_model + + def _slice_output_rows(output: Any, num_tokens: int) -> Any: """Slice every tensor leaf in a transformer-body output by token rows. @@ -454,21 +478,15 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner): # not the LM head + logits_processor — the eager tail keeps the captured # graph bs-invariant so req_slots is not bound by an (req_slots, vocab) buffer. if isinstance(self.backend, (BreakableCudaGraphBackend, FullCudaGraphBackend)): - language_model = getattr( - self.model_runner.model, "language_model", self.model_runner.model - ) - if hasattr(language_model, "model") and hasattr( - language_model.model, "layers" - ): - self.layer_model = language_model.model - elif hasattr(language_model, "layers"): - self.layer_model = language_model - else: - raise RuntimeError( - f"{type(self.backend).__name__} could not resolve inner " - f"layer_model on {type(language_model).__name__}; " - f"this backend is unsupported for this model architecture." + try: + self.layer_model = _resolve_transformer_layer_model( + self.model_runner.model ) + except RuntimeError as exc: + raise RuntimeError( + f"{type(self.backend).__name__} {exc} This backend is " + f"unsupported for this model architecture." + ) from exc params = list(inspect.signature(self.layer_model.forward).parameters) self._input_embeds_arg_idx = ( params.index("input_embeds") if "input_embeds" in params else None diff --git a/test/registered/unit/model_executor/test_prefill_cuda_graph_runner_helpers.py b/test/registered/unit/model_executor/test_prefill_cuda_graph_runner_helpers.py new file mode 100644 index 000000000..b01de41d1 --- /dev/null +++ b/test/registered/unit/model_executor/test_prefill_cuda_graph_runner_helpers.py @@ -0,0 +1,63 @@ +"""Unit tests for prefill CUDA graph wrapper helpers.""" + +import unittest +from types import SimpleNamespace + +from sglang.srt.model_executor.runner.prefill_cuda_graph_runner import ( + _resolve_transformer_layer_model, +) +from sglang.srt.model_loader.utils import resolve_language_model +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + + +class _LayerModel: + def __init__(self): + self.layers = [object()] + + def forward(self, input_ids, positions, forward_batch, input_embeds=None): + return input_embeds + + +class TestPrefillCudaGraphRunnerHelpers(CustomTestCase): + def test_resolve_layer_model_from_language_model_wrapper(self): + layer_model = _LayerModel() + model = SimpleNamespace(language_model=SimpleNamespace(model=layer_model)) + + self.assertIs(_resolve_transformer_layer_model(model), layer_model) + + def test_resolve_layer_model_from_nested_model_wrapper(self): + layer_model = _LayerModel() + model = SimpleNamespace(model=SimpleNamespace(model=layer_model)) + + self.assertIs(_resolve_transformer_layer_model(model), layer_model) + + def test_resolve_layer_model_rejects_wrapper_without_layers(self): + model = SimpleNamespace() + model.model = model + + with self.assertRaisesRegex(RuntimeError, "without layers"): + _resolve_transformer_layer_model(model) + + def test_resolve_language_model_accepts_asr_style_wrapper(self): + language_model = object() + self.assertIs( + resolve_language_model(SimpleNamespace(language_model=language_model)), + language_model, + ) + + def test_resolve_language_model_accepts_omni_style_wrapper(self): + language_model = object() + omni_model = type("Qwen3OmniMoeForConditionalGeneration", (), {})() + omni_model.thinker = SimpleNamespace(model=language_model) + self.assertIs(resolve_language_model(omni_model), language_model) + + def test_resolve_language_model_rejects_non_language_wrapper(self): + with self.assertRaises(AttributeError): + resolve_language_model(SimpleNamespace()) + + +if __name__ == "__main__": + unittest.main()