Enable multimodal prefill BCG for VL and audio models (#30872)

This commit is contained in:
Xiaoyu Zhang
2026-07-29 06:47:40 +08:00
committed by GitHub
parent 85618cc798
commit c9947b087b
3 changed files with 96 additions and 14 deletions
@@ -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)
@@ -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
@@ -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()