[diffusion] endpoint: fix vertex generate (#17611)
This commit is contained in:
@@ -183,19 +183,28 @@ async def vertex_generate(vertex_req: VertexGenerateReqInput):
|
|||||||
image_input = inst.get("image") or inst.get("image_url")
|
image_input = inst.get("image") or inst.get("image_url")
|
||||||
seed_val = params.get("seed", DEFAULT_SEED)
|
seed_val = params.get("seed", DEFAULT_SEED)
|
||||||
|
|
||||||
|
# Create a dictionary of provided parameters
|
||||||
|
# This filters out None values so the dataclass defaults kick in
|
||||||
|
user_params = {
|
||||||
|
"num_frames": params.get("num_frames"),
|
||||||
|
"fps": params.get("fps"),
|
||||||
|
"width": params.get("width"),
|
||||||
|
"height": params.get("height"),
|
||||||
|
"guidance_scale": params.get("guidance_scale"),
|
||||||
|
"save_output": params.get("save_output"),
|
||||||
|
}
|
||||||
|
|
||||||
|
# Remove None values to allow SamplingParams defaults to take over
|
||||||
|
valid_params = {k: v for k, v in user_params.items() if v is not None}
|
||||||
|
|
||||||
sp = SamplingParams.from_user_sampling_params_args(
|
sp = SamplingParams.from_user_sampling_params_args(
|
||||||
model_path=server_args.model_path,
|
model_path=server_args.model_path,
|
||||||
request_id=rid,
|
request_id=rid,
|
||||||
prompt=prompt,
|
prompt=prompt,
|
||||||
image_path=image_input,
|
image_path=image_input,
|
||||||
num_frames=params.get("num_frames"),
|
|
||||||
fps=params.get("fps"),
|
|
||||||
width=params.get("width"),
|
|
||||||
height=params.get("height"),
|
|
||||||
guidance_scale=params.get("guidance_scale"),
|
|
||||||
seed=seed_val,
|
seed=seed_val,
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
save_output=params.get("save_output"),
|
**valid_params, # Unpack the filtered dictionary
|
||||||
)
|
)
|
||||||
|
|
||||||
backend_req = prepare_request(server_args, sampling_params=sp)
|
backend_req = prepare_request(server_args, sampling_params=sp)
|
||||||
|
|||||||
Reference in New Issue
Block a user