Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
59c63b4c75 | ||
|
|
bb208e9bab | ||
|
|
18b2c2673d | ||
|
|
65d47c5d29 | ||
|
|
ed02c87a4e | ||
|
|
d14cd2b07a | ||
|
|
86fde63d5d |
@@ -10,8 +10,9 @@ def main():
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
model_id = "FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers"
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
|
||||
model_id,
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
@@ -25,7 +26,7 @@ def main():
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained("FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers")
|
||||
sampling_param = SamplingParam.from_pretrained(model_id)
|
||||
sampling_param.num_frames = 81
|
||||
sampling_param.width = 832
|
||||
sampling_param.height = 480
|
||||
|
||||
@@ -6,8 +6,10 @@ 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
|
||||
|
||||
@@ -15,7 +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",
|
||||
"CausalWan2.2-I2V-A14B-Preview": "FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers",
|
||||
}
|
||||
|
||||
|
||||
@@ -85,28 +87,39 @@ def save_video_from_base64(video_data: str, output_dir: str, prompt: str) -> str
|
||||
return None
|
||||
|
||||
|
||||
def encode_image_to_base64(image_path: str) -> str:
|
||||
"""Encode an image file to base64 string."""
|
||||
if not image_path or not os.path.exists(image_path):
|
||||
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
|
||||
|
||||
mime_types = {
|
||||
'.jpg': 'image/jpeg',
|
||||
'.jpeg': 'image/jpeg',
|
||||
'.png': 'image/png',
|
||||
'.gif': 'image/gif',
|
||||
'.webp': 'image/webp',
|
||||
}
|
||||
|
||||
try:
|
||||
with open(image_path, 'rb') as f:
|
||||
image_bytes = f.read()
|
||||
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')
|
||||
|
||||
# Determine image type from extension
|
||||
ext = os.path.splitext(image_path)[1].lower()
|
||||
mime_types = {
|
||||
'.jpg': 'image/jpeg',
|
||||
'.jpeg': 'image/jpeg',
|
||||
'.png': 'image/png',
|
||||
'.gif': 'image/gif',
|
||||
'.webp': 'image/webp',
|
||||
}
|
||||
mime_type = mime_types.get(ext, 'image/jpeg')
|
||||
|
||||
return f"data:{mime_type};base64,{image_base64}"
|
||||
|
||||
except Exception as e:
|
||||
@@ -426,7 +439,7 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
|
||||
gr.Markdown("**Please make sure you upload a 480x832 image**")
|
||||
input_image = gr.Image(
|
||||
label="",
|
||||
type="filepath",
|
||||
type="pil",
|
||||
height=400,
|
||||
)
|
||||
|
||||
@@ -512,7 +525,17 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
|
||||
if example_label and example_label in example_labels:
|
||||
index = example_labels.index(example_label)
|
||||
selected_prompt = examples[index]
|
||||
selected_image = example_images[index] if index < len(example_images) else None
|
||||
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
|
||||
|
||||
@@ -615,7 +638,7 @@ def main():
|
||||
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",
|
||||
default="FastVideo/CausalWan2.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")
|
||||
@@ -733,6 +756,7 @@ def main():
|
||||
app,
|
||||
demo,
|
||||
path="/gradio",
|
||||
root_path="/gradio",
|
||||
allowed_paths=[
|
||||
os.path.abspath("outputs"),
|
||||
os.path.abspath("fastvideo-logos"),
|
||||
|
||||
@@ -12,4 +12,4 @@ Man dressed in 80's style dances very happily in his kitchen while listening to
|
||||
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.
|
||||
Loving couple sitting on a log on the shore of a lake outside, sharing an affectionate hug.
|
||||
|
||||
@@ -26,7 +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",
|
||||
"FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers",
|
||||
]
|
||||
|
||||
MODEL_CONFIGS = {
|
||||
@@ -272,7 +272,7 @@ class T2VModelDeployment(BaseModelDeployment):
|
||||
|
||||
|
||||
@serve.deployment(
|
||||
ray_actor_options={"num_cpus": 16, "num_gpus": 1, "runtime_env": {"conda": "demo-fv"}},
|
||||
ray_actor_options={"num_cpus": 2, "num_gpus": 1, "runtime_env": {"conda": "fv"}},
|
||||
)
|
||||
class T2V14BModelDeployment(BaseModelDeployment):
|
||||
def __init__(self, t2v_14b_model_path: str, output_path: str = "outputs"):
|
||||
@@ -284,7 +284,7 @@ class T2V14BModelDeployment(BaseModelDeployment):
|
||||
|
||||
|
||||
@serve.deployment(
|
||||
ray_actor_options={"num_cpus": 15, "num_gpus": 1, "runtime_env": {"conda": "demo-fv"}},
|
||||
ray_actor_options={"num_cpus": 2, "num_gpus": 1, "runtime_env": {"conda": "fv"}},
|
||||
)
|
||||
class I2VModelDeployment(BaseModelDeployment):
|
||||
def __init__(self, i2v_model_path: str, output_path: str = "outputs"):
|
||||
@@ -445,7 +445,7 @@ if __name__ == "__main__":
|
||||
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",
|
||||
default="FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers",
|
||||
help="Comma separated list of paths to the I2V model(s)")
|
||||
parser.add_argument("--i2v_model_replicas",
|
||||
type=str,
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
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_paths "FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers" \
|
||||
--i2v_model_replicas "1"
|
||||
|
||||
@@ -159,7 +159,7 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
# 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())
|
||||
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):
|
||||
@@ -201,6 +201,7 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
**pos_cond_kwargs,
|
||||
)
|
||||
if boundary_timestep is not None:
|
||||
assert self.transformer_2 is not None, "transformer_2 is not provided, but boundary_timestep is not None"
|
||||
self.transformer_2(
|
||||
first_frame_latent.to(target_dtype),
|
||||
prompt_embeds,
|
||||
@@ -283,6 +284,7 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
(latent_model_input.shape[0], 1),
|
||||
device=latent_model_input.device,
|
||||
dtype=torch.long)
|
||||
assert current_model is not None
|
||||
pred_noise_btchw = current_model(
|
||||
latent_model_input,
|
||||
prompt_embeds,
|
||||
@@ -372,6 +374,7 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
t_expanded_context = t_context.unsqueeze(1)
|
||||
|
||||
if boundary_timestep is not None:
|
||||
assert self.transformer_2 is not None, "transformer_2 is not provided, but boundary_timestep is not None"
|
||||
self.transformer_2(
|
||||
context_bcthw,
|
||||
prompt_embeds,
|
||||
@@ -384,7 +387,6 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
**image_kwargs,
|
||||
**pos_cond_kwargs,
|
||||
)
|
||||
|
||||
self.transformer(
|
||||
context_bcthw,
|
||||
prompt_embeds,
|
||||
@@ -397,7 +399,6 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
**image_kwargs,
|
||||
**pos_cond_kwargs,
|
||||
)
|
||||
|
||||
start_index += current_num_frames
|
||||
|
||||
if boundary_timestep is not None:
|
||||
|
||||
@@ -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__)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user