[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:
@@ -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)
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user