Compare commits

...
Author SHA1 Message Date
Y-aang c1d7c51bc7 fixi2v demo bug 2025-11-19 20:53:43 +00:00
Y-aang 5de2d80ebd add full i2v demo 2025-11-19 15:40:04 +00:00
RandNMR73 6cd227349c add demo prompts 2025-11-19 05:49:01 +00:00
RandNMR73 7e15e5dab1 add demo images 2025-11-19 05:42:37 +00:00
RandNMR73 c9ee4f8bf8 add i2v demo 2025-11-18 09:40:05 +00:00
SolitaryThinker f764b43aaa nit 2025-11-16 23:10:47 +00:00
SolitaryThinker 592af8c954 fix lint 2025-11-16 23:09:04 +00:00
SolitaryThinker 4225a96a40 use new ckpt 2025-11-16 12:49:06 +00:00
JerryZhou54 733c14a2cb Change to fp32 inference 2025-11-16 04:53:57 +00:00
JerryZhou54 81b302409e Add i2v images 2025-11-16 04:51:14 +00:00
JerryZhou54 7fbd0cfec7 Fi lint 2025-11-15 23:50:04 +00:00
RandNMR73 85b9934079 Add inference for MoE SF 2025-11-15 23:44:21 +00:00
44 changed files with 922 additions and 444 deletions
+1
View File
@@ -10,6 +10,7 @@ exclude: |
demo/.*|
predict\.py|
scripts/.*|
prompts/.*|
fastvideo/data_preprocess/.*|
fastvideo/dataset/.*|
fastvideo/models/.*|
@@ -0,0 +1,44 @@
# NOTE: This is still a work in progress, and the checkpoints are not released yet.
from fastvideo import VideoGenerator, SamplingParam
import json
# from fastvideo.configs.sample import SamplingParam
OUTPUT_PATH = "video_samples_self_forcing_causal_wan2_2_14B_i2v"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
generator = VideoGenerator.from_pretrained(
"FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=True,
dit_cpu_offload=True, # DiT need to be offloaded for MoE
dit_precision="fp32",
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
dmd_denoising_steps=[1000, 850, 700, 550, 350, 275, 200, 125],
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
pin_cpu_memory=True,
# image_encoder_cpu_offload=False,
)
sampling_param = SamplingParam.from_pretrained("FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers")
sampling_param.num_frames = 73
sampling_param.width = 832
sampling_param.height = 480
sampling_param.seed = 1000
with open("prompts/mixkit_i2v.jsonl", "r") as f:
prompt_image_pairs = json.load(f)
for prompt_image_pair in prompt_image_pairs:
prompt = prompt_image_pair["prompt"]
image_path = prompt_image_pair["image_path"]
_ = generator.generate_video(prompt, image_path=image_path, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
if __name__ == "__main__":
main()
@@ -3,8 +3,13 @@ import os
import requests
import base64
import time
import json
from pathlib import Path
import tempfile
from io import BytesIO
import gradio as gr
from PIL import Image
from fastvideo.configs.sample.base import SamplingParam
@@ -12,6 +17,7 @@ from fastvideo.configs.sample.base import SamplingParam
MODEL_PATH_MAPPING = {
"FastWan2.1-T2V-1.3B": "FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
"FastWan2.2-TI2V-5B-FullAttn": "FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
"CausalWan2.2-I2V-A14B-Preview": "FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
}
@@ -37,7 +43,7 @@ class RayServeClient:
f"{self.backend_url}/generate_video",
json=request_data,
headers=headers,
timeout=300
timeout=900 # 15 minutes timeout for longer video generation
)
round_trip_time = time.time() - start_time
@@ -81,49 +87,89 @@ def save_video_from_base64(video_data: str, output_dir: str, prompt: str) -> str
return None
def create_timing_display(inference_time, encoding_time, network_time, total_time, stage_execution_times, num_frames):
dit_denoising_time = f"{stage_execution_times[5]:.2f}s" if len(stage_execution_times) > 5 else "N/A"
def encode_image_to_base64(image_input) -> str:
"""Encode an image file path or in-memory image to a base64 string."""
if image_input is None:
return None
timing_html = f"""
<div style="margin: 10px 0;">
<h3 style="text-align: center; margin-bottom: 10px;">⏱️ Timing Breakdown</h3>
<div style="display: grid; grid-template-columns: repeat(5, 1fr); gap: 10px; margin-bottom: 10px;">
<div class="timing-card timing-card-highlight">
<div style="font-size: 20px;">🚀</div>
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">DiT Denoising</div>
<div style="font-size: 16px; color: #ffa200; font-weight: bold;">{dit_denoising_time}</div>
</div>
<div class="timing-card">
<div style="font-size: 20px;">🧠</div>
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">E2E (w. vae/text encoder)</div>
<div style="font-size: 16px; color: #2563eb;">{inference_time:.2f}s</div>
</div>
<div class="timing-card">
<div style="font-size: 20px;">🎬</div>
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Video Encoding</div>
<div style="font-size: 16px; color: #dc2626;">{encoding_time:.2f}s</div>
</div>
<div class="timing-card">
<div style="font-size: 20px;">🌐</div>
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Network Transfer</div>
<div style="font-size: 16px; color: #059669;">{network_time:.2f}s</div>
</div>
<div class="timing-card">
<div style="font-size: 20px;">📊</div>
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Total Processing</div>
<div style="font-size: 18px; color: #0277bd;">{total_time:.2f}s</div>
</div>
</div>"""
mime_types = {
'.jpg': 'image/jpeg',
'.jpeg': 'image/jpeg',
'.png': 'image/png',
'.gif': 'image/gif',
'.webp': 'image/webp',
}
if inference_time > 0:
fps = num_frames / inference_time
timing_html += f"""
<div class="performance-card" style="margin-top: 15px;">
<span style="font-weight: bold;">Generation Speed: </span>
<span style="font-size: 18px; color: #6366f1; font-weight: bold;">{fps:.1f} frames/second</span>
</div>"""
return timing_html + "</div>"
try:
if isinstance(image_input, str):
if not os.path.exists(image_input):
return None
with open(image_input, 'rb') as f:
image_bytes = f.read()
ext = os.path.splitext(image_input)[1].lower()
mime_type = mime_types.get(ext, 'image/jpeg')
elif isinstance(image_input, Image.Image):
buffer = BytesIO()
image_to_save = image_input.convert("RGB")
image_to_save.save(buffer, format="PNG")
image_bytes = buffer.getvalue()
mime_type = 'image/png'
else:
return None
image_base64 = base64.b64encode(image_bytes).decode('utf-8')
return f"data:{mime_type};base64,{image_base64}"
except Exception as e:
print(f"Failed to encode image: {e}")
return None
# def create_timing_display(inference_time, encoding_time, network_time, total_time, stage_execution_times, num_frames):
# dit_denoising_time = f"{stage_execution_times[5]:.2f}s" if len(stage_execution_times) > 5 else "N/A"
#
# timing_html = f"""
# <div style="margin: 10px 0;">
# <h3 style="text-align: center; margin-bottom: 10px;">⏱️ Timing Breakdown</h3>
# <div style="display: grid; grid-template-columns: repeat(5, 1fr); gap: 10px; margin-bottom: 10px;">
# <div class="timing-card timing-card-highlight">
# <div style="font-size: 20px;">🚀</div>
# <div style="font-weight: bold; margin: 3px 0; font-size: 14px;">DiT Denoising</div>
# <div style="font-size: 16px; color: #ffa200; font-weight: bold;">{dit_denoising_time}</div>
# </div>
# <div class="timing-card">
# <div style="font-size: 20px;">🧠</div>
# <div style="font-weight: bold; margin: 3px 0; font-size: 14px;">E2E (w. vae/text encoder)</div>
# <div style="font-size: 16px; color: #2563eb;">{inference_time:.2f}s</div>
# </div>
# <div class="timing-card">
# <div style="font-size: 20px;">🎬</div>
# <div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Video Encoding</div>
# <div style="font-size: 16px; color: #dc2626;">{encoding_time:.2f}s</div>
# </div>
# <div class="timing-card">
# <div style="font-size: 20px;">🌐</div>
# <div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Network Transfer</div>
# <div style="font-size: 16px; color: #059669;">{network_time:.2f}s</div>
# </div>
# <div class="timing-card">
# <div style="font-size: 20px;">📊</div>
# <div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Total Processing</div>
# <div style="font-size: 18px; color: #0277bd;">{total_time:.2f}s</div>
# </div>
# </div>"""
#
# if inference_time > 0:
# fps = num_frames / inference_time
# timing_html += f"""
# <div class="performance-card" style="margin-top: 15px;">
# <span style="font-weight: bold;">Generation Speed: </span>
# <span style="font-size: 18px; color: #6366f1; font-weight: bold;">{fps:.1f} frames/second</span>
# </div>"""
#
# return timing_html + "</div>"
def load_example_prompts():
@@ -144,26 +190,83 @@ def load_example_prompts():
print(f"Warning: Could not read {filepath}: {e}")
return prompts, labels
examples, example_labels = load_from_file("prompts/prompts_final.txt")
# Load prompts from prompts.txt
examples, example_labels = load_from_file("examples/inference/gradio/serving/prompts.txt")
if not examples:
examples = ["A crowded rooftop bar buzzes with energy, the city skyline twinkling like a field of stars in the background."]
example_labels = ["Crowded rooftop bar at night"]
return examples, example_labels
# Load image mappings from JSON file
prompt_to_image = {}
# Try to find the JSON file relative to project root
possible_json_paths = [
Path("prompts/mixkit_i2v.jsonl"),
Path(__file__).parent.parent.parent.parent / "prompts" / "mixkit_i2v.jsonl",
]
json_path = None
for path in possible_json_paths:
if path.exists():
json_path = path
break
if json_path and json_path.exists():
try:
with open(json_path, "r", encoding='utf-8') as f:
data = json.load(f)
# Get the project root directory (parent of prompts directory)
project_root = json_path.parent.parent
for item in data:
prompt_text = item.get("prompt", "").strip()
image_path = item.get("image_path", "")
if prompt_text and image_path:
# Resolve image path relative to project root
full_image_path = project_root / image_path
if full_image_path.exists():
prompt_to_image[prompt_text] = str(full_image_path.absolute())
except Exception as e:
print(f"Warning: Could not load image mappings from {json_path}: {e}")
# Create image paths list matching the prompts
example_images = []
for prompt in examples:
# Try exact match first
image_path = prompt_to_image.get(prompt)
if not image_path:
# Try fuzzy match (case-insensitive, whitespace normalized)
normalized_prompt = " ".join(prompt.split())
for json_prompt, img_path in prompt_to_image.items():
normalized_json = " ".join(json_prompt.split())
if normalized_prompt.lower() == normalized_json.lower():
image_path = img_path
break
example_images.append(image_path if image_path and os.path.exists(image_path) else None)
return examples, example_labels, example_images
def create_gradio_interface(backend_url: str, default_params: dict[str, SamplingParam]):
client = RayServeClient(backend_url)
def is_i2v_model(model_name: str) -> bool:
"""Check if the model is an I2V model."""
return "I2V" in model_name
def generate_video(
prompt, negative_prompt, use_negative_prompt, seed, guidance_scale,
num_frames, height, width, randomize_seed, model_selection, progress
prompt, negative_prompt, use_negative_prompt, guidance_scale,
num_frames, height, width, model_selection, input_image, progress
):
# Use default seed value (randomize_seed disabled)
seed = 1000
randomize_seed = False
if not client.check_health():
return None, f"Backend is not available. Please check if Ray Serve is running at {backend_url}", ""
# Check if I2V model requires an image
if is_i2v_model(model_selection) and not input_image:
return None, "I2V models require an input image. Please upload an image.", ""
# Validate dimensions
max_pixels = 720 * 1280
if height * width > max_pixels:
@@ -172,6 +275,15 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
if progress:
progress(0.1, desc="Checking backend health...")
# Encode image if provided
image_data = None
if input_image:
if progress:
progress(0.2, desc="Encoding input image...")
image_data = encode_image_to_base64(input_image)
if not image_data:
return None, "Failed to encode input image", ""
request_data = {
"prompt": prompt,
"negative_prompt": negative_prompt,
@@ -183,7 +295,7 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
"width": width,
"randomize_seed": randomize_seed,
"return_frames": False,
"image_path": None,
"image_data": image_data,
"model_path": MODEL_PATH_MAPPING.get(model_selection, "FastVideo/FastWan2.1-T2V-1.3B-Diffusers")
}
@@ -198,16 +310,16 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
if response.get("success", False):
video_data = response.get("video_data", "")
used_seed = response.get("seed", seed)
inference_time = response.get("inference_time", 0.0)
encoding_time = response.get("encoding_time", 0.0)
total_time = response.get("total_time", 0.0)
network_time = response.get("network_time", 0.0)
stage_execution_times = response.get("stage_execution_times", [])
# inference_time = response.get("inference_time", 0.0)
# encoding_time = response.get("encoding_time", 0.0)
# total_time = response.get("total_time", 0.0)
# network_time = response.get("network_time", 0.0)
# stage_execution_times = response.get("stage_execution_times", [])
timing_details = create_timing_display(
inference_time, encoding_time, network_time, total_time,
stage_execution_times, num_frames
)
# timing_details = create_timing_display(
# inference_time, encoding_time, network_time, total_time,
# stage_execution_times, num_frames
# )
if video_data:
if progress:
@@ -219,7 +331,7 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
progress(1.0, desc="Generation complete!")
if video_path and os.path.exists(video_path):
return video_path, used_seed, timing_details
return video_path, used_seed, ""
else:
return None, "Failed to save video", ""
else:
@@ -228,7 +340,7 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
error_msg = response.get("error_message", "Unknown error occurred")
return None, f"Generation failed: {error_msg}", ""
examples, example_labels = load_example_prompts()
examples, example_labels, example_images = load_example_prompts()
theme = gr.themes.Base().set(
button_primary_background_fill="#2563eb",
@@ -239,33 +351,39 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
)
def get_default_values(model_name):
model_path = MODEL_PATH_MAPPING.get(model_name)
if model_path and model_path in default_params:
params = default_params[model_path]
return {
'height': params.height,
'width': params.width,
'num_frames': params.num_frames,
'guidance_scale': params.guidance_scale,
'seed': params.seed,
}
# model_path = MODEL_PATH_MAPPING.get(model_name)
# if model_path and model_path in default_params:
# params = default_params[model_path]
# return {
# 'height': params.height,
# 'width': params.width,
# 'num_frames': params.num_frames,
# 'guidance_scale': params.guidance_scale,
# }
return {
'height': 448,
'height': 480,
'width': 832,
'num_frames': 61,
'guidance_scale': 3.0,
'seed': 1024,
'num_frames': 73,
}
initial_values = get_default_values("FastWan2.1-T2V-1.3B")
# Get available models based on what's loaded
available_models = []
for model_name, model_path in MODEL_PATH_MAPPING.items():
if model_path in default_params:
available_models.append(model_name)
with gr.Blocks(title="FastWan", theme=theme) as demo:
# Select first available model as default
default_model = available_models[0] if available_models else "FastWan2.1-T2V-1.3B"
initial_values = get_default_values(default_model)
initial_show_image = is_i2v_model(default_model)
with gr.Blocks(title="CausalWan", theme=theme) as demo:
gr.Image("assets/logos/logo.svg", show_label=False, container=False, height=80)
gr.HTML("""
<div style="text-align: center; margin-bottom: 10px;">
<p style="font-size: 18px;"> Make Video Generation Go Blurrrrrrr </p>
<p style="font-size: 18px;"> <a href="https://github.com/hao-ai-lab/FastVideo/tree/main" target="_blank">Code</a> | <a href="https://hao-ai-lab.github.io/blogs/fastvideo_post_training/" target="_blank">Blog</a> | <a href="https://hao-ai-lab.github.io/FastVideo/" target="_blank">Docs</a> </p>
<p style="font-size: 18px;"> <a href="https://github.com/hao-ai-lab/FastVideo/tree/main" target="_blank">Code</a> | <a href="https://hao-ai-lab.github.io/blogs/fastvideo_causalwan_preview/" target="_blank">Blog</a> | <a href="https://hao-ai-lab.github.io/FastVideo/" target="_blank">Docs</a> </p>
</div>
""")
@@ -280,8 +398,8 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
with gr.Row():
model_selection = gr.Dropdown(
choices=list(MODEL_PATH_MAPPING.keys()),
value="FastWan2.1-T2V-1.3B",
choices=available_models,
value=default_model,
label="Select Model",
interactive=True
)
@@ -312,69 +430,70 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
with gr.Row():
with gr.Column():
error_output = gr.Text(label="Error", visible=False)
timing_display = gr.Markdown(label="Timing Breakdown", visible=False)
# timing_display = gr.Markdown(label="Timing Breakdown", visible=False)
with gr.Row(equal_height=True, elem_classes="main-content-row"):
with gr.Column(scale=1, elem_classes="advanced-options-column"):
with gr.Group():
gr.HTML("<div style='margin: 0 0 15px 0; text-align: center; font-size: 16px;'>Advanced Options</div>")
with gr.Row():
height = gr.Number(
label="Height",
value=initial_values['height'],
interactive=False,
container=True
)
width = gr.Number(
label="Width",
value=initial_values['width'],
interactive=False,
container=True
with gr.Row(equal_height=False):
with gr.Column(scale=1):
with gr.Tabs():
with gr.Tab("Input Image", visible=initial_show_image) as image_tab:
gr.Markdown("**Please make sure you upload a 480x832 image**")
input_image = gr.Image(
label="",
type="pil",
height=400,
)
with gr.Row():
num_frames = gr.Number(
label="Number of Frames",
value=initial_values['num_frames'],
interactive=False,
container=True
)
guidance_scale = gr.Slider(
label="Guidance Scale",
minimum=1,
maximum=12,
value=initial_values['guidance_scale'],
)
with gr.Row():
use_negative_prompt = gr.Checkbox(
label="Use negative prompt", value=False)
negative_prompt = gr.Text(
label="Negative prompt",
max_lines=3,
lines=3,
placeholder="Enter a negative prompt",
visible=False,
)
with gr.Tab("Advanced Options"):
with gr.Group():
with gr.Row():
height = gr.Number(
label="Height",
value=initial_values['height'],
interactive=False,
container=True
)
width = gr.Number(
label="Width",
value=initial_values['width'],
interactive=False,
container=True
)
with gr.Row():
num_frames = gr.Number(
label="Number of Frames",
value=initial_values['num_frames'],
interactive=False,
container=True
)
guidance_scale = gr.Number(
label="Guidance Scale",
value=1.0,
interactive=False,
container=True
)
with gr.Row():
use_negative_prompt = gr.Checkbox(
label="Use negative prompt", value=False)
negative_prompt = gr.Text(
label="Negative prompt",
max_lines=3,
lines=3,
placeholder="Enter a negative prompt",
visible=False,
)
seed = gr.Slider(
label="Seed",
minimum=0,
maximum=1000000,
step=1,
value=initial_values['seed'],
)
randomize_seed = gr.Checkbox(label="Randomize seed", value=False)
seed_output = gr.Number(label="Used Seed")
# randomize_seed = gr.Checkbox(label="Randomize seed", value=False)
seed_output = gr.Number(label="Used Seed", value=1000)
with gr.Column(scale=1, elem_classes="video-column"):
with gr.Column(scale=1):
result = gr.Video(
label="Generated Video",
show_label=True,
height=466,
width=600,
height=500,
container=True,
elem_classes="video-component"
autoplay=True,
)
gr.HTML("""
@@ -387,116 +506,10 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
}
.gradio-container {
max-width: 1200px !important;
max-width: 1400px !important;
margin: 0 auto !important;
}
.main {
max-width: 1200px !important;
margin: 0 auto !important;
}
.gr-form, .gr-box, .gr-group {
max-width: 1200px !important;
}
.gr-video {
max-width: 500px !important;
margin: 0 auto !important;
}
.main-content-row {
display: flex !important;
align-items: flex-start !important;
min-height: 500px !important;
gap: 20px !important;
}
.advanced-options-column,
.video-column {
display: flex !important;
flex-direction: column !important;
flex: 1 !important;
min-height: 400px !important;
align-items: stretch !important;
}
.video-column > * {
margin-top: 0 !important;
}
.video-column .gr-video,
.video-component {
margin-top: 0 !important;
padding-top: 0 !important;
}
.video-column .gr-video .gr-form {
margin-top: 0 !important;
}
.advanced-options-column .gr-group,
.video-column .gr-video {
margin-top: 0 !important;
vertical-align: top !important;
}
.advanced-options-column > *:last-child,
.video-column > *:last-child {
flex-grow: 0 !important;
}
@media (max-width: 1400px) {
.main-content-row {
min-height: 600px !important;
}
.advanced-options-column,
.video-column {
min-height: 600px !important;
}
}
@media (max-width: 1200px) {
.main-content-row {
flex-direction: column !important;
align-items: stretch !important;
}
.advanced-options-column,
.video-column {
min-height: auto !important;
width: 100% !important;
}
}
.timing-card {
background: var(--background-fill-secondary) !important;
border: 1px solid var(--border-color-primary) !important;
color: var(--body-text-color) !important;
padding: 10px;
border-radius: 8px;
text-align: center;
min-height: 80px;
display: flex;
flex-direction: column;
justify-content: center;
}
.timing-card-highlight {
background: var(--background-fill-primary) !important;
border: 2px solid var(--color-accent) !important;
}
.performance-card {
background: var(--background-fill-secondary) !important;
border: 1px solid var(--border-color-primary) !important;
color: var(--body-text-color) !important;
padding: 10px;
border-radius: 6px;
text-align: center;
}
.gr-number input[readonly] {
background-color: var(--background-fill-secondary) !important;
border: 1px solid var(--border-color-primary) !important;
@@ -511,18 +524,30 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
def on_example_select(example_label):
if example_label and example_label in example_labels:
index = example_labels.index(example_label)
return examples[index]
return ""
selected_prompt = examples[index]
selected_image_path = example_images[index] if index < len(example_images) else None
if selected_image_path and os.path.exists(selected_image_path):
try:
with Image.open(selected_image_path) as img:
selected_image = img.convert("RGB")
except Exception:
selected_image = None
else:
selected_image = None
return selected_prompt, selected_image
return "", None
example_dropdown.change(
fn=on_example_select,
inputs=example_dropdown,
outputs=prompt,
outputs=[prompt, input_image],
)
gr.HTML("""
<div style="text-align: center; margin-top: 10px; margin-bottom: 15px;">
<p style="font-size: 16px; margin: 0;">The compute for this demo is generously provided by <a href="https://www.gmicloud.ai/" target="_blank">GMI Cloud</a>. Note that this demo is meant to showcase FastWan's quality and that under a large number of requests, generation speed may be affected. We are also rate-limiting users to 3 requests per minute.</p>
<p style="font-size: 16px; margin: 0;">The compute for this demo is generously provided by <a href="https://www.gmicloud.ai/" target="_blank">GMI Cloud</a>. Note that this demo is meant as a preview of our distilled I2V model. Outside of few-step distillation, we have not yet fully optimized it for speed. Stay tuned for updates!</p>
</div>
""")
@@ -537,6 +562,7 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
selected_model = "FastWan2.1-T2V-1.3B"
model_path = MODEL_PATH_MAPPING.get(selected_model)
show_image_input = is_i2v_model(selected_model)
if model_path and model_path in default_params:
params = default_params[model_path]
@@ -545,29 +571,29 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
gr.update(value=params.width),
gr.update(value=params.num_frames),
gr.update(value=params.guidance_scale),
gr.update(value=params.seed),
gr.update(visible=show_image_input),
)
return (
gr.update(value=448),
gr.update(value=832),
gr.update(value=61),
gr.update(value=20),
gr.update(value=3.0),
gr.update(value=1024),
gr.update(visible=show_image_input),
)
model_selection.change(
fn=on_model_selection_change,
inputs=model_selection,
outputs=[height, width, num_frames, guidance_scale, seed],
outputs=[height, width, num_frames, guidance_scale, image_tab],
)
def handle_generation(*args, progress=None, request: gr.Request = None):
model_selection, prompt, negative_prompt, use_negative_prompt, seed, guidance_scale, num_frames, height, width, randomize_seed = args
model_selection, prompt, negative_prompt, use_negative_prompt, guidance_scale, num_frames, height, width, input_image = args
result_path, seed_or_error, timing_details = generate_video(
prompt, negative_prompt, use_negative_prompt, seed, guidance_scale,
num_frames, height, width, randomize_seed, model_selection, progress
result_path, seed_or_error, _ = generate_video(
prompt, negative_prompt, use_negative_prompt, guidance_scale,
num_frames, height, width, model_selection, input_image, progress
)
if result_path and os.path.exists(result_path):
@@ -575,14 +601,12 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
result_path,
seed_or_error,
gr.update(visible=False),
gr.update(visible=True, value=timing_details),
)
else:
return (
None,
seed_or_error,
gr.update(visible=True, value=seed_or_error),
gr.update(visible=False),
)
run_button.click(
@@ -592,14 +616,14 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
prompt,
negative_prompt,
use_negative_prompt,
seed,
guidance_scale,
num_frames,
height,
width,
randomize_seed,
# randomize_seed,
input_image,
],
outputs=[result, seed_output, error_output, timing_display],
outputs=[result, seed_output, error_output], # timing_display removed
concurrency_limit=20,
)
@@ -611,8 +635,11 @@ def main():
parser.add_argument("--backend_url", type=str, default="http://localhost:8000",
help="URL of the Ray Serve backend")
parser.add_argument("--t2v_model_paths", type=str,
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.1-T2V-14B-Diffusers",
default="",
help="Comma separated list of paths to the T2V model(s)")
parser.add_argument("--i2v_model_paths", type=str,
default="FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
help="Comma separated list of paths to the I2V model(s)")
parser.add_argument("--host", type=str, default="0.0.0.0",
help="Host to bind to")
parser.add_argument("--port", type=int, default=7860,
@@ -621,8 +648,15 @@ def main():
args = parser.parse_args()
default_params = {}
model_paths = args.t2v_model_paths.split(",")
for model_path in model_paths:
# Load T2V model params
t2v_paths = [p.strip() for p in args.t2v_model_paths.split(",") if p.strip()]
for model_path in t2v_paths:
default_params[model_path] = SamplingParam.from_pretrained(model_path)
# Load I2V model params
i2v_paths = [p.strip() for p in args.i2v_model_paths.split(",") if p.strip()]
for model_path in i2v_paths:
default_params[model_path] = SamplingParam.from_pretrained(model_path)
demo = create_gradio_interface(args.backend_url, default_params)
@@ -630,6 +664,8 @@ def main():
print(f"Starting Gradio frontend at http://{args.host}:{args.port}")
print(f"Backend URL: {args.backend_url}")
print(f"T2V Models: {args.t2v_model_paths}")
if args.i2v_model_paths:
print(f"I2V Models: {args.i2v_model_paths}")
from fastapi import FastAPI, Request, HTTPException
from fastapi.responses import HTMLResponse, FileResponse
@@ -674,23 +710,23 @@ def main():
<meta charset="UTF-8" />
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
<title>FastWan</title>
<meta name="title" content="FastWan">
<title>CausalWan</title>
<meta name="title" content="CausalWan">
<meta name="description" content="Make video generation go blurrrrrrr">
<meta name="keywords" content="FastVideo, video generation, AI, machine learning, FastWan">
<meta name="keywords" content="FastVideo, video generation, AI, machine learning, CausalWan">
<meta property="og:type" content="website">
<meta property="og:url" content="{base_url}/">
<meta property="og:title" content="FastWan">
<meta property="og:title" content="CausalWan">
<meta property="og:description" content="Make video generation go blurrrrrrr">
<meta property="og:image" content="{base_url}/logo.svg">
<meta property="og:image:width" content="1200">
<meta property="og:image:height" content="630">
<meta property="og:site_name" content="FastWan">
<meta property="og:site_name" content="CausalWan">
<meta property="twitter:card" content="summary_large_image">
<meta property="twitter:url" content="{base_url}/">
<meta property="twitter:title" content="FastWan">
<meta property="twitter:title" content="CausalWan">
<meta property="twitter:description" content="Make video generation go blurrrrrrr">
<meta property="twitter:image" content="{base_url}/logo.svg">
<link rel="icon" type="image/png" sizes="32x32" href="/favicon.ico">
@@ -720,7 +756,15 @@ def main():
app,
demo,
path="/gradio",
allowed_paths=[os.path.abspath("outputs"), os.path.abspath("fastvideo-logos")]
root_path="/gradio",
allowed_paths=[
os.path.abspath("outputs"),
os.path.abspath("fastvideo-logos"),
os.path.abspath("prompts"),
os.path.abspath("images"),
os.path.abspath(tempfile.gettempdir()),
os.path.abspath(os.path.join(tempfile.gettempdir(), "gradio")),
]
)
uvicorn.run(app, host=args.host, port=args.port)
@@ -0,0 +1,15 @@
Some friends dancing and having fun together in circles, at a party surrounded by colored lights at a party, in a fancy old place, in a view from below them.
A man wearing grey shorts jumps rope in a gym, weights and gym equipment in the background.
Flying over a peninsula covered in bushy trees, while discovering the sea around it, painted a beautiful turquoise blue, on a sunny day.
Skillful cyclist doing a wheelie on a bike while riding through a forest, on a dirt road, surrounded by many trees, in the morning.
In the video, a lone rider guides a majestic horse across an expansive, open field as the sun sets in the background. The rider, dressed in a classic blue shirt and wide-brimmed hat, sits confidently in the saddle, silhouetted against the warm glow of the evening sky. The horse moves gracefully, its mane and tail flowing with each step, creating a sense of harmony between horse and rider. Surrounding the pair, towering trees form a natural border, their leaves gently rustling in the breeze. The shadows lengthen on the ground, accentuating the serene and timeless feel of the scene. The distant hills and wooden fences frame the horizon, adding depth to the tranquil landscape. A few horses graze peacefully in the background, blending into the pastoral setting. The overall ambiance evokes a sense of calmness and quietude, capturing a perfect moment in the golden light of dusk.
A silver SUV drives along a winding, snow-covered mountain road, with dense pine trees blanketed in snow lining both sides. The scene is serene, with the vehicle moving smoothly, possibly on a winter journey or vacation. As the SUV disappears around the bend, another, darker SUV follows, creating a sense of motion and perspective on the snow-dusted asphalt. The towering, snow-laden rock formation to the right contrasts with the dark green of the pines, highlighting the peacefulness of the wintry landscape.
A man and a woman playing in a field with grass, during a bright afternoon, while cars pass by in the distance.
A saxophonist wearing a blazer dances while playing a song in a park.
Romantic couple embracing and looking at each other in the middle of a forest, during a break on a road trip through nature.
Young woman cleaning her house decorated with plants and decorations, while dancing happily to music in her headphones.
Man dressed in 80's style dances very happily in his kitchen while listening to music on his radio and drinking wine.
Pair of jazz musicians performing a song with their saxophone and trombone on an abandoned train.
A young woman with short hair wearing pink sunglasses chews gum and makes a bubble gum with the city in the background.
Natural aerial landscape with a relief covered with abundant trees and vegetation and a thick layer of mist.
Loving couple sitting on a log on the shore of a lake outside, sharing an affectionate hug.
@@ -26,6 +26,7 @@ SEED_RANGE_MAX = 1_000_000
SUPPORTED_MODELS = [
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
"FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
"FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
]
MODEL_CONFIGS = {
@@ -42,6 +43,13 @@ MODEL_CONFIGS = {
"dit_cpu_offload": True,
"vae_cpu_offload": False,
"VSA_sparsity": 0.9,
},
"I2V-A14B": {
"num_cpus": 15,
"text_encoder_cpu_offload": True,
"dit_cpu_offload": True,
"vae_cpu_offload": False,
"VSA_sparsity": 0.0,
}
}
@@ -58,6 +66,7 @@ class VideoGenerationRequest(BaseModel):
randomize_seed: bool = False
return_frames: bool = False
model_path: Optional[str] = None
image_data: Optional[str] = None # Base64 encoded image for I2V
class VideoGenerationResponse(BaseModel):
@@ -91,11 +100,38 @@ def encode_video_to_base64(frames: List[np.ndarray], fps: int = DEFAULT_FPS) ->
return ""
def save_image_from_base64(image_data: str, output_dir: str) -> Optional[str]:
"""Save base64 image data to a temporary file and return the path."""
if not image_data:
return None
try:
# Remove data URL prefix if present
if image_data.startswith('data:image/'):
image_data = image_data.split(',')[1]
image_bytes = base64.b64decode(image_data)
# Save to temporary file
os.makedirs(output_dir, exist_ok=True)
temp_image_path = os.path.join(output_dir, f"temp_input_{int(time.time() * 1000)}.png")
with open(temp_image_path, 'wb') as f:
f.write(image_bytes)
return temp_image_path
except Exception as e:
print(f"Warning: Failed to save image: {e}")
return None
def setup_model_environment(model_path: str) -> None:
if "fullattn" in model_path.lower():
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
else:
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
# if "fullattn" in model_path.lower():
# os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
# else:
# os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
os.environ["FASTVIDEO_STAGE_LOGGING"] = "1"
@@ -157,22 +193,41 @@ class BaseModelDeployment:
num_gpus=1,
use_fsdp_inference=True,
text_encoder_cpu_offload=config["text_encoder_cpu_offload"],
dmd_denoising_steps=[1000, 850, 700, 550, 350, 275, 200, 125], # TODO: hardocde for I2V
dit_precision="fp32", # TODO: hardocde for I2V
dit_cpu_offload=config["dit_cpu_offload"],
vae_cpu_offload=config["vae_cpu_offload"],
VSA_sparsity=config["VSA_sparsity"],
enable_stage_verification=False,
)
self.default_params = SamplingParam.from_pretrained(self.model_path)
self.default_params.seed = 1000
self.default_params.num_frames = 73
self.default_params.width = 832
self.default_params.height = 480
def generate_video(self, video_request: VideoGenerationRequest) -> VideoGenerationResponse:
total_start_time = time.time()
params = prepare_sampling_params(video_request, self.default_params)
# Save image if provided (for I2V)
image_path = None
if video_request.image_data:
image_path = save_image_from_base64(video_request.image_data, self.output_path)
if image_path is None:
return VideoGenerationResponse(
video_data=None,
seed=params.seed,
success=False,
error_message="Failed to save input image",
)
inference_start_time = time.time()
result = self.generator.generate_video(
prompt=video_request.prompt,
sampling_param=params,
image_path=image_path,
save_video=False,
return_frames=False,
)
@@ -185,6 +240,13 @@ class BaseModelDeployment:
encoding_time = time.time() - encoding_start_time
total_time = time.time() - total_start_time
# Clean up temporary image file
if image_path and os.path.exists(image_path):
try:
os.remove(image_path)
except Exception as e:
print(f"Warning: Failed to remove temporary image file {image_path}: {e}")
return VideoGenerationResponse(
video_data=video_data,
@@ -200,7 +262,7 @@ class BaseModelDeployment:
@serve.deployment(
ray_actor_options={"num_cpus": 2, "num_gpus": 1, "runtime_env": {"conda": "fv"}},
ray_actor_options={"num_cpus": 2, "num_gpus": 1, "runtime_env": {"conda": "demo-fv"}},
)
class T2VModelDeployment(BaseModelDeployment):
def __init__(self, t2v_model_path: str, output_path: str = "outputs"):
@@ -210,7 +272,7 @@ class T2VModelDeployment(BaseModelDeployment):
@serve.deployment(
ray_actor_options={"num_cpus": 16, "num_gpus": 1, "runtime_env": {"conda": "fv"}},
ray_actor_options={"num_cpus": 16, "num_gpus": 1, "runtime_env": {"conda": "demo-fv"}},
)
class T2V14BModelDeployment(BaseModelDeployment):
def __init__(self, t2v_14b_model_path: str, output_path: str = "outputs"):
@@ -221,18 +283,32 @@ class T2V14BModelDeployment(BaseModelDeployment):
print("✅ T2V 14B model initialized successfully")
@serve.deployment(
ray_actor_options={"num_cpus": 15, "num_gpus": 1, "runtime_env": {"conda": "demo-fv"}},
)
class I2VModelDeployment(BaseModelDeployment):
def __init__(self, i2v_model_path: str, output_path: str = "outputs"):
super().__init__(i2v_model_path, output_path)
# Override environment for I2V model
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
self._initialize_generator(MODEL_CONFIGS["I2V-A14B"])
print("✅ I2V model initialized successfully")
app = FastAPI()
limiter = Limiter(key_func=get_remote_address)
app.state.limiter = limiter
app.add_exception_handler(RateLimitExceeded, _rate_limit_exceeded_handler)
@serve.deployment(num_replicas=50, ray_actor_options={"num_cpus": 2})
@serve.deployment(num_replicas=1, ray_actor_options={"num_cpus": 1})
@serve.ingress(app)
class FastVideoAPI:
def __init__(self, t2v_deployments: Dict[str, DeploymentHandle]):
def __init__(self, t2v_deployments: Dict[str, DeploymentHandle], i2v_deployments: Dict[str, DeploymentHandle] = None):
self.t2v_deployments = t2v_deployments
self.i2v_deployments = i2v_deployments or {}
self.all_deployments = {**self.t2v_deployments, **self.i2v_deployments}
# Initialize Prometheus metrics
self.request_count = Counter('fastvideo_requests_total', 'Total FastVideo requests', ['model_type', 'status'])
@@ -257,10 +333,10 @@ class FastVideoAPI:
model_name = self._get_model_name(video_request.model_path)
try:
if video_request.model_path not in self.t2v_deployments:
if video_request.model_path not in self.all_deployments:
raise ValueError(f"Model {video_request.model_path} not found")
response_ref = self.t2v_deployments[video_request.model_path].generate_video.remote(video_request)
response_ref = self.all_deployments[video_request.model_path].generate_video.remote(video_request)
response = await response_ref
self._record_metrics(model_name, "success", time.time() - start_time, response)
@@ -291,18 +367,21 @@ class FastVideoAPI:
def validate_configuration(model_paths: List[str], replicas: List[int]) -> None:
assert len(model_paths) > 0, "At least one model must be specified"
assert len(model_paths) == len(replicas), "Number of models and replicas must match"
assert sum(replicas) <= NUM_GPUS, f"Total replicas ({sum(replicas)}) must be <= {NUM_GPUS}"
for model, replica_count in zip(model_paths, replicas):
assert model in SUPPORTED_MODELS, f"Model {model} not supported"
assert model in SUPPORTED_MODELS, f"Model {model} not supported. Supported models: {SUPPORTED_MODELS}"
assert replica_count > 0, f"Replicas must be greater than 0"
def start_ray_serve(
*,
t2v_model_paths: str,
t2v_model_replicas: str,
t2v_model_paths: str = "",
t2v_model_replicas: str = "",
i2v_model_paths: str = "",
i2v_model_replicas: str = "",
output_path: str = "outputs",
host: str = "0.0.0.0",
port: int = 8000,
@@ -310,21 +389,39 @@ def start_ray_serve(
if not ray.is_initialized():
ray.init()
model_paths = t2v_model_paths.split(",")
replicas = [int(r) for r in t2v_model_replicas.split(",")]
validate_configuration(model_paths, replicas)
# Parse T2V models
t2v_paths = [p.strip() for p in t2v_model_paths.split(",") if p.strip()]
t2v_reps = [int(r.strip()) for r in t2v_model_replicas.split(",") if r.strip()] if t2v_model_replicas else []
# Parse I2V models
i2v_paths = [p.strip() for p in i2v_model_paths.split(",") if p.strip()]
i2v_reps = [int(r.strip()) for r in i2v_model_replicas.split(",") if r.strip()] if i2v_model_replicas else []
# Validate configurations
all_paths = t2v_paths + i2v_paths
all_replicas = t2v_reps + i2v_reps
validate_configuration(all_paths, all_replicas)
# Create T2V deployments
t2v_deps = {}
for model_path, replica_count in zip(model_paths, replicas):
for model_path, replica_count in zip(t2v_paths, t2v_reps):
t2v_dep = T2VModelDeployment.options(num_replicas=replica_count).bind(model_path, output_path)
t2v_deps[model_path] = t2v_dep
api = FastVideoAPI.bind(t2v_deps)
# Create I2V deployments
i2v_deps = {}
for model_path, replica_count in zip(i2v_paths, i2v_reps):
i2v_dep = I2VModelDeployment.options(num_replicas=replica_count).bind(model_path, output_path)
i2v_deps[model_path] = i2v_dep
api = FastVideoAPI.bind(t2v_deps, i2v_deps)
serve.run(api, route_prefix="/", name="fast_video")
print(f"Ray Serve backend started at http://{host}:{port}")
for model_path, replica_count in zip(model_paths, replicas):
for model_path, replica_count in zip(t2v_paths, t2v_reps):
print(f"T2V Model: {model_path} | Replicas: {replica_count}")
for model_path, replica_count in zip(i2v_paths, i2v_reps):
print(f"I2V Model: {model_path} | Replicas: {replica_count}")
print(f"Health check: http://{host}:{port}/health")
print(f"Video generation endpoint: http://{host}:{port}/generate_video")
@@ -340,12 +437,20 @@ if __name__ == "__main__":
parser = argparse.ArgumentParser(description="FastVideo Ray Serve Backend")
parser.add_argument("--t2v_model_paths",
type=str,
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
default="",
help="Comma separated list of paths to the T2V model(s)")
parser.add_argument("--t2v_model_replicas",
type=str,
default="4,4",
default="",
help="Comma separated list of number of replicas for the T2V model(s)")
parser.add_argument("--i2v_model_paths",
type=str,
default="FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
help="Comma separated list of paths to the I2V model(s)")
parser.add_argument("--i2v_model_replicas",
type=str,
default="1",
help="Comma separated list of number of replicas for the I2V model(s)")
parser.add_argument("--output_path",
type=str,
default="outputs",
@@ -361,13 +466,21 @@ if __name__ == "__main__":
args = parser.parse_args()
model_paths = args.t2v_model_paths.split(",")
replicas = [int(r) for r in args.t2v_model_replicas.split(",")]
validate_configuration(model_paths, replicas)
# Parse and validate all models
t2v_paths = [p.strip() for p in args.t2v_model_paths.split(",") if p.strip()]
t2v_reps = [int(r.strip()) for r in args.t2v_model_replicas.split(",") if r.strip()] if args.t2v_model_replicas else []
i2v_paths = [p.strip() for p in args.i2v_model_paths.split(",") if p.strip()]
i2v_reps = [int(r.strip()) for r in args.i2v_model_replicas.split(",") if r.strip()] if args.i2v_model_replicas else []
all_paths = t2v_paths + i2v_paths
all_replicas = t2v_reps + i2v_reps
validate_configuration(all_paths, all_replicas)
start_ray_serve(
t2v_model_paths=args.t2v_model_paths,
t2v_model_replicas=args.t2v_model_replicas,
i2v_model_paths=args.i2v_model_paths,
i2v_model_replicas=args.i2v_model_replicas,
output_path=args.output_path,
host=args.host,
port=args.port,
@@ -376,4 +489,4 @@ if __name__ == "__main__":
setup_signal_handlers()
print("✅ FastVideo backend is running. Press Ctrl-C to stop.")
while True:
time.sleep(3600)
time.sleep(3600)
+5 -3
View File
@@ -1,3 +1,5 @@
python examples/inference/gradio/start_ray_serve_app.py \
--t2v_model_paths "FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers" \
--t2v_model_replicas "4,4"
python examples/inference/gradio/serving/start_ray_serve_app.py \
--t2v_model_paths "" \
--t2v_model_replicas "" \
--i2v_model_paths "FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers" \
--i2v_model_replicas "1"
@@ -20,8 +20,10 @@ DEFAULT_BACKEND_PORT = 8000
DEFAULT_FRONTEND_HOST = "0.0.0.0"
DEFAULT_FRONTEND_PORT = 7860
DEFAULT_OUTPUT_PATH = "outputs"
DEFAULT_T2V_MODELS = "FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers"
DEFAULT_T2V_REPLICAS = "4,4"
DEFAULT_T2V_MODELS = ""
DEFAULT_T2V_REPLICAS = ""
DEFAULT_I2V_MODELS = "FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers"
DEFAULT_I2V_REPLICAS = "1"
HEALTH_CHECK_TIMEOUT = 5
HEALTH_CHECK_MAX_RETRIES = 100
@@ -100,6 +102,12 @@ class ServiceManager:
"port": self.args.backend_port
}
# Add I2V parameters if provided
if self.args.i2v_model_paths:
backend_args["i2v_model_paths"] = self.args.i2v_model_paths
if self.args.i2v_model_replicas:
backend_args["i2v_model_replicas"] = self.args.i2v_model_replicas
self.backend_process = self._start_service("ray_serve_backend.py", backend_args, "backend")
return self.backend_process
@@ -111,6 +119,10 @@ class ServiceManager:
"port": self.args.frontend_port
}
# Add I2V parameters if provided
if self.args.i2v_model_paths:
frontend_args["i2v_model_paths"] = self.args.i2v_model_paths
self.frontend_process = self._start_service("gradio_frontend.py", frontend_args, "frontend")
return self.frontend_process
@@ -173,6 +185,9 @@ def print_startup_info(args: argparse.Namespace) -> None:
print("=" * 50)
print(f"T2V Models: {args.t2v_model_paths}")
print(f"T2V Model Replicas: {args.t2v_model_replicas}")
if args.i2v_model_paths:
print(f"I2V Models: {args.i2v_model_paths}")
print(f"I2V Model Replicas: {args.i2v_model_replicas}")
print(f"Output: {args.output_path}")
print(f"Backend: http://{args.backend_host}:{args.backend_port}")
print(f"Frontend: http://{args.frontend_host}:{args.frontend_port}")
@@ -190,6 +205,14 @@ def parse_arguments() -> argparse.Namespace:
type=str,
default=DEFAULT_T2V_REPLICAS,
help="Comma separated list of number of replicas for the T2V model(s)")
parser.add_argument("--i2v_model_paths",
type=str,
default=DEFAULT_I2V_MODELS,
help="Comma separated list of paths to the I2V model(s)")
parser.add_argument("--i2v_model_replicas",
type=str,
default=DEFAULT_I2V_REPLICAS,
help="Comma separated list of number of replicas for the I2V model(s)")
parser.add_argument("--output_path",
type=str,
default=DEFAULT_OUTPUT_PATH,
+2
View File
@@ -39,6 +39,8 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
"FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers": WanT2V720PConfig,
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers": SelfForcingWanT2V480PConfig,
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers": SelfForcingWan2_2_T2V480PConfig,
"FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers":
SelfForcingWan2_2_T2V480PConfig,
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_Config,
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_Config,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_Config,
+4
View File
@@ -186,3 +186,7 @@ class SelfForcingWan2_2_T2V480PConfig(Wan2_2_T2V_A14B_Config):
dmd_denoising_steps: list[int] | None = field(
default_factory=lambda: [1000, 850, 700, 550, 350, 275, 200, 125])
warp_denoising_step: bool = True
def __post_init__(self) -> None:
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
+2
View File
@@ -78,6 +78,8 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
# Causal Self-Forcing Wan2.2
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers":
SelfForcingWan2_2_T2V_A14B_480P_SamplingParam,
"FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers":
SelfForcingWan2_2_T2V_A14B_480P_SamplingParam,
# Cosmos2
"nvidia/Cosmos-Predict2-2B-Video2World":
-2
View File
@@ -191,8 +191,6 @@ class SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam(
@dataclass
class SelfForcingWan2_2_T2V_A14B_480P_SamplingParam(
Wan2_2_T2V_A14B_SamplingParam):
guidance_scale: float = 2.0
guidance_scale_2: float = 2.0
num_inference_steps: int = 8
num_frames: int = 81
height: int = 448
@@ -86,6 +86,12 @@ class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
return (prev_sample, )
return SelfForcingFlowMatchSchedulerOutput(prev_sample=prev_sample)
@staticmethod
def calculate_alpha_beta_high(sigma, sigma_bound):
alpha = (1 - sigma) / (1 - sigma_bound)
beta = torch.sqrt(sigma ** 2 - (alpha * sigma_bound) ** 2)
return alpha, beta
def add_noise(self, original_samples, noise, timestep):
"""
Diffusion forward corruption process.
@@ -105,6 +111,32 @@ class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
sample = (1 - sigma) * original_samples + sigma * noise
return sample.type_as(noise)
def add_noise_high(self, original_samples, noise, timestep, boundary_timestep):
"""
Diffusion forward corruption process.
Input:
- clean_latent: the clean latent with shape [B*T, C, H, W]
- noise: the noise with shape [B*T, C, H, W]
- timestep: the timestep with shape [B*T]
Output: the corrupted latent with shape [B*T, C, H, W]
"""
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
if boundary_timestep.ndim == 2:
boundary_timestep = boundary_timestep.flatten(0, 1)
self.sigmas = self.sigmas.to(noise.device)
self.timesteps = self.timesteps.to(noise.device)
timestep_id = torch.argmin(
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
boundary_timestep_id = torch.argmin(
(self.timesteps.unsqueeze(0) - boundary_timestep.unsqueeze(1)).abs(), dim=1)
sigma_boundary = self.sigmas[boundary_timestep_id].reshape(-1, 1, 1, 1)
alpha, beta = self.calculate_alpha_beta_high(sigma, sigma_boundary)
sample = alpha * original_samples + beta * noise
return sample.type_as(noise)
def training_target(self, sample, noise, timestep):
target = noise - sample
return target
+48
View File
@@ -180,3 +180,51 @@ def pred_noise_to_pred_video(pred_noise: torch.Tensor,
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
pred_video = noise_input_latent - sigma_t * pred_noise
return pred_video.to(dtype)
def pred_noise_to_x_bound(pred_noise: torch.Tensor,
noise_input_latent: torch.Tensor,
timestep: torch.Tensor,
boundary_timestep: torch.Tensor,
scheduler: Any) -> torch.Tensor:
"""
Convert predicted noise to clean latent.
Args:
pred_noise: the predicted noise with shape [B, C, H, W]
where B is batch_size or batch_size * num_frames
noise_input_latent: the noisy latent with shape [B, C, H, W],
timestep: the timestep with shape [1] or [bs * num_frames] or [bs, num_frames]
boundary_timestep: the boundary timestep with shape [1] or [bs * num_frames] or [bs, num_frames]
scheduler: the scheduler
Returns:
the predicted video with shape [B, C, H, W]
"""
# If timestep is [bs, num_frames]
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
assert timestep.numel() == noise_input_latent.shape[0]
elif timestep.ndim == 1:
# If timestep is [1]
if timestep.shape[0] == 1:
timestep = timestep.expand(noise_input_latent.shape[0])
else:
assert timestep.numel() == noise_input_latent.shape[0]
else:
raise ValueError(f"[pred_noise_to_pred_video] Invalid timestep shape: {timestep.shape}")
# timestep shape should be [B]
dtype = pred_noise.dtype
device = pred_noise.device
pred_noise = pred_noise.double().to(device)
noise_input_latent = noise_input_latent.double().to(device)
sigmas = scheduler.sigmas.double().to(device)
timesteps = scheduler.timesteps.double().to(device)
timestep_id = torch.argmin(
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
boundary_timestep_id = torch.argmin(
(timesteps.unsqueeze(0) - boundary_timestep.unsqueeze(1)).abs(), dim=1)
sigma_t_boundary = sigmas[boundary_timestep_id].reshape(-1, 1, 1, 1)
pred_video = noise_input_latent - (sigma_t - sigma_t_boundary) * pred_noise
return pred_video.to(dtype)
@@ -50,7 +50,8 @@ class WanCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
stage=CausalDMDDenosingStage(
transformer=self.get_module("transformer"),
transformer_2=self.get_module("transformer_2", None),
scheduler=self.get_module("scheduler")))
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
+163 -126
View File
@@ -4,7 +4,7 @@ from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.utils import pred_noise_to_pred_video
from fastvideo.models.utils import pred_noise_to_pred_video, pred_noise_to_x_bound
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.denoising import DenoisingStage
from fastvideo.pipelines.stages.validators import StageValidators as V
@@ -34,13 +34,16 @@ class CausalDMDDenosingStage(DenoisingStage):
Denoising stage for causal diffusion.
"""
def __init__(self, transformer, scheduler, transformer_2=None) -> None:
def __init__(self,
transformer,
scheduler,
transformer_2=None,
vae=None) -> None:
super().__init__(transformer, scheduler, transformer_2)
# KV and cross-attention cache state (initialized on first forward)
self.transformer = transformer
self.transformer_2 = transformer_2
self.kv_cache1: list | None = None
self.crossattn_cache: list | None = None
self.vae = vae
# Model-dependent constants (aligned with causal_inference.py assumptions)
self.num_transformer_blocks = len(self.transformer.blocks)
self.num_frames_per_block = self.transformer.config.arch_config.num_frames_per_block
@@ -80,6 +83,13 @@ class CausalDMDDenosingStage(DenoisingStage):
timesteps = scheduler_timesteps[1000 - timesteps]
timesteps = timesteps.to(get_local_torch_device())
if fastvideo_args.pipeline_config.dit_config.boundary_ratio is not None:
boundary_timestep = fastvideo_args.pipeline_config.dit_config.boundary_ratio * self.scheduler.num_train_timesteps
else:
boundary_timestep = None
high_noise_timesteps = timesteps[timesteps >= boundary_timestep]
# Image kwargs (kept empty unless caller provides compatible args)
image_kwargs: dict = {}
@@ -103,113 +113,102 @@ class CausalDMDDenosingStage(DenoisingStage):
assert torch.isnan(prompt_embeds[0]).sum() == 0
# Initialize or reset caches
if self.kv_cache1 is None:
self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
self._initialize_crossattn_cache(
batch_size=latents.shape[0],
max_text_len=fastvideo_args.pipeline_config.
text_encoder_configs[0].arch_config.text_len,
dtype=target_dtype,
device=latents.device)
else:
assert self.crossattn_cache is not None
# reset cross-attention cache
for block_index in range(self.num_transformer_blocks):
self.crossattn_cache[block_index][
"is_init"] = False # type: ignore
# reset kv cache pointers
for block_index in range(len(self.kv_cache1)):
self.kv_cache1[block_index][
"global_end_index"] = torch.tensor( # type: ignore
[0],
dtype=torch.long,
device=latents.device)
self.kv_cache1[block_index][
"local_end_index"] = torch.tensor( # type: ignore
[0],
dtype=torch.long,
device=latents.device)
kv_cache1 = self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
kv_cache2 = None
if boundary_timestep is not None:
# Initialize the low noise kv cache
kv_cache2 = self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
# Optional: cache context features from provided image latents prior to generation
current_start_frame = 0
if getattr(batch, "image_latent", None) is not None:
image_latent = batch.image_latent
assert image_latent is not None
input_frames = image_latent.shape[2]
# timestep zero (or configured context noise) for cache warm-up
t_zero = torch.zeros([latents.shape[0]],
device=latents.device,
dtype=torch.long)
if independent_first_frame and input_frames >= 1:
# warm-up with the very first frame independently
image_first_btchw = image_latent[:, :, :1, :, :].to(
target_dtype).permute(0, 2, 1, 3, 4)
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled):
_ = self.transformer(
image_first_btchw,
prompt_embeds,
t_zero,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_start=current_start_frame *
self.frame_seq_length,
**image_kwargs,
**pos_cond_kwargs,
)
current_start_frame += 1
remaining_frames = input_frames - 1
else:
remaining_frames = input_frames
def _get_kv_cache(timestep: float) -> list[dict]:
if boundary_timestep is not None:
if timestep >= boundary_timestep:
return kv_cache1
else:
assert kv_cache2 is not None, "kv_cache2 is not initialized"
return kv_cache2
return kv_cache1
# process remaining input frames in blocks of num_frame_per_block
while remaining_frames > 0:
block = min(self.num_frames_per_block, remaining_frames)
ref_btchw = image_latent[:, :, current_start_frame:
current_start_frame +
block, :, :].to(target_dtype).permute(
0, 2, 1, 3, 4)
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled):
_ = self.transformer(
ref_btchw,
prompt_embeds,
t_zero,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_start=current_start_frame *
self.frame_seq_length,
**image_kwargs,
**pos_cond_kwargs,
)
current_start_frame += block
remaining_frames -= block
crossattn_cache = self._initialize_crossattn_cache(
batch_size=latents.shape[0],
max_text_len=fastvideo_args.pipeline_config.text_encoder_configs[0].
arch_config.text_len,
dtype=target_dtype,
device=latents.device)
# Base position offset from any cache warm-up
pos_start_base = current_start_frame
pos_start_base = 0
# Determine block sizes
if not independent_first_frame or (independent_first_frame
and batch.image_latent is not None):
if t % self.num_frames_per_block != 0:
raise ValueError(
"num_frames must be divisible by num_frames_per_block for causal DMD denoising"
block_sizes = [self.num_frames_per_block] * 7
block_sizes[0] = 1
start_index = 0
first_frame_latent = None
if batch.pil_image is not None:
# Causal video gen directly replaces the first frame of the latent with
# the image latent instead of appending along the channel dim
assert self.vae is not None, "VAE is not provided for causal video gen task"
self.vae = self.vae.to(get_local_torch_device())
first_frame_latent = self.vae.encode(batch.pil_image).mean.float()
if (hasattr(self.vae, "shift_factor")
and self.vae.shift_factor is not None):
if isinstance(self.vae.shift_factor, torch.Tensor):
first_frame_latent -= self.vae.shift_factor.to(
first_frame_latent.device, first_frame_latent.dtype)
else:
first_frame_latent -= self.vae.shift_factor
if isinstance(self.vae.scaling_factor, torch.Tensor):
first_frame_latent = first_frame_latent * self.vae.scaling_factor.to(
first_frame_latent.device, first_frame_latent.dtype)
else:
first_frame_latent = first_frame_latent * self.vae.scaling_factor
if fastvideo_args.vae_cpu_offload:
self.vae = self.vae.to("cpu")
# Fill the low noise and high noise kv cache with first_frame_latent and timestep 0
t_zero = torch.zeros([latents.shape[0], 1],
device=latents.device,
dtype=torch.long)
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled), \
set_forward_context(current_timestep=0,
attn_metadata=None,
forward_batch=batch):
self.transformer(
first_frame_latent.to(target_dtype),
prompt_embeds,
t_zero,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**pos_cond_kwargs,
)
num_blocks = t // self.num_frames_per_block
block_sizes = [self.num_frames_per_block] * num_blocks
start_index = 0
else:
if (t - 1) % self.num_frames_per_block != 0:
raise ValueError(
"(num_frames - 1) must be divisible by num_frame_per_block when independent_first_frame=True"
)
num_blocks = (t - 1) // self.num_frames_per_block
block_sizes = [1] + [self.num_frames_per_block] * num_blocks
start_index = 0
if boundary_timestep is not None:
self.transformer_2(
first_frame_latent.to(target_dtype),
prompt_embeds,
t_zero,
kv_cache=kv_cache2,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**pos_cond_kwargs,
)
start_index += 1
block_sizes.pop(0)
latents[:, :, :1, :, :] = first_frame_latent
# DMD loop in causal blocks
with self.progress_bar(total=len(block_sizes) *
@@ -222,7 +221,7 @@ class CausalDMDDenosingStage(DenoisingStage):
video_raw_latent_shape = noise_latents_btchw.shape
for i, t_cur in enumerate(timesteps):
if self.transformer_2 is not None and fastvideo_args.pipeline_config.dit_config.boundary_ratio is not None and t_cur < fastvideo_args.pipeline_config.dit_config.boundary_ratio * self.scheduler.num_train_timesteps:
if boundary_timestep is not None and t_cur < boundary_timestep:
current_model = self.transformer_2
else:
current_model = self.transformer
@@ -280,8 +279,8 @@ class CausalDMDDenosingStage(DenoisingStage):
latent_model_input,
prompt_embeds,
t_expanded_noise,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
kv_cache=_get_kv_cache(t_cur),
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
@@ -290,12 +289,22 @@ class CausalDMDDenosingStage(DenoisingStage):
).permute(0, 2, 1, 3, 4)
# Convert pred noise to pred video with FM Euler scheduler utilities
pred_video_btchw = pred_noise_to_pred_video(
pred_noise=pred_noise_btchw.flatten(0, 1),
noise_input_latent=noise_latents.flatten(0, 1),
timestep=t_expand,
scheduler=self.scheduler).unflatten(
0, pred_noise_btchw.shape[:2])
if boundary_timestep is not None and t_cur >= boundary_timestep:
pred_video_btchw = pred_noise_to_x_bound(
pred_noise=pred_noise_btchw.flatten(0, 1),
noise_input_latent=noise_latents.flatten(0, 1),
timestep=t_expand,
boundary_timestep=torch.ones_like(t_expand) *
boundary_timestep,
scheduler=self.scheduler).unflatten(
0, pred_noise_btchw.shape[:2])
else:
pred_video_btchw = pred_noise_to_pred_video(
pred_noise=pred_noise_btchw.flatten(0, 1),
noise_input_latent=noise_latents.flatten(0, 1),
timestep=t_expand,
scheduler=self.scheduler).unflatten(
0, pred_noise_btchw.shape[:2])
if i < len(timesteps) - 1:
next_timestep = timesteps[i + 1] * torch.ones(
@@ -309,11 +318,23 @@ class CausalDMDDenosingStage(DenoisingStage):
batch.generator, list) else
batch.generator)).to(self.device)
noise_btchw = noise
noise_latents_btchw = self.scheduler.add_noise(
pred_video_btchw.flatten(0, 1),
noise_btchw.flatten(0, 1),
next_timestep).unflatten(0,
pred_video_btchw.shape[:2])
if boundary_timestep is not None and i < len(
high_noise_timesteps) - 1:
noise_latents_btchw = self.scheduler.add_noise_high(
pred_video_btchw.flatten(0, 1),
noise_btchw.flatten(0, 1), next_timestep,
torch.ones_like(next_timestep) *
boundary_timestep).unflatten(
0, pred_video_btchw.shape[:2])
elif boundary_timestep is not None and i == len(
high_noise_timesteps) - 1:
noise_latents_btchw = pred_video_btchw
else:
noise_latents_btchw = self.scheduler.add_noise(
pred_video_btchw.flatten(0, 1),
noise_btchw.flatten(0, 1),
next_timestep).unflatten(
0, pred_video_btchw.shape[:2])
current_latents = noise_latents_btchw.permute(
0, 2, 1, 3, 4)
else:
@@ -341,24 +362,40 @@ class CausalDMDDenosingStage(DenoisingStage):
attn_metadata=attn_metadata,
forward_batch=batch):
t_expanded_context = t_context.unsqueeze(1)
_ = current_model(
if boundary_timestep is not None:
self.transformer_2(
context_bcthw,
prompt_embeds,
t_expanded_context,
kv_cache=kv_cache2,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**pos_cond_kwargs,
)
self.transformer(
context_bcthw,
prompt_embeds,
t_expanded_context,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**pos_cond_kwargs,
)
start_index += current_num_frames
batch.latents = latents
return batch
def _initialize_kv_cache(self, batch_size, dtype, device) -> None:
def _initialize_kv_cache(self, batch_size, dtype, device) -> list[dict]:
"""
Initialize a Per-GPU KV cache aligned with the Wan model assumptions.
"""
@@ -392,10 +429,10 @@ class CausalDMDDenosingStage(DenoisingStage):
torch.tensor([0], dtype=torch.long, device=device),
})
self.kv_cache1 = kv_cache1
return kv_cache1
def _initialize_crossattn_cache(self, batch_size, max_text_len, dtype,
device) -> None:
device) -> list[dict]:
"""
Initialize a Per-GPU cross-attention cache aligned with the Wan model assumptions.
"""
@@ -421,7 +458,7 @@ class CausalDMDDenosingStage(DenoisingStage):
"is_init":
False,
})
self.crossattn_cache = crossattn_cache
return crossattn_cache
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
+22 -20
View File
@@ -5,7 +5,6 @@ Input validation stage for diffusion pipelines.
import torch
import torchvision.transforms.functional as TF
from PIL import Image
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
@@ -14,7 +13,6 @@ from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import (StageValidators,
VerificationResult)
from fastvideo.utils import best_output_size
logger = init_logger(__name__)
@@ -106,27 +104,31 @@ class InputValidationStage(PipelineStage):
batch.pil_image = image
# further processing for ti2v task
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
if (fastvideo_args.pipeline_config.ti2v_task
or fastvideo_args.pipeline_config.is_causal
) and batch.pil_image is not None:
img = batch.pil_image
ih, iw = img.height, img.width
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
vae_stride = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
dh, dw = patch_size[1] * vae_stride, patch_size[2] * vae_stride
max_area = 704 * 1280
ow, oh = best_output_size(iw, ih, dw, dh, max_area)
# ih, iw = img.height, img.width
# patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
# vae_stride = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
# dh, dw = patch_size[1] * vae_stride, patch_size[2] * vae_stride
# max_area = 720 * 1280
# ow, oh = best_output_size(iw, ih, dw, dh, max_area)
scale = max(ow / iw, oh / ih)
img = img.resize((round(iw * scale), round(ih * scale)),
Image.LANCZOS)
logger.info("resized img height: %s, img width: %s", img.height,
img.width)
# center-crop
x1 = (img.width - ow) // 2
y1 = (img.height - oh) // 2
img = img.crop((x1, y1, x1 + ow, y1 + oh))
assert img.width == ow and img.height == oh
# scale = max(ow / iw, oh / ih)
# img = img.resize((round(iw * scale), round(ih * scale)),
# Image.LANCZOS)
# logger.info("resized img height: %s, img width: %s", img.height,
# img.width)
# # center-crop
# x1 = (img.width - ow) // 2
# y1 = (img.height - oh) // 2
# img = img.crop((x1, y1, x1 + ow, y1 + oh))
# assert img.width == ow and img.height == oh
logger.info("img height: %s, img width: %s", img.height, img.width)
oh = img.height
ow = img.width
# to tensor
img = TF.to_tensor(img).sub_(0.5).div_(0.5).to(
self.device).unsqueeze(1)
Binary file not shown.

After

Width:  |  Height:  |  Size: 113 KiB

BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 229 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 168 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 103 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 148 KiB

BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 155 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 723 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 723 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 875 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 664 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 62 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 686 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 957 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 585 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 558 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 942 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 890 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 433 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 595 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 781 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 783 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 762 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 68 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 147 KiB

BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 89 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 133 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 213 KiB

+110
View File
@@ -0,0 +1,110 @@
[
{
"prompt": "Young man skating with a skateboard on the ramps with graffiti of a park with trees, on a sunny day.",
"image_path": "images/mixkit-boy-skating-with-a-skateboard-in-a-park-with-ramps-34389.png"
},
{
"prompt": "In the midst of the joyous New Year's Eve celebration, the cheerful group of friends, their spirits lifted by the festivities, decides to immortalize the moment with a vibrant snapshot",
"image_path": "images/mixkit-a-cheerful-group-of-friends-celebrate-new-years-eve-and-51525.png"
},
{
"prompt": "A man and a woman playing in a field with grass, during a bright afternoon, while cars pass by in the distance.",
"image_path": "images/mixkit-a-cute-couple-playing-on-the-grass-4688.png"
},
{
"prompt": "Aerial view of a rocky mountain in the forest at a sunny day drone flight footage",
"image_path": "images/mixkit-aerial-view-of-a-rocky-mountain-in-the-forest-50589.png"
},
{
"prompt": "A little girl wearing a pink security helmet and denim overall discovers the art of cycling amidst the serene park, as the camera captures her graceful progress.",
"image_path": "images/mixkit-a-little-girl-cruises-through-the-forest-path-on-her-50088.png"
},
{
"prompt": "Aerial shot of a beach shore with sea waves. Big rocks on the sand at an alone beach.",
"image_path": "images/mixkit-aerial-shot-of-a-beach-with-sea-waves-1087.png"
},
{
"prompt": "Young woman cleaning her house decorated with plants and decorations, while dancing happily to music in her headphones.",
"image_path": "images/mixkit-woman-cleaning-her-house-dancing-happy-43379.png"
},
{
"prompt": "Aerial tour in a meadow surrounded by hills on the horizon, while some birds fly low over a lake.",
"image_path": "images/mixkit-birds-flying-low-over-a-lake-in-a-meadow-41417.png"
},
{
"prompt": "In the video, a lone rider guides a majestic horse across an expansive, open field as the sun sets in the background. The rider, dressed in a classic blue shirt and wide-brimmed hat, sits confidently in the saddle, silhouetted against the warm glow of the evening sky. The horse moves gracefully, its mane and tail flowing with each step, creating a sense of harmony between horse and rider. Surrounding the pair, towering trees form a natural border, their leaves gently rustling in the breeze. The shadows lengthen on the ground, accentuating the serene and timeless feel of the scene. The distant hills and wooden fences frame the horizon, adding depth to the tranquil landscape. A few horses graze peacefully in the background, blending into the pastoral setting. The overall ambiance evokes a sense of calmness and quietude, capturing a perfect moment in the golden light of dusk.",
"image_path": "images/mixkit-a-rancher-riding-a-horse-at-sunset-1143.png"
},
{
"prompt": "In the video, a martial artist dressed in a traditional white uniform with a black belt demonstrates a series of precise movements against a stark black background. The individual gracefully transitions between stances, embodying a sense of focused discipline and control. Each motion is executed with a deliberate pace, showcasing the fluidity of martial arts techniques. The soft lighting creates subtle highlights on the uniform, adding depth to the figure as it moves. The practitioner begins with an open-hand pose, feet firmly grounded, gradually shifting to a powerful forward punch. The fluidity of the sequence displays a mastery of balance and poise. Every trajectory of the limbs is precise and deliberate, capturing the elegance and strength of martial arts. The serene, isolated setting enhances the intensity and concentration of the practitioner. This visual presentation is an elegant interplay of motion and stillness, displaying the art form's discipline and grace.",
"image_path": "images/mixkit-a-young-man-practicing-his-karate-moves-49635.png"
},
{
"prompt": "In a serene and softly lit yoga studio, three individuals engage in a yoga session, each performing an upward-facing stretch. The central figure is a woman with shoulder-length brown hair, dressed in a light cropped top and green leggings, her posture reflecting grace and concentration. To her right, another participant, a woman in a purple outfit, mirrors the pose with equal poise. On her left, a person with a bun focuses intently, supported slightly by yoga blocks beneath their hands. The warm-colored wooden floor contrasts soothingly with the soft pastel mural on the back wall, featuring an abstract design and partial visage of a serene face. Natural light floods the space from a large window on the right, where lush greens peek through, adding an element of tranquility. In the corner of the room, a collection of meditation instruments, including a gong and a Buddha statue, subtly frame the peaceful setting. The mood is calm yet focused, as all three participants are deeply engaged in their practice. The scene combines elements of balance, harmony, and a shared journey towards mindfulness. This depiction captures the essence of a yoga session that blends personal growth with collective experience.",
"image_path": "images/mixkit-small-group-of-people-doing-yoga-together-43730.png"
},
{
"prompt": "In the deep blue expanse of the ocean, two dolphins glide effortlessly, their sleek bodies reflecting the sunlight filtering through the water. The prominent shadows and caustics create a shimmering effect on their skin, capturing the beauty of their natural habitat. Each dolphin moves with a fluid grace, occasionally interacting with gentle nudges, showcasing their playful and social nature. The scene is vibrant and dynamic, with the clear blue background accentuating the dolphins' movements, making it an ideal subject for AI recreation.",
"image_path": "images/mixkit-dolphins-underwater-4133.png"
},
{
"prompt": "A bustling ski slope comes alive with skiers descending a pristine, snow-covered hill, surrounded by towering, snow-draped evergreens. Several figures stand atop the slope, silhouetted against a clear blue sky, preparing to embark on their ski run. The chair lift on the right continuously drops off eager adventurers, adding to the excitement at the hilltop. Each skier, clad in colorful winter gear, carves distinct paths into the textured snow as they weave their way down. The interplay of sunlight and shadows accentuates the myriad tracks etched into the slope, creating a dynamic visual rhythm. The scene captures a vibrant winter wonderland, full of action and the thrill of a perfect ski day.",
"image_path": "images/mixkit-skiers-on-a-snowy-slope-3327.png"
},
{
"prompt": "A determined climber is scaling a massive rock face, showcasing exceptional strength and skill. The person, clad in a teal shirt and dark pants, climbs with precision, their movements measured and deliberate. They are secured by climbing gear, which includes ropes and a harness, emphasizing their commitment to safety. The rugged texture of the sandy-colored rock provides an imposing backdrop, adding drama and scale to the climb. In the distance, other large rock formations and sparse vegetation can be seen under a bright, overcast sky, contributing to the natural and adventurous atmosphere. The scene captures a moment of focus and challenge, highlighting the climber's tenacity and the breathtaking environment.",
"image_path": "images/mixkit-alpinist-climbing-a-huge-rock-in-a-desert-43306.png"
},
{
"prompt": "A silver SUV drives along a winding, snow-covered mountain road, with dense pine trees blanketed in snow lining both sides. The scene is serene, with the vehicle moving smoothly, possibly on a winter journey or vacation. As the SUV disappears around the bend, another, darker SUV follows, creating a sense of motion and perspective on the snow-dusted asphalt. The towering, snow-laden rock formation to the right contrasts with the dark green of the pines, highlighting the peacefulness of the wintry landscape.",
"image_path": "images/mixkit-curve-on-a-snowy-forest-road-3317.png"
},
{
"prompt": "A solitary boat glides across the expansive, tranquil expanse of a serene lake. The vessel leaves a gentle wake behind, creating delicate ripples across the mirror-like surface. The water appears a rich shade of teal, seamlessly blending with the sky at the horizon. Silhouettes of distant trees are faintly visible, creating a picturesque backdrop that enhances the solitary journey of the boat. The sky is a calm gradient, shifting from soft oranges near the shore to the pale blues above. In the distance, a few slender poles emerge from the water, remnants of an old structure or natural formation. The mood of the scene is one of peace and solitude, with the boat journeying steadily through the quiet landscape. There is a sense of endless possibilities as the boat moves toward the unseen beyond the frame. The simplicity and stillness of the scene invite contemplation and reflection, encapsulating a perfect moment of quietude on the water.",
"image_path": "images/mixkit-motorboat-on-a-large-lake-with-turquoise-blue-waters-4996.png"
},
{
"prompt": "A man wearing grey shorts jumps rope in a gym, weights and gym equipment in the background.",
"image_path": "images/gray_short_man.jpg"
},
{
"prompt": "Flying over a peninsula covered in bushy trees, while discovering the sea around it, painted a beautiful turquoise blue, on a sunny day.",
"image_path": "images/peninsula.jpg"
},
{
"prompt": "Skillful cyclist doing a wheelie on a bike while riding through a forest, on a dirt road, surrounded by many trees, in the morning.",
"image_path": "images/cyclist.jpg"
},
{
"prompt": "Some friends dancing and having fun together in circles, at a party surrounded by colored lights at a party, in a fancy old place, in a view from below them.",
"image_path": "images/friends.jpg"
},
{
"prompt": "A saxophonist wearing a blazer dances while playing a song in a park.",
"image_path": "images/saxophonist.jpg"
},
{
"prompt": "Romantic couple embracing and looking at each other in the middle of a forest, during a break on a road trip through nature.",
"image_path": "images/romance.jpg"
},
{
"prompt": "Man dressed in 80's style dances very happily in his kitchen while listening to music on his radio and drinking wine.",
"image_path": "images/80s_dance.jpg"
},
{
"prompt": "Pair of jazz musicians performing a song with their saxophone and trombone on an abandoned train.",
"image_path": "images/jazz.jpg"
},
{
"prompt": "A young woman with short hair wearing pink sunglasses chews gum and makes a bubble gum with the city in the background.",
"image_path": "images/pink.jpg"
},
{
"prompt": "Natural aerial landscape with a relief covered with abundant trees and vegetation and a thick layer of mist.",
"image_path": "images/natural.jpg"
},
{
"prompt": "Loving couple sitting on a log on the shore of a lake outside, sharing an affectionate hug.",
"image_path": "images/couple.jpg"
}
]