Enable multimodal prefill BCG for VL and audio models (#30872)
This commit is contained in:
@@ -243,6 +243,7 @@ class MRotaryEmbedding(RotaryEmbedding):
|
|||||||
fused_set_kv_buffer_arg=None,
|
fused_set_kv_buffer_arg=None,
|
||||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
assert positions.ndim == 1 or positions.ndim == 2
|
assert positions.ndim == 1 or positions.ndim == 2
|
||||||
|
self._match_cos_sin_cache_dtype(query)
|
||||||
if positions.ndim == 2 and self.mrope_section:
|
if positions.ndim == 2 and self.mrope_section:
|
||||||
return self.forward_triton(positions, query, key)
|
return self.forward_triton(positions, query, key)
|
||||||
return self.forward_native(positions, query, key, fused_set_kv_buffer_arg)
|
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 (
|
from sglang.srt.model_executor.runner_utils.buffers import (
|
||||||
PrefillInputBuffers,
|
PrefillInputBuffers,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.model_loader.utils import resolve_language_model
|
||||||
from sglang.srt.runtime_context import get_parallel
|
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.speculative.eagle_utils import get_draft_input_from_target_hidden_dim
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
@@ -135,6 +136,29 @@ def _ceil_div(a: int, b: int) -> int:
|
|||||||
return -(-a // b)
|
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:
|
def _slice_output_rows(output: Any, num_tokens: int) -> Any:
|
||||||
"""Slice every tensor leaf in a transformer-body output by token rows.
|
"""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
|
# 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.
|
# graph bs-invariant so req_slots is not bound by an (req_slots, vocab) buffer.
|
||||||
if isinstance(self.backend, (BreakableCudaGraphBackend, FullCudaGraphBackend)):
|
if isinstance(self.backend, (BreakableCudaGraphBackend, FullCudaGraphBackend)):
|
||||||
language_model = getattr(
|
try:
|
||||||
self.model_runner.model, "language_model", self.model_runner.model
|
self.layer_model = _resolve_transformer_layer_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."
|
|
||||||
)
|
)
|
||||||
|
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)
|
params = list(inspect.signature(self.layer_model.forward).parameters)
|
||||||
self._input_embeds_arg_idx = (
|
self._input_embeds_arg_idx = (
|
||||||
params.index("input_embeds") if "input_embeds" in params else None
|
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()
|
||||||
Reference in New Issue
Block a user