From b91fb8393e6b1a5f7ad71f1e3b3fab646ab21cdc Mon Sep 17 00:00:00 2001 From: Lancer Date: Sat, 7 Mar 2026 13:01:59 +0800 Subject: [PATCH] [diffusion] fix: fix multi-prompt generation and support multiple prompts in cli (#19960) Signed-off-by: Lancer --- .../multimodal_gen/configs/sample/sampling_params.py | 3 ++- .../runtime/entrypoints/diffusion_generator.py | 9 ++++++++- 2 files changed, 10 insertions(+), 2 deletions(-) diff --git a/python/sglang/multimodal_gen/configs/sample/sampling_params.py b/python/sglang/multimodal_gen/configs/sample/sampling_params.py index fd00f82dc..ec387da2d 100644 --- a/python/sglang/multimodal_gen/configs/sample/sampling_params.py +++ b/python/sglang/multimodal_gen/configs/sample/sampling_params.py @@ -624,8 +624,9 @@ class SamplingParams: parser.add_argument( "--prompt", type=str, + nargs="+", default=SamplingParams.prompt, - help="Text prompt for generation", + help="Text prompt(s) for generation. Use space-separated values for multiple prompts, e.g., --prompt 'prompt 1' 'prompt 2'", ) parser.add_argument( "--negative-prompt", diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py b/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py index 12631f137..162df5f89 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py @@ -166,12 +166,19 @@ class DiffGenerator: """ # 1. prepare requests prompts = self._resolve_prompts(sampling_params_kwargs.get("prompt")) + user_output_file_name = sampling_params_kwargs.get("output_file_name") + + if len(prompts) > 1 and user_output_file_name is not None: + raise ValueError( + "Cannot use multiple prompts with a fixed output_file_name. " + "Either remove --output-file-name or use a single prompt." + ) + sampling_params_orig = SamplingParams.from_user_sampling_params_args( self.server_args.model_path, server_args=self.server_args, **sampling_params_kwargs, ) - user_output_file_name = sampling_params_kwargs.get("output_file_name") requests: list[Req] = [] for p in prompts: