From 7f9f85d4c81c0574e448d65eabcf48f18f9c6434 Mon Sep 17 00:00:00 2001 From: xingsy97 <87063252+xingsy97@users.noreply.github.com> Date: Sun, 8 Mar 2026 14:26:05 +0800 Subject: [PATCH] [diffusion] feat: make QwenImageLayered resolution configurable (#20044) --- .../configs/pipeline_configs/base.py | 7 ++++++ .../configs/pipeline_configs/qwen_image.py | 2 +- .../qwen_image_layered.py | 2 +- .../test/unit/test_server_args_unit.py | 24 +++++++++++++++++++ 4 files changed, 33 insertions(+), 2 deletions(-) diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py index 455635467..2f6f47eb8 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py @@ -436,6 +436,13 @@ class PipelineConfig: default=PipelineConfig.flow_shift, help="Flow shift parameter", ) + parser.add_argument( + f"--{prefix_with_dot}resolution", + type=int, + dest=f"{prefix_with_dot.replace('-', '_')}resolution", + default=None, + help="Override the selected pipeline config's resolution setting. Only applies to pipelines that define a resolution field.", + ) # DiT configuration parser.add_argument( diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/qwen_image.py b/python/sglang/multimodal_gen/configs/pipeline_configs/qwen_image.py index 40f42c0c7..08d20c958 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/qwen_image.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/qwen_image.py @@ -500,7 +500,7 @@ class QwenImageEditPlus_2511_PipelineConfig(QwenImageEditPlusPipelineConfig): @dataclass class QwenImageLayeredPipelineConfig(QwenImageEditPipelineConfig): - resolution: int = 640 # TODO: allow user to set resolution + resolution: int = 640 vae_precision: str = "bf16" def _prepare_edit_cond_kwargs( 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 4d64759c5..2cb35dead 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 @@ -422,7 +422,7 @@ the image\n<|vision_start|><|image_pad|><|vision_end|><|im_end|>\n<|im_start|>as image = load_image(batch.image_path[0]) image = image.convert("RGBA") image_size = image.size - resolution = 640 # TODO: support user-specified resolution + resolution = server_args.pipeline_config.resolution calculated_width, calculated_height = calculate_dimensions( resolution * resolution, image_size[0] / image_size[1] ) diff --git a/python/sglang/multimodal_gen/test/unit/test_server_args_unit.py b/python/sglang/multimodal_gen/test/unit/test_server_args_unit.py index 59d7584d6..cee6e2cf3 100644 --- a/python/sglang/multimodal_gen/test/unit/test_server_args_unit.py +++ b/python/sglang/multimodal_gen/test/unit/test_server_args_unit.py @@ -1,8 +1,11 @@ import os +import sys import unittest +from unittest.mock import patch from sglang.multimodal_gen.registry import _get_config_info from sglang.multimodal_gen.runtime.server_args import ServerArgs +from sglang.multimodal_gen.utils import FlexibleArgumentParser class TestServerArgsPathExpansion(unittest.TestCase): @@ -46,5 +49,26 @@ class TestModelIdResolution(unittest.TestCase): _get_config_info("/data/no-such-model", model_id="NonExistentModelXYZ") +class TestPipelineResolutionCliOverride(unittest.TestCase): + def setUp(self): + _get_config_info.cache_clear() + + def test_resolution_flag_overrides_qwen_image_layered_pipeline_config(self): + parser = FlexibleArgumentParser() + ServerArgs.add_cli_args(parser) + argv = [ + "--model-path", + "Qwen/Qwen-Image-Layered", + "--resolution", + "768", + ] + + with patch.object(sys, "argv", ["sglang"] + argv): + args, unknown_args = parser.parse_known_args(argv) + server_args = ServerArgs.from_cli_args(args, unknown_args) + + self.assertEqual(server_args.pipeline_config.resolution, 768) + + if __name__ == "__main__": unittest.main()