From da5528db3079170d1b3306e48373309536740930 Mon Sep 17 00:00:00 2001 From: Mick Date: Wed, 29 Jul 2026 15:25:04 +0800 Subject: [PATCH] fix(vlm): materialize Qwen3-VL features on the vision device (#31596) --- python/sglang/srt/models/qwen3_vl.py | 19 +++-- .../test_qwen3_vl_feature_materialization.py | 75 +++++++++++++++++++ 2 files changed, 87 insertions(+), 7 deletions(-) create mode 100644 test/registered/unit/models/test_qwen3_vl_feature_materialization.py diff --git a/python/sglang/srt/models/qwen3_vl.py b/python/sglang/srt/models/qwen3_vl.py index 7036cd717..3e9d5d211 100644 --- a/python/sglang/srt/models/qwen3_vl.py +++ b/python/sglang/srt/models/qwen3_vl.py @@ -68,7 +68,10 @@ from sglang.srt.models.utils import ( WeightsMapper, compute_cu_seqlens_from_grid_numpy, ) -from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model +from sglang.srt.multimodal.mm_utils import ( + materialize_multimodal_features, + run_dp_sharded_mrope_vision_model, +) from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.utils import ( @@ -1330,9 +1333,10 @@ class Qwen3VLForConditionalGeneration(nn.Module): return pattern.pad_input_tokens(input_ids, mm_inputs) def get_image_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor: - # in qwen-vl, last dim is the same - pixel_values = torch.cat([item.feature for item in items], dim=0).type( - self.visual.dtype + pixel_values = materialize_multimodal_features( + [item.feature for item in items], + device=self.visual.device, + dtype=self.visual.dtype, ) image_grid_thw = torch.concat([item.image_grid_thw for item in items], dim=0) assert pixel_values.dim() == 2, pixel_values.dim() @@ -1349,9 +1353,10 @@ class Qwen3VLForConditionalGeneration(nn.Module): return self.visual(pixel_values, grid_thw=image_grid_thw) def get_video_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor: - # in qwen-vl, last dim is the same - pixel_values = torch.cat([item.feature for item in items], dim=0).type( - self.visual.dtype + pixel_values = materialize_multimodal_features( + [item.feature for item in items], + device=self.visual.device, + dtype=self.visual.dtype, ) video_grid_thw = torch.concat([item.video_grid_thw for item in items], dim=0) assert pixel_values.dim() == 2, pixel_values.dim() diff --git a/test/registered/unit/models/test_qwen3_vl_feature_materialization.py b/test/registered/unit/models/test_qwen3_vl_feature_materialization.py new file mode 100644 index 000000000..bb9eaf0cd --- /dev/null +++ b/test/registered/unit/models/test_qwen3_vl_feature_materialization.py @@ -0,0 +1,75 @@ +"""Regression tests for Qwen3-VL multimodal feature materialization.""" + +import unittest +from types import SimpleNamespace + +import torch + +from sglang.srt.models.qwen3_vl import Qwen3VLForConditionalGeneration +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + + +class _RecordingVisual: + device = torch.device("meta") + dtype = torch.bfloat16 + + def __init__(self): + self.pixel_values = None + self.grid_thw = None + + def __call__(self, pixel_values, *, grid_thw): + self.pixel_values = pixel_values + self.grid_thw = grid_thw + return pixel_values + + +class TestQwen3VLFeatureMaterialization(CustomTestCase): + def test_image_features_are_packed_on_the_visual_device(self): + visual = _RecordingVisual() + model = SimpleNamespace(visual=visual, use_data_parallel=False) + items = [ + SimpleNamespace( + feature=torch.ones(2, 3), + image_grid_thw=torch.tensor([[1, 1, 2]]), + ), + SimpleNamespace( + feature=torch.ones(1, 3), + image_grid_thw=torch.tensor([[1, 1, 1]]), + ), + ] + output = Qwen3VLForConditionalGeneration.get_image_feature(model, items) + + self.assertIs(visual.pixel_values, output) + self.assertEqual(output.shape, (3, 3)) + self.assertEqual(output.device, visual.device) + self.assertEqual(output.dtype, visual.dtype) + + def test_video_features_are_packed_on_the_visual_device(self): + visual = _RecordingVisual() + model = SimpleNamespace(visual=visual, use_data_parallel=False) + items = [ + SimpleNamespace( + feature=torch.ones(3, 4), + video_grid_thw=torch.tensor([[1, 1, 3]]), + ), + SimpleNamespace( + feature=torch.ones(2, 4), + video_grid_thw=torch.tensor([[1, 1, 2]]), + ), + ] + output = Qwen3VLForConditionalGeneration.get_video_feature(model, items) + + self.assertIs(visual.pixel_values, output) + self.assertEqual(output.shape, (5, 4)) + self.assertEqual(output.device, visual.device) + self.assertEqual(output.dtype, visual.dtype) + self.assertTrue( + torch.equal(visual.grid_thw, torch.tensor([[1, 1, 3], [1, 1, 2]])) + ) + + +if __name__ == "__main__": + unittest.main(verbosity=2)