Compare commits

..
12 changed files with 161 additions and 90 deletions
+1
View File
@@ -15,6 +15,7 @@ FastVideo features an end-to-end unified pipeline for accelerating diffusion mod
</div>
## NEWS
- ```2025/11/19```: Release [CausalWan2.2 I2V A14B Preview](https://huggingface.co/FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers) models, [Blog](https://hao-ai-lab.github.io/blogs/fastvideo_causalwan_preview/) and [Inference Code!](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_self_forcing_causal_wan2_2_i2v.py)
- ```2025/08/04```: Release [FastWan](https://hao-ai-lab.github.io/FastVideo/distillation/dmd.html) models and [Sparse-Distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
- ```2025/06/14```: Release finetuning and inference code for [VSA](https://arxiv.org/pdf/2505.13389)
- ```2025/04/24```: [FastVideo V1](https://hao-ai-lab.github.io/blogs/fastvideo/) is released!
@@ -26,7 +26,7 @@ def main():
)
sampling_param = SamplingParam.from_pretrained("FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers")
sampling_param.num_frames = 73
sampling_param.num_frames = 81
sampling_param.width = 832
sampling_param.height = 480
sampling_param.seed = 1000
@@ -6,10 +6,8 @@ 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
@@ -87,39 +85,28 @@ def save_video_from_base64(video_data: str, output_dir: str, prompt: str) -> str
return None
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:
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):
return None
mime_types = {
'.jpg': 'image/jpeg',
'.jpeg': 'image/jpeg',
'.png': 'image/png',
'.gif': 'image/gif',
'.webp': 'image/webp',
}
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
with open(image_path, 'rb') as f:
image_bytes = f.read()
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:
@@ -439,7 +426,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="pil",
type="filepath",
height=400,
)
@@ -525,17 +512,7 @@ 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_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
selected_image = example_images[index] if index < len(example_images) else None
return selected_prompt, selected_image
return "", None
@@ -756,7 +733,6 @@ def main():
app,
demo,
path="/gradio",
root_path="/gradio",
allowed_paths=[
os.path.abspath("outputs"),
os.path.abspath("fastvideo-logos"),
+1
View File
@@ -45,6 +45,7 @@ class PipelineConfig:
embedded_cfg_scale: float = 6.0
flow_shift: float | None = None
disable_autocast: bool = False
is_causal: bool = False
# Model configuration
dit_config: DiTConfig = field(default_factory=DiTConfig)
+21 -4
View File
@@ -101,12 +101,19 @@ class BaseLayerWithLoRA(nn.Module):
def set_lora_weights(self,
A: torch.Tensor,
B: torch.Tensor,
lora_alpha: float | None = None,
training_mode: bool = False,
lora_path: str | None = None) -> None:
self.lora_A = torch.nn.Parameter(
A) # share storage with weights in the pipeline
self.lora_B = torch.nn.Parameter(B)
self.disable_lora = False
# Store rank and alpha directly
rank = A.shape[0] # rank is the first dimension of A
self.lora_rank = rank
self.lora_alpha = int(lora_alpha) if lora_alpha is not None else rank
if not training_mode:
self.merge_lora_weights()
self.lora_path = lora_path
@@ -134,8 +141,13 @@ class BaseLayerWithLoRA(nn.Module):
current_device = self.base_layer.weight.data.device
data = self.base_layer.weight.data.to(
get_local_torch_device()).full_tensor()
data += (self.slice_lora_b_weights(self.lora_B).to(data)
@ self.slice_lora_a_weights(self.lora_A).to(data))
# Apply LoRA with alpha scaling
lora_delta = (self.slice_lora_b_weights(self.lora_B).to(data)
@ self.slice_lora_a_weights(self.lora_A).to(data))
if self.lora_alpha and self.lora_rank and self.lora_alpha != self.lora_rank:
lora_delta *= (self.lora_alpha / self.lora_rank)
data += lora_delta
unsharded_base_layer.weight = nn.Parameter(data.to(current_device))
if isinstance(getattr(self.base_layer, "bias", None), DTensor):
unsharded_base_layer.bias = nn.Parameter(
@@ -154,8 +166,13 @@ class BaseLayerWithLoRA(nn.Module):
else:
current_device = self.base_layer.weight.data.device
data = self.base_layer.weight.data.to(get_local_torch_device())
data += \
(self.slice_lora_b_weights(self.lora_B.to(data)) @ self.slice_lora_a_weights(self.lora_A.to(data)))
# Apply LoRA with alpha scaling
lora_delta = (self.slice_lora_b_weights(self.lora_B.to(data))
@ self.slice_lora_a_weights(self.lora_A.to(data)))
if self.lora_alpha and self.lora_rank and self.lora_alpha != self.lora_rank:
lora_delta *= (self.lora_alpha / self.lora_rank)
data += lora_delta
self.base_layer.weight.data = data.to(current_device,
non_blocking=True)
+28 -3
View File
@@ -28,7 +28,8 @@ class LoRAPipeline(ComposedPipelineBase):
TODO: support training.
"""
lora_adapters: dict[str, dict[str, torch.Tensor]] = defaultdict(
dict) # state dicts of loaded lora adapters
dict
) # state dicts of loaded lora adapters (includes lora_A, lora_B, and lora_alpha)
cur_adapter_name: str = ""
cur_adapter_path: str = ""
lora_layers: dict[str, BaseLayerWithLoRA] = {}
@@ -183,11 +184,26 @@ class LoRAPipeline(ComposedPipelineBase):
lora_param_names_mapping_fn = get_param_names_mapping(
self.modules["transformer"].lora_param_names_mapping)
# Extract alpha values and weights in a single pass
to_merge_params: defaultdict[Hashable,
dict[Any, Any]] = defaultdict(dict)
for name, weight in lora_state_dict.items():
# Extract weights (lora_A, lora_B, and lora_alpha)
name = name.replace("diffusion_model.", "")
name = name.replace(".weight", "")
if "lora_alpha" in name:
# Store alpha with minimal mapping - same processing as lora_A/lora_B
# but store in lora_adapters with ".lora_alpha" suffix
layer_name = name.replace(".lora_alpha", "")
layer_name, _, _ = lora_param_names_mapping_fn(layer_name)
target_name, _, _ = param_names_mapping_fn(layer_name)
# Store alpha alongside weights with same target_name base
alpha_key = target_name + ".lora_alpha"
self.lora_adapters[lora_nickname][alpha_key] = weight.item(
) if weight.numel() == 1 else float(weight.mean())
continue
name, _, _ = lora_param_names_mapping_fn(name)
target_name, merge_index, num_params_to_merge = param_names_mapping_fn(
name)
@@ -225,11 +241,20 @@ class LoRAPipeline(ComposedPipelineBase):
for name, layer in self.lora_layers.items():
lora_A_name = name + ".lora_A"
lora_B_name = name + ".lora_B"
lora_alpha_name = name + ".lora_alpha"
if lora_A_name in self.lora_adapters[lora_nickname]\
and lora_B_name in self.lora_adapters[lora_nickname]:
# Get alpha value for this layer (defaults to None if not present)
lora_A = self.lora_adapters[lora_nickname][lora_A_name]
lora_B = self.lora_adapters[lora_nickname][lora_B_name]
# Simple lookup - alpha stored with same naming scheme as lora_A/lora_B
alpha = self.lora_adapters[lora_nickname].get(
lora_alpha_name) if adapter_updated else None
layer.set_lora_weights(
self.lora_adapters[lora_nickname][lora_A_name],
self.lora_adapters[lora_nickname][lora_B_name],
lora_A,
lora_B,
lora_alpha=alpha,
training_mode=self.fastvideo_args.training_mode,
lora_path=lora_path)
adapted_count += 1
+17 -5
View File
@@ -85,10 +85,10 @@ class CausalDMDDenosingStage(DenoisingStage):
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
high_noise_timesteps = timesteps[timesteps >= boundary_timestep]
else:
boundary_timestep = None
high_noise_timesteps = timesteps[timesteps >= boundary_timestep]
high_noise_timesteps = None
# Image kwargs (kept empty unless caller provides compatible args)
image_kwargs: dict = {}
@@ -142,10 +142,18 @@ class CausalDMDDenosingStage(DenoisingStage):
pos_start_base = 0
# Determine block sizes
block_sizes = [self.num_frames_per_block] * 7
block_sizes[0] = 1
if t % self.num_frames_per_block != 0:
raise ValueError(
"num_frames must be divisible by num_frames_per_block for causal DMD denoising"
)
num_blocks = t // self.num_frames_per_block
block_sizes = [self.num_frames_per_block] * num_blocks
start_index = 0
# For now hardcode the first block to be 1 frame assuming the model is Wan2.2-MoE
if boundary_timestep is not None:
block_sizes[0] = 1
first_frame_latent = None
if batch.pil_image is not None:
# Causal video gen directly replaces the first frame of the latent with
@@ -392,6 +400,10 @@ class CausalDMDDenosingStage(DenoisingStage):
start_index += current_num_frames
if boundary_timestep is not None:
num_frames_to_remove = self.num_frames_per_block - 1
latents = latents[:, :, :-num_frames_to_remove, :, :]
batch.latents = latents
return batch
@@ -482,4 +494,4 @@ class CausalDMDDenosingStage(DenoisingStage):
result.add_check(
"negative_prompt_embeds", batch.negative_prompt_embeds, lambda x:
not batch.do_classifier_free_guidance or V.list_not_empty(x))
return result
return result
+19 -19
View File
@@ -5,6 +5,7 @@ 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
@@ -13,6 +14,7 @@ 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__)
@@ -108,27 +110,25 @@ class InputValidationStage(PipelineStage):
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 = 720 * 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 = 480 * 832
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)
scale = max(ow / iw, oh / ih)
img = img.resize((round(iw * scale), round(ih * scale)),
Image.LANCZOS)
# 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("final processed 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)
@@ -1,6 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
import json
import os
import re
import pytest
@@ -52,12 +53,38 @@ LORA_CONFIGS = [
"negative_prompt": "bad quality video,色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走",
"ssim_threshold": 0.79
}
# TODO: Add a LoRA with lora_alpha values to test alpha scaling
#
# Context: This change is mainly for an in-progress ticket porting over LongCat-Video,
# where they used an alpha value that is two times smaller than their rank. This fix
# ensures that LoRA weights are correctly scaled by the alpha/rank ratio when merged.
#
# Issue: Currently, we cannot add a test for LoRA adapters with alpha values because:
# - The existing public LoRAs for Wan-AI/Wan2.1-T2V-1.3B-Diffusers don't store lora_alpha
# - No publicly available LoRA for this model includes lora_alpha tensors in their weights
# - This is why the alpha/rank scaling bug wasn't caught by existing tests
#
# The fix has been validated with:
# - LongCat-Video distilled LoRA (which includes alpha values)
# - Manual testing shows correct alpha/rank scaling behavior
# - Backward compatibility confirmed with LoRAs without alpha values
#
# Future work:
# - Add a synthetic LoRA test fixture with alpha values when feasible
# - Or wait for public Wan LoRAs with alpha to become available
]
MODEL_TO_PARAMS = {
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WAN_LORA_PARAMS,
}
def _sanitize_filename_component(name: str) -> str:
"""Sanitize filename to remove invalid characters (same logic as VideoGenerator)"""
sanitized = re.sub(r'[\\/:*?"<>|]', '', name)
sanitized = sanitized.strip().strip('.')
sanitized = re.sub(r'\s+', ' ', sanitized)
return sanitized or "video"
@pytest.mark.parametrize("model_id", list(MODEL_TO_PARAMS.keys()))
def test_merge_lora_weights(model_id):
lora_config = LORA_CONFIGS[0] # test only one
@@ -137,14 +164,16 @@ def test_lora_inference_similarity(ATTENTION_BACKEND, model_id):
generation_kwargs["negative_prompt"] = lora_config["negative_prompt"]
generator.set_lora_adapter(lora_nickname=lora_nickname, lora_path=lora_path)
# Sanitize the filename before adding .mp4 extension to match VideoGenerator's behavior
output_video_name = f"{lora_path.split('/')[-1]}_{prompt[:50]}"
generation_kwargs["output_path"] = output_dir
generation_kwargs["output_video_name"] = output_video_name
output_video_name = _sanitize_filename_component(output_video_name)
generated_video_path = os.path.join(output_dir, f"{output_video_name}.mp4")
generation_kwargs["output_path"] = generated_video_path
generator.generate_video(prompt, **generation_kwargs)
assert os.path.exists(
output_dir), f"Output video was not generated at {output_dir}"
generated_video_path), f"Output video was not generated at {generated_video_path}"
reference_folder = os.path.join(script_dir, 'L40S_reference_videos', model_id.split('/')[-1], ATTENTION_BACKEND)
@@ -153,13 +182,25 @@ def test_lora_inference_similarity(ATTENTION_BACKEND, model_id):
raise FileNotFoundError(
f"Reference video folder does not exist: {reference_folder}")
# Find the matching reference video for the switched LoRA
# Find the matching reference video - try exact match first, then fuzzy match
# The reference might have different sanitization (e.g., trailing spaces)
reference_video_name = None
unsanitized_prefix = f"{lora_path.split('/')[-1]}_{prompt[:50]}"
for filename in os.listdir(reference_folder):
# Check if the filename starts with the expected output_video_name and ends with .mp4
if filename.startswith(output_video_name) and filename.endswith('.mp4'):
reference_video_name = filename # Remove .mp4 extension to match the logic below
if not filename.endswith('.mp4'):
continue
# Try exact match with sanitized name
if filename.startswith(output_video_name):
reference_video_name = filename
break
# Try match with unsanitized prefix (for legacy reference videos)
# Remove .mp4 and compare the base names after sanitization
base_filename = filename[:-4] # Remove .mp4
if _sanitize_filename_component(base_filename) == output_video_name:
reference_video_name = filename
break
if not reference_video_name:
@@ -167,7 +208,6 @@ def test_lora_inference_similarity(ATTENTION_BACKEND, model_id):
raise FileNotFoundError(f"Reference video missing for adapter {lora_path}")
reference_video_path = os.path.join(reference_folder, reference_video_name)
generated_video_path = os.path.join(output_dir, output_video_name + ".mp4")
logger.info(
f"Computing SSIM between {reference_video_path} and {generated_video_path}"
@@ -729,9 +729,6 @@ class DistillationPipeline(TrainingPipeline):
self.num_train_timestep, [1],
device=self.device,
dtype=torch.long)
world_group = get_world_group()
if world_group.world_size > 1:
world_group.broadcast(timestep, src=0)
timestep = shift_timestep(
timestep,
@@ -844,9 +841,6 @@ class DistillationPipeline(TrainingPipeline):
self.num_train_timestep, [1],
device=self.device,
dtype=torch.long)
world_group = get_world_group()
if world_group.world_size > 1:
world_group.broadcast(fake_score_timestep, src=0)
fake_score_timestep = shift_timestep(
fake_score_timestep,
+4
View File
@@ -467,6 +467,10 @@ class WorkerMultiprocProc:
"output_batch": output_batch.output.cpu(),
"logging_info": logging_info
})
else:
result = self.worker.execute_method(
method, *args, **kwargs)
self.pipe.send(result)
else:
result = self.worker.execute_method(method, *args, **kwargs)
self.pipe.send(result)
+1
View File
@@ -14,6 +14,7 @@ edit_uri: edit/main/docs/
# Configuration
theme:
name: material
favicon: assets/logos/icon_simple.svg
palette:
- scheme: default
toggle: