[diffusion]Support url image input (#15262)
This commit is contained in:
@@ -21,8 +21,9 @@ from sglang.multimodal_gen.runtime.entrypoints.openai.protocol import (
|
|||||||
from sglang.multimodal_gen.runtime.entrypoints.openai.stores import IMAGE_STORE
|
from sglang.multimodal_gen.runtime.entrypoints.openai.stores import IMAGE_STORE
|
||||||
from sglang.multimodal_gen.runtime.entrypoints.openai.utils import (
|
from sglang.multimodal_gen.runtime.entrypoints.openai.utils import (
|
||||||
_parse_size,
|
_parse_size,
|
||||||
_save_upload_to_path,
|
merge_image_input_list,
|
||||||
process_generation_batch,
|
process_generation_batch,
|
||||||
|
save_image_to_path,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.entrypoints.utils import prepare_request
|
from sglang.multimodal_gen.runtime.entrypoints.utils import prepare_request
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
|
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
|
||||||
@@ -154,6 +155,8 @@ async def generations(
|
|||||||
async def edits(
|
async def edits(
|
||||||
image: Optional[List[UploadFile]] = File(None),
|
image: Optional[List[UploadFile]] = File(None),
|
||||||
image_array: Optional[List[UploadFile]] = File(None, alias="image[]"),
|
image_array: Optional[List[UploadFile]] = File(None, alias="image[]"),
|
||||||
|
url: Optional[List[str]] = Form(None),
|
||||||
|
url_array: Optional[List[str]] = Form(None, alias="url[]"),
|
||||||
prompt: str = Form(...),
|
prompt: str = Form(...),
|
||||||
mask: Optional[UploadFile] = File(None),
|
mask: Optional[UploadFile] = File(None),
|
||||||
model: Optional[str] = Form(None),
|
model: Optional[str] = Form(None),
|
||||||
@@ -173,20 +176,30 @@ async def edits(
|
|||||||
request_id = generate_request_id()
|
request_id = generate_request_id()
|
||||||
# Resolve images from either `image` or `image[]` (OpenAI SDK sends `image[]` when list is provided)
|
# Resolve images from either `image` or `image[]` (OpenAI SDK sends `image[]` when list is provided)
|
||||||
images = image or image_array
|
images = image or image_array
|
||||||
if not images or len(images) == 0:
|
urls = url or url_array
|
||||||
raise HTTPException(status_code=422, detail="Field 'image' is required")
|
|
||||||
|
if (not images or len(images) == 0) and (not urls or len(urls) == 0):
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=422, detail="Field 'image' or 'url' is required"
|
||||||
|
)
|
||||||
|
|
||||||
# Save all input images; additional images beyond the first are saved for potential future use
|
# Save all input images; additional images beyond the first are saved for potential future use
|
||||||
uploads_dir = os.path.join("outputs", "uploads")
|
uploads_dir = os.path.join("outputs", "uploads")
|
||||||
os.makedirs(uploads_dir, exist_ok=True)
|
os.makedirs(uploads_dir, exist_ok=True)
|
||||||
if images is not None and not isinstance(images, list):
|
image_list = merge_image_input_list(images, urls)
|
||||||
images = [images]
|
|
||||||
input_paths = []
|
input_paths = []
|
||||||
for idx, img in enumerate(images):
|
try:
|
||||||
filename = img.filename or f"image_{idx}"
|
for idx, img in enumerate(image_list):
|
||||||
input_path = os.path.join(uploads_dir, f"{request_id}_{idx}_{filename}")
|
filename = img.filename if hasattr(img, "filename") else f"image_{idx}"
|
||||||
await _save_upload_to_path(img, input_path)
|
input_path = await save_image_to_path(
|
||||||
input_paths.append(input_path)
|
img, os.path.join(uploads_dir, f"{request_id}_{idx}_{filename}")
|
||||||
|
)
|
||||||
|
input_paths.append(input_path)
|
||||||
|
except Exception as e:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=400, detail=f"Failed to process image source: {str(e)}"
|
||||||
|
)
|
||||||
|
|
||||||
sampling = _build_sampling_params_from_request(
|
sampling = _build_sampling_params_from_request(
|
||||||
request_id=request_id,
|
request_id=request_id,
|
||||||
|
|||||||
@@ -1,9 +1,12 @@
|
|||||||
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
||||||
|
import base64
|
||||||
import dataclasses
|
import dataclasses
|
||||||
import os
|
import os
|
||||||
|
import re
|
||||||
import time
|
import time
|
||||||
from typing import Optional
|
from typing import Any, List, Optional, Union
|
||||||
|
|
||||||
|
import httpx
|
||||||
from fastapi import UploadFile
|
from fastapi import UploadFile
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.entrypoints.utils import post_process_sample
|
from sglang.multimodal_gen.runtime.entrypoints.utils import post_process_sample
|
||||||
@@ -44,6 +47,13 @@ def _parse_size(size: str) -> tuple[int, int] | tuple[None, None]:
|
|||||||
return None, None
|
return None, None
|
||||||
|
|
||||||
|
|
||||||
|
async def save_image_to_path(image: Union[UploadFile, str], target_path: str) -> str:
|
||||||
|
input_path = await _maybe_url_image(image, target_path)
|
||||||
|
if input_path is None:
|
||||||
|
input_path = await _save_upload_to_path(image, target_path)
|
||||||
|
return input_path
|
||||||
|
|
||||||
|
|
||||||
# Helpers
|
# Helpers
|
||||||
async def _save_upload_to_path(upload: UploadFile, target_path: str) -> str:
|
async def _save_upload_to_path(upload: UploadFile, target_path: str) -> str:
|
||||||
os.makedirs(os.path.dirname(target_path), exist_ok=True)
|
os.makedirs(os.path.dirname(target_path), exist_ok=True)
|
||||||
@@ -53,6 +63,98 @@ async def _save_upload_to_path(upload: UploadFile, target_path: str) -> str:
|
|||||||
return target_path
|
return target_path
|
||||||
|
|
||||||
|
|
||||||
|
async def _maybe_url_image(img_url: str, target_path: str) -> str:
|
||||||
|
if not isinstance(img_url, str):
|
||||||
|
return None
|
||||||
|
|
||||||
|
if img_url.lower().startswith(("http://", "https://")):
|
||||||
|
# Download image from URL
|
||||||
|
input_path = await _save_url_image_to_path(img_url, target_path)
|
||||||
|
return input_path
|
||||||
|
elif img_url.startswith("data:image"):
|
||||||
|
# encode image base64 url
|
||||||
|
input_path = await _save_base64_image_to_path(img_url, target_path)
|
||||||
|
return input_path
|
||||||
|
else:
|
||||||
|
raise ValueError("Unsupported image url format")
|
||||||
|
|
||||||
|
|
||||||
|
async def _save_url_image_to_path(image_url: str, target_path: str) -> str:
|
||||||
|
"""Download image from URL and save to target path."""
|
||||||
|
|
||||||
|
os.makedirs(os.path.dirname(target_path), exist_ok=True)
|
||||||
|
|
||||||
|
try:
|
||||||
|
async with httpx.AsyncClient() as client:
|
||||||
|
response = await client.get(image_url, timeout=10.0)
|
||||||
|
response.raise_for_status()
|
||||||
|
|
||||||
|
# Determine file extension from content type or URL after downloading
|
||||||
|
if not os.path.splitext(target_path)[1]:
|
||||||
|
content_type = response.headers.get("content-type", "")
|
||||||
|
if not content_type.startswith("image/"):
|
||||||
|
raise ValueError(
|
||||||
|
f"URL does not point to an image. Content-Type: {content_type}"
|
||||||
|
)
|
||||||
|
if "jpeg" in content_type or "jpg" in content_type:
|
||||||
|
ext = ".jpg"
|
||||||
|
elif "png" in content_type:
|
||||||
|
ext = ".png"
|
||||||
|
elif "webp" in content_type:
|
||||||
|
ext = ".webp"
|
||||||
|
else:
|
||||||
|
ext = ".jpg" # Default to jpg
|
||||||
|
target_path = f"{target_path}{ext}"
|
||||||
|
|
||||||
|
with open(target_path, "wb") as f:
|
||||||
|
f.write(response.content)
|
||||||
|
|
||||||
|
return target_path
|
||||||
|
except Exception as e:
|
||||||
|
raise Exception(f"Failed to download image from URL: {str(e)}")
|
||||||
|
|
||||||
|
|
||||||
|
async def _save_base64_image_to_path(base64_data: str, target_path: str) -> str:
|
||||||
|
"""Decode base64 image data and save to target path."""
|
||||||
|
|
||||||
|
# split `data:[<media-type>][;base64],<data>` to media-type base64 data
|
||||||
|
pattern = r"data:(.*?)(;base64)?,(.*)"
|
||||||
|
match = re.match(pattern, base64_data)
|
||||||
|
if not match:
|
||||||
|
raise ValueError(
|
||||||
|
f"Failed to decoding base64 image, please make sure the url format `data:[<media-type>][;base64],<data>` "
|
||||||
|
)
|
||||||
|
media_type = match.group(1)
|
||||||
|
is_base64 = match.group(2)
|
||||||
|
if not is_base64:
|
||||||
|
raise ValueError(
|
||||||
|
f"Failed to decoding base64 image, please make sure the url format `data:[<media-type>][;base64],<data>` "
|
||||||
|
)
|
||||||
|
data = match.group(3)
|
||||||
|
if not data:
|
||||||
|
raise ValueError(
|
||||||
|
f"Failed to decoding base64 image, please make sure the url format `data:[<media-type>][;base64],<data>` "
|
||||||
|
)
|
||||||
|
# get ext from url
|
||||||
|
if media_type.startswith("image/"):
|
||||||
|
ext = media_type.split("/")[-1].lower()
|
||||||
|
if ext == "jpeg":
|
||||||
|
ext = "jpg"
|
||||||
|
else:
|
||||||
|
ext = "jpg"
|
||||||
|
target_path = f"{target_path}.{ext}"
|
||||||
|
os.makedirs(os.path.dirname(target_path), exist_ok=True)
|
||||||
|
|
||||||
|
try:
|
||||||
|
image_data = base64.b64decode(data)
|
||||||
|
with open(target_path, "wb") as f:
|
||||||
|
f.write(image_data)
|
||||||
|
|
||||||
|
return target_path
|
||||||
|
except Exception as e:
|
||||||
|
raise Exception(f"Failed to decode base64 image: {str(e)}")
|
||||||
|
|
||||||
|
|
||||||
async def process_generation_batch(
|
async def process_generation_batch(
|
||||||
scheduler_client,
|
scheduler_client,
|
||||||
batch,
|
batch,
|
||||||
@@ -77,3 +179,31 @@ async def process_generation_batch(
|
|||||||
log_batch_completion(logger, 1, total_time)
|
log_batch_completion(logger, 1, total_time)
|
||||||
|
|
||||||
return save_file_path
|
return save_file_path
|
||||||
|
|
||||||
|
|
||||||
|
def merge_image_input_list(*inputs: Union[List, Any, None]) -> List:
|
||||||
|
"""
|
||||||
|
Merge multiple image input sources into a single list.
|
||||||
|
|
||||||
|
This function handles both single items and lists of items, merging them
|
||||||
|
into a single flattened list. Useful for processing images, URLs, or other
|
||||||
|
multimedia inputs that can come as either single items or lists.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
*inputs: Variable number of inputs, each can be None, single item, or list
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
List: Flattened list of all non-None inputs
|
||||||
|
|
||||||
|
Example:
|
||||||
|
>>> merge_image_input_list(["img1", "img2"], "img3", None)
|
||||||
|
["img1", "img2", "img3"]
|
||||||
|
"""
|
||||||
|
result = []
|
||||||
|
for input_item in inputs:
|
||||||
|
if input_item is not None:
|
||||||
|
if isinstance(input_item, list):
|
||||||
|
result.extend(input_item)
|
||||||
|
else:
|
||||||
|
result.append(input_item)
|
||||||
|
return result
|
||||||
|
|||||||
@@ -30,8 +30,9 @@ from sglang.multimodal_gen.runtime.entrypoints.openai.protocol import (
|
|||||||
from sglang.multimodal_gen.runtime.entrypoints.openai.stores import VIDEO_STORE
|
from sglang.multimodal_gen.runtime.entrypoints.openai.stores import VIDEO_STORE
|
||||||
from sglang.multimodal_gen.runtime.entrypoints.openai.utils import (
|
from sglang.multimodal_gen.runtime.entrypoints.openai.utils import (
|
||||||
_parse_size,
|
_parse_size,
|
||||||
_save_upload_to_path,
|
merge_image_input_list,
|
||||||
process_generation_batch,
|
process_generation_batch,
|
||||||
|
save_image_to_path,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.entrypoints.utils import prepare_request
|
from sglang.multimodal_gen.runtime.entrypoints.utils import prepare_request
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
|
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
|
||||||
@@ -136,6 +137,7 @@ async def create_video(
|
|||||||
# multipart/form-data fields (optional; used only when content-type is multipart)
|
# multipart/form-data fields (optional; used only when content-type is multipart)
|
||||||
prompt: Optional[str] = Form(None),
|
prompt: Optional[str] = Form(None),
|
||||||
input_reference: Optional[UploadFile] = File(None),
|
input_reference: Optional[UploadFile] = File(None),
|
||||||
|
reference_url: Optional[str] = Form(None),
|
||||||
model: Optional[str] = Form(None),
|
model: Optional[str] = Form(None),
|
||||||
seconds: Optional[int] = Form(None),
|
seconds: Optional[int] = Form(None),
|
||||||
size: Optional[str] = Form(None),
|
size: Optional[str] = Form(None),
|
||||||
@@ -155,17 +157,24 @@ async def create_video(
|
|||||||
if "multipart/form-data" in content_type:
|
if "multipart/form-data" in content_type:
|
||||||
if not prompt:
|
if not prompt:
|
||||||
raise HTTPException(status_code=400, detail="prompt is required")
|
raise HTTPException(status_code=400, detail="prompt is required")
|
||||||
if input_reference is None:
|
if input_reference is None and reference_url is None:
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
status_code=400, detail="input_reference file is required"
|
status_code=400,
|
||||||
|
detail="input_reference file or reference_url is required",
|
||||||
)
|
)
|
||||||
|
image_list = merge_image_input_list(input_reference, reference_url)
|
||||||
|
# Save first input image
|
||||||
|
image = image_list[0]
|
||||||
uploads_dir = os.path.join("outputs", "uploads")
|
uploads_dir = os.path.join("outputs", "uploads")
|
||||||
os.makedirs(uploads_dir, exist_ok=True)
|
os.makedirs(uploads_dir, exist_ok=True)
|
||||||
input_path = os.path.join(
|
filename = image.filename if hasattr(image, "filename") else f"url_image"
|
||||||
uploads_dir, f"{request_id}_{input_reference.filename}"
|
input_path = os.path.join(uploads_dir, f"{request_id}_{filename}")
|
||||||
)
|
try:
|
||||||
await _save_upload_to_path(input_reference, input_path)
|
input_path = await save_image_to_path(image, input_path)
|
||||||
|
except Exception as e:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=400, detail=f"Failed to process image source: {str(e)}"
|
||||||
|
)
|
||||||
|
|
||||||
# Parse extra_body JSON (if provided in multipart form) to get fps/num_frames overrides
|
# Parse extra_body JSON (if provided in multipart form) to get fps/num_frames overrides
|
||||||
extra_from_form: Dict[str, Any] = {}
|
extra_from_form: Dict[str, Any] = {}
|
||||||
@@ -207,6 +216,29 @@ async def create_video(
|
|||||||
if isinstance(extra, dict):
|
if isinstance(extra, dict):
|
||||||
# Shallow-merge: only keys like fps/num_frames are expected
|
# Shallow-merge: only keys like fps/num_frames are expected
|
||||||
payload.update(extra)
|
payload.update(extra)
|
||||||
|
# openai may turn extra_body to extra_json
|
||||||
|
extra_json = payload.pop("extra_json", None)
|
||||||
|
if isinstance(extra_json, dict):
|
||||||
|
payload.update(extra_json)
|
||||||
|
# for not multipart/form-data type
|
||||||
|
if payload.get("reference_url"):
|
||||||
|
image_list = merge_image_input_list(payload.get("reference_url"))
|
||||||
|
# Save first input image
|
||||||
|
image = image_list[0]
|
||||||
|
uploads_dir = os.path.join("outputs", "uploads")
|
||||||
|
os.makedirs(uploads_dir, exist_ok=True)
|
||||||
|
filename = (
|
||||||
|
image.filename if hasattr(image, "filename") else f"url_image"
|
||||||
|
)
|
||||||
|
input_path = os.path.join(uploads_dir, f"{request_id}_{filename}")
|
||||||
|
try:
|
||||||
|
input_path = await save_image_to_path(image, input_path)
|
||||||
|
except Exception as e:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=400,
|
||||||
|
detail=f"Failed to process image source: {str(e)}",
|
||||||
|
)
|
||||||
|
payload["input_reference"] = input_path
|
||||||
req = VideoGenerationsRequest(**payload)
|
req = VideoGenerationsRequest(**payload)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
raise HTTPException(status_code=400, detail=f"Invalid request body: {e}")
|
raise HTTPException(status_code=400, detail=f"Invalid request body: {e}")
|
||||||
|
|||||||
@@ -671,6 +671,7 @@ def get_generate_fn(
|
|||||||
prompt: str | None = None,
|
prompt: str | None = None,
|
||||||
seconds: int | None = None,
|
seconds: int | None = None,
|
||||||
input_reference: Any | None = None,
|
input_reference: Any | None = None,
|
||||||
|
extra_body: dict[Any] | None = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
"""
|
"""
|
||||||
Create a video job via /v1/videos, poll until completion,
|
Create a video job via /v1/videos, poll until completion,
|
||||||
@@ -687,6 +688,8 @@ def get_generate_fn(
|
|||||||
create_kwargs["seconds"] = seconds
|
create_kwargs["seconds"] = seconds
|
||||||
if input_reference is not None:
|
if input_reference is not None:
|
||||||
create_kwargs["input_reference"] = input_reference # triggers multipart
|
create_kwargs["input_reference"] = input_reference # triggers multipart
|
||||||
|
if extra_body is not None:
|
||||||
|
create_kwargs["extra_body"] = extra_body
|
||||||
|
|
||||||
job = client.videos.create(**create_kwargs) # type: ignore[attr-defined]
|
job = client.videos.create(**create_kwargs) # type: ignore[attr-defined]
|
||||||
video_id = job.id
|
video_id = job.id
|
||||||
@@ -839,6 +842,53 @@ def get_generate_fn(
|
|||||||
|
|
||||||
return rid
|
return rid
|
||||||
|
|
||||||
|
def generate_image_edit_url(case_id, client) -> str:
|
||||||
|
"""TI2I: Text + Image ? Image edit using direct URL transfer (no pre-download)."""
|
||||||
|
if not sampling_params.prompt or not sampling_params.image_path:
|
||||||
|
pytest.skip(f"{id}: no edit config")
|
||||||
|
|
||||||
|
# Handle both single URL and list of URLs
|
||||||
|
image_urls = sampling_params.image_path
|
||||||
|
if not isinstance(image_urls, list):
|
||||||
|
image_urls = [image_urls]
|
||||||
|
|
||||||
|
# Validate all URLs
|
||||||
|
for url in image_urls:
|
||||||
|
if not is_image_url(url):
|
||||||
|
pytest.skip(
|
||||||
|
f"{id}: image_path must be a URL for URL direct test: {url}"
|
||||||
|
)
|
||||||
|
|
||||||
|
response = client.images.with_raw_response.edit(
|
||||||
|
model=model_path,
|
||||||
|
prompt=sampling_params.prompt,
|
||||||
|
image=[], # Only for OpenAI verification
|
||||||
|
n=1,
|
||||||
|
size=sampling_params.output_size,
|
||||||
|
response_format="b64_json",
|
||||||
|
extra_body={"url": image_urls},
|
||||||
|
)
|
||||||
|
|
||||||
|
rid = response.headers.get("x-request-id", "")
|
||||||
|
result = response.parse()
|
||||||
|
validate_image(result.data[0].b64_json)
|
||||||
|
|
||||||
|
# Save and upload result for verification
|
||||||
|
img_data = base64.b64decode(result.data[0].b64_json)
|
||||||
|
tmp_path = f"{rid}.png"
|
||||||
|
with open(tmp_path, "wb") as f:
|
||||||
|
f.write(img_data)
|
||||||
|
upload_file_to_slack(
|
||||||
|
case_id=case_id,
|
||||||
|
model=model_path,
|
||||||
|
prompt=sampling_params.prompt,
|
||||||
|
file_path=tmp_path,
|
||||||
|
origin_file_path=str(sampling_params.image_path),
|
||||||
|
)
|
||||||
|
os.remove(tmp_path)
|
||||||
|
|
||||||
|
return rid
|
||||||
|
|
||||||
def generate_video(case_id, client) -> str:
|
def generate_video(case_id, client) -> str:
|
||||||
"""T2V: Text ? Video."""
|
"""T2V: Text ? Video."""
|
||||||
if not sampling_params.prompt:
|
if not sampling_params.prompt:
|
||||||
@@ -876,6 +926,19 @@ def get_generate_fn(
|
|||||||
input_reference=fh,
|
input_reference=fh,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def generate_text_url_image_to_video(case_id, client) -> str:
|
||||||
|
if not sampling_params.prompt or not sampling_params.image_path:
|
||||||
|
pytest.skip(f"{id}: no edit config")
|
||||||
|
return _create_and_download_video(
|
||||||
|
client,
|
||||||
|
case_id,
|
||||||
|
model=model_path,
|
||||||
|
prompt=sampling_params.prompt,
|
||||||
|
size=sampling_params.output_size,
|
||||||
|
seconds=video_seconds,
|
||||||
|
extra_body={"reference_url": sampling_params.image_path},
|
||||||
|
)
|
||||||
|
|
||||||
def generate_text_image_to_video(case_id, client) -> str:
|
def generate_text_image_to_video(case_id, client) -> str:
|
||||||
"""TI2V: Text + Image ? Video."""
|
"""TI2V: Text + Image ? Video."""
|
||||||
if not sampling_params.prompt or not sampling_params.image_path:
|
if not sampling_params.prompt or not sampling_params.image_path:
|
||||||
@@ -901,13 +964,19 @@ def get_generate_fn(
|
|||||||
|
|
||||||
if modality == "video":
|
if modality == "video":
|
||||||
if sampling_params.image_path and sampling_params.prompt:
|
if sampling_params.image_path and sampling_params.prompt:
|
||||||
fn = generate_text_image_to_video
|
if getattr(sampling_params, "direct_url_test", False):
|
||||||
|
fn = generate_text_url_image_to_video
|
||||||
|
else:
|
||||||
|
fn = generate_text_image_to_video
|
||||||
elif sampling_params.image_path:
|
elif sampling_params.image_path:
|
||||||
fn = generate_image_to_video
|
fn = generate_image_to_video
|
||||||
else:
|
else:
|
||||||
fn = generate_video
|
fn = generate_video
|
||||||
elif sampling_params.prompt and sampling_params.image_path:
|
elif sampling_params.prompt and sampling_params.image_path:
|
||||||
fn = generate_image_edit
|
if getattr(sampling_params, "direct_url_test", False):
|
||||||
|
fn = generate_image_edit_url
|
||||||
|
else:
|
||||||
|
fn = generate_image_edit
|
||||||
else:
|
else:
|
||||||
fn = generate_image
|
fn = generate_image
|
||||||
|
|
||||||
|
|||||||
@@ -141,6 +141,9 @@ class DiffusionSamplingParams:
|
|||||||
num_frames: int | None = None # for video: number of frames
|
num_frames: int | None = None # for video: number of frames
|
||||||
fps: int | None = None # for video: frames per second
|
fps: int | None = None # for video: frames per second
|
||||||
|
|
||||||
|
# URL direct test flag - if True, don't pre-download URL images
|
||||||
|
direct_url_test: bool = False
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
class DiffusionTestCase:
|
class DiffusionTestCase:
|
||||||
@@ -233,6 +236,7 @@ MULTI_IMAGE_TI2I_sampling_params = DiffusionSamplingParams(
|
|||||||
"https://qianwen-res.oss-cn-beijing.aliyuncs.com/Qwen-Image/edit2509/edit2509_1.jpg",
|
"https://qianwen-res.oss-cn-beijing.aliyuncs.com/Qwen-Image/edit2509/edit2509_1.jpg",
|
||||||
"https://qianwen-res.oss-cn-beijing.aliyuncs.com/Qwen-Image/edit2509/edit2509_2.jpg",
|
"https://qianwen-res.oss-cn-beijing.aliyuncs.com/Qwen-Image/edit2509/edit2509_2.jpg",
|
||||||
],
|
],
|
||||||
|
direct_url_test=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
T2V_PROMPT = "A curious raccoon"
|
T2V_PROMPT = "A curious raccoon"
|
||||||
@@ -240,6 +244,7 @@ T2V_PROMPT = "A curious raccoon"
|
|||||||
TI2V_sampling_params = DiffusionSamplingParams(
|
TI2V_sampling_params = DiffusionSamplingParams(
|
||||||
prompt="The man in the picture slowly turns his head, his expression enigmatic and otherworldly. The camera performs a slow, cinematic dolly out, focusing on his face. Moody lighting, neon signs glowing in the background, shallow depth of field.",
|
prompt="The man in the picture slowly turns his head, his expression enigmatic and otherworldly. The camera performs a slow, cinematic dolly out, focusing on his face. Moody lighting, neon signs glowing in the background, shallow depth of field.",
|
||||||
image_path="https://is1-ssl.mzstatic.com/image/thumb/Music114/v4/5f/fa/56/5ffa56c2-ea1f-7a17-6bad-192ff9b6476d/825646124206.jpg/600x600bb.jpg",
|
image_path="https://is1-ssl.mzstatic.com/image/thumb/Music114/v4/5f/fa/56/5ffa56c2-ea1f-7a17-6bad-192ff9b6476d/825646124206.jpg/600x600bb.jpg",
|
||||||
|
direct_url_test=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
# All test cases with clean default values
|
# All test cases with clean default values
|
||||||
|
|||||||
Reference in New Issue
Block a user