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,
|
||||
) -> 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
|
||||
|
||||
Reference in New Issue
Block a user