From e3f51e823e098384438945dc5c37fa3bb5287b3b Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <35585791+BBuf@users.noreply.github.com> Date: Sun, 14 Dec 2025 19:44:03 +0800 Subject: [PATCH] [diffusion] feat: add support for additional sampling parameters in video generation API (#15062) Co-authored-by: Mick --- .../configs/sample/sampling_params.py | 10 +++++- .../runtime/entrypoints/openai/protocol.py | 4 +++ .../runtime/entrypoints/openai/video_api.py | 33 ++++++++++++------- 3 files changed, 35 insertions(+), 12 deletions(-) diff --git a/python/sglang/multimodal_gen/configs/sample/sampling_params.py b/python/sglang/multimodal_gen/configs/sample/sampling_params.py index 98bdc2a50..86e743995 100644 --- a/python/sglang/multimodal_gen/configs/sample/sampling_params.py +++ b/python/sglang/multimodal_gen/configs/sample/sampling_params.py @@ -123,6 +123,7 @@ class SamplingParams: # Denoising parameters num_inference_steps: int = None guidance_scale: float = None + guidance_scale_2: float = None guidance_rescale: float = 0.0 boundary_ratio: float | None = None @@ -494,6 +495,13 @@ class SamplingParams: default=SamplingParams.guidance_scale, help="Classifier-free guidance scale", ) + parser.add_argument( + "--guidance-scale-2", + type=float, + default=SamplingParams.guidance_scale_2, + dest="guidance_scale_2", + help="Secondary guidance scale for dual-guidance models (e.g., Wan low-noise expert)", + ) parser.add_argument( "--guidance-rescale", type=float, @@ -589,7 +597,7 @@ class SamplingParams: attrs = sampling_params_fields & args_attrs args.height_not_provided = False args.width_not_provided = False - return {attr: getattr(args, attr) for attr in attrs} + return {attr: getattr(args, attr) for attr in attrs if hasattr(args, attr)} def output_file_path(self): return os.path.join(self.output_path, self.output_file_name) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py index bafb97cce..8675b038c 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py @@ -58,6 +58,10 @@ class VideoGenerationsRequest(BaseModel): num_frames: Optional[int] = None seed: Optional[int] = 1024 generator_device: Optional[str] = "cuda" + num_inference_steps: Optional[int] = None + guidance_scale: Optional[float] = None + guidance_scale_2: Optional[float] = None + negative_prompt: Optional[str] = None class VideoListResponse(BaseModel): diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/video_api.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/video_api.py index 0b68efaea..72b08e6ed 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/video_api.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/video_api.py @@ -61,20 +61,31 @@ def _build_sampling_params_from_request( request.num_frames if request.num_frames is not None else derived_num_frames ) server_args = get_global_server_args() + sampling_kwargs = { + "request_id": request_id, + "prompt": request.prompt, + "num_frames": num_frames, + "fps": fps, + "width": width, + "height": height, + "image_path": request.input_reference, + "save_output": True, + "output_file_name": request_id, + "seed": request.seed, + "generator_device": request.generator_device, + } + if request.num_inference_steps is not None: + sampling_kwargs["num_inference_steps"] = request.num_inference_steps + if request.guidance_scale is not None: + sampling_kwargs["guidance_scale"] = request.guidance_scale + if request.guidance_scale_2 is not None: + sampling_kwargs["guidance_scale_2"] = request.guidance_scale_2 + if request.negative_prompt is not None: + sampling_kwargs["negative_prompt"] = request.negative_prompt sampling_params = SamplingParams.from_user_sampling_params_args( model_path=server_args.model_path, - request_id=request_id, - prompt=request.prompt, - num_frames=num_frames, - fps=fps, - width=width, - height=height, - image_path=request.input_reference, - save_output=True, server_args=server_args, - output_file_name=request_id, - seed=request.seed, - generator_device=request.generator_device, + **sampling_kwargs, ) return sampling_params