[diffusion] feat: add support for additional sampling parameters in video generation API (#15062)

Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
Xiaoyu Zhang
2025-12-14 19:44:03 +08:00
committed by GitHub
parent fdfabb7afc
commit e3f51e823e
3 changed files with 35 additions and 12 deletions

View File

@@ -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)

View File

@@ -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):

View File

@@ -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