diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/qwen_image_layered.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/qwen_image_layered.py index 6d31fbebc..1bfb7444e 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/qwen_image_layered.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/qwen_image_layered.py @@ -56,6 +56,14 @@ def _seq_lens_from_optional_mask( return [int(x) for x in prompt_embeds_mask.sum(dim=1).tolist()] +def _resolve_layered_image_path(image_path: str | list[str]) -> str: + if isinstance(image_path, str): + return image_path + if isinstance(image_path, list) and image_path: + return image_path[0] + raise ValueError("Qwen-Image-Layered requires a non-empty image_path.") + + # Copied from diffusers.pipelines.qwenimage.pipeline_qwenimage_edit_plus.calculate_dimensions def calculate_dimensions(target_area, ratio): width = math.sqrt(target_area * ratio) @@ -490,7 +498,7 @@ the image\n<|vision_start|><|image_pad|><|vision_end|><|im_end|>\n<|im_start|>as generator = batch.generator assert batch.image_path is not None - image = load_image(batch.image_path[0]) + image = load_image(_resolve_layered_image_path(batch.image_path)) image = image.convert("RGBA") image_size = image.size resolution = server_args.pipeline_config.resolution diff --git a/python/sglang/multimodal_gen/test/unit/test_disagg_roles.py b/python/sglang/multimodal_gen/test/unit/test_disagg_roles.py index 05a0cf917..a15a06b7f 100644 --- a/python/sglang/multimodal_gen/test/unit/test_disagg_roles.py +++ b/python/sglang/multimodal_gen/test/unit/test_disagg_roles.py @@ -59,6 +59,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.m ) from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.qwen_image_layered import ( QwenImageLayeredBeforeDenoisingStage, + _resolve_layered_image_path, _resolve_text_encoder_dtype, ) from sglang.multimodal_gen.runtime.server_args import set_global_server_args @@ -429,6 +430,20 @@ class TestPipelineSpecificExtraModules(unittest.TestCase): class TestQwenImageLayeredDtype(_GlobalStageArgsMixin, unittest.TestCase): + def test_layered_image_path_accepts_string_and_list(self): + self.assertEqual( + _resolve_layered_image_path("/tmp/input.png"), + "/tmp/input.png", + ) + self.assertEqual( + _resolve_layered_image_path(["/tmp/input.png"]), + "/tmp/input.png", + ) + + def test_layered_image_path_rejects_empty_list(self): + with self.assertRaisesRegex(ValueError, "non-empty image_path"): + _resolve_layered_image_path([]) + def test_text_encoder_dtype_uses_parameter_dtype_without_dtype_attr(self): text_encoder = torch.nn.Linear(1, 1, bias=False).to(dtype=torch.bfloat16) self.assertEqual(