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
@@ -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()