Compare commits
8
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
de3436458c | ||
|
|
ed2222f1d1 | ||
|
|
84d1c5c2ec | ||
|
|
6c2ea4e46c | ||
|
|
f61c0685eb | ||
|
|
539a90b2b5 | ||
|
|
84f8e34efe | ||
|
|
f9ef86c649 |
@@ -85,6 +85,9 @@ class PipelineConfig:
|
||||
# DMD parameters
|
||||
dmd_denoising_steps: list[int] | None = field(default=None)
|
||||
|
||||
# Wan2.2 TI2V parameters
|
||||
ti2v_task: bool = False
|
||||
|
||||
# Compilation
|
||||
# enable_torch_compile: bool = False
|
||||
|
||||
|
||||
@@ -8,6 +8,7 @@ from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
from fastvideo.configs.pipelines.wan import (FastWanT2V480PConfig,
|
||||
Wan2_2_TI2V_5B_Config,
|
||||
WanI2V480PConfig, WanI2V720PConfig,
|
||||
WanT2V480PConfig, WanT2V720PConfig)
|
||||
from fastvideo.logger import init_logger
|
||||
@@ -26,8 +27,12 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V720PConfig,
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V720PConfig,
|
||||
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers": FastWanT2V480PConfig,
|
||||
"FastVideo/FastWan2.1-T2V-14B-480P-Diffusers": FastWanT2V480PConfig,
|
||||
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VConfig,
|
||||
"FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers": WanT2V720PConfig,
|
||||
"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,
|
||||
# Add other specific weight variants
|
||||
}
|
||||
|
||||
|
||||
@@ -111,3 +111,26 @@ class FastWanT2V480PConfig(WanT2V480PConfig):
|
||||
def __post_init__(self) -> None:
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_TI2V_5B_Config(WanT2V480PConfig):
|
||||
"""Base configuration for FastWan T2V 1.3B 480P pipeline architecture with DMD"""
|
||||
|
||||
# Denoising stage
|
||||
flow_shift: int = 5
|
||||
ti2v_task: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_T2V_A14B_Config(WanT2V480PConfig):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_I2V_A14B_Config(WanT2V480PConfig):
|
||||
pass
|
||||
|
||||
@@ -6,7 +6,8 @@ from typing import Any
|
||||
from fastvideo.configs.sample.hunyuan import (FastHunyuanSamplingParam,
|
||||
HunyuanSamplingParam)
|
||||
from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
|
||||
from fastvideo.configs.sample.wan import (FastWanT2V480PConfig,
|
||||
from fastvideo.configs.sample.wan import (Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
Wan2_2_TI2V_5B_SamplingParam,
|
||||
WanI2V_14B_480P_SamplingParam,
|
||||
WanI2V_14B_720P_SamplingParam,
|
||||
WanT2V_1_3B_SamplingParam,
|
||||
@@ -24,8 +25,18 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V_14B_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V_14B_480P_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V_14B_720P_SamplingParam,
|
||||
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers":
|
||||
Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VSamplingParam,
|
||||
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers": FastWanT2V480PConfig,
|
||||
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers":
|
||||
Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
"FastVideo/FastWan2.1-T2V-14B-480P-Diffusers":
|
||||
Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
|
||||
# "Wan-AI/Wan2.2-T2V-A14B-Diffusers":
|
||||
# Wan2_2_T2V_A14B_SamplingParam,
|
||||
# "Wan-AI/Wan2.2-I2V-A14B-Diffusers":
|
||||
# Wan2_2_I2V_A14B_SamplingParam,
|
||||
# Add other specific weight variants
|
||||
}
|
||||
|
||||
|
||||
@@ -105,3 +105,48 @@ class FastWanT2V480PConfig(WanT2V_1_3B_SamplingParam):
|
||||
height: int = 448
|
||||
width: int = 832
|
||||
fps: int = 16
|
||||
|
||||
|
||||
# =============================================
|
||||
# ============= Wan2.1 Fun Models =============
|
||||
# =============================================
|
||||
@dataclass
|
||||
class Wan2_1_Fun_1_3B_InP_SamplingParam(SamplingParam):
|
||||
"""Sampling parameters for Wan2.1 Fun 1.3B InP model."""
|
||||
height: int = 480
|
||||
width: int = 832
|
||||
num_frames: int = 81
|
||||
fps: int = 16
|
||||
negative_prompt: str | None = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
|
||||
guidance_scale: float = 6.0
|
||||
num_inference_steps: int = 50
|
||||
|
||||
|
||||
# =============================================
|
||||
# ============= Wan2.2 TI2V Models =============
|
||||
# =============================================
|
||||
@dataclass
|
||||
class Wan2_2_Base_SamplingParam(SamplingParam):
|
||||
"""Sampling parameters for Wan2.2 TI2V 5B model."""
|
||||
negative_prompt: str | None = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_TI2V_5B_SamplingParam(Wan2_2_Base_SamplingParam):
|
||||
"""Sampling parameters for Wan2.2 TI2V 5B model."""
|
||||
height: int = 704
|
||||
width: int = 1280
|
||||
num_frames: int = 121
|
||||
fps: int = 24
|
||||
guidance_scale: float = 5.0
|
||||
num_inference_steps: int = 50
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_T2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_I2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
|
||||
pass
|
||||
|
||||
@@ -9,6 +9,7 @@ diffusion models.
|
||||
import math
|
||||
import os
|
||||
import time
|
||||
from copy import deepcopy
|
||||
from typing import Any
|
||||
|
||||
import imageio
|
||||
@@ -202,6 +203,8 @@ class VideoGenerator:
|
||||
if sampling_param is None:
|
||||
sampling_param = SamplingParam.from_pretrained(
|
||||
fastvideo_args.model_path)
|
||||
else:
|
||||
sampling_param = deepcopy(sampling_param)
|
||||
|
||||
kwargs["prompt"] = prompt
|
||||
sampling_param.update(kwargs)
|
||||
|
||||
@@ -20,6 +20,7 @@ logger = init_logger(__name__)
|
||||
_PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
|
||||
"WanPipeline": "wan",
|
||||
"WanImageToVideoPipeline": "wan",
|
||||
"WanDMDPipeline": "wan",
|
||||
"StepVideoPipeline": "stepvideo",
|
||||
"HunyuanVideoPipeline": "hunyuan",
|
||||
}
|
||||
|
||||
@@ -30,7 +30,7 @@ from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
from fastvideo.utils import dict_to_3d_list
|
||||
from fastvideo.utils import dict_to_3d_list, masks_like
|
||||
|
||||
try:
|
||||
from fastvideo.attention.backends.sliding_tile_attn import (
|
||||
@@ -194,9 +194,21 @@ class DenoisingStage(PipelineStage):
|
||||
# Expand latents for I2V
|
||||
latent_model_input = latents.to(target_dtype)
|
||||
if batch.image_latent is not None:
|
||||
assert not fastvideo_args.pipeline_config.ti2v_task, "image latents should not be provided for TI2V task"
|
||||
latent_model_input = torch.cat(
|
||||
[latent_model_input, batch.image_latent],
|
||||
dim=1).to(target_dtype)
|
||||
elif fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
|
||||
# TI2V directly replaces the first frame of the latent with
|
||||
# the image latent instead of appending along the channel dim
|
||||
assert batch.image_latent is None, "TI2V task should not have image latents"
|
||||
z = self.get_module("vae").encode([batch.pil_image])
|
||||
mask1, mask2 = masks_like([latent_model_input],
|
||||
zero=True,
|
||||
generator=batch.generator)
|
||||
latent_model_input = (
|
||||
1. - mask2[0]) * z[0] + mask2[0] * latent_model_input
|
||||
|
||||
assert torch.isnan(latent_model_input).sum() == 0
|
||||
latent_model_input = self.scheduler.scale_model_input(
|
||||
latent_model_input, t)
|
||||
|
||||
@@ -812,3 +812,37 @@ def set_random_seed(seed: int) -> None:
|
||||
@lru_cache(maxsize=1)
|
||||
def is_vsa_available() -> bool:
|
||||
return importlib.util.find_spec("vsa") is not None
|
||||
|
||||
|
||||
def masks_like(tensor,
|
||||
zero=False,
|
||||
generator=None,
|
||||
p=0.2) -> tuple[list[torch.Tensor], list[torch.Tensor]]:
|
||||
assert isinstance(tensor, list)
|
||||
out1 = [torch.ones(u.shape, dtype=u.dtype, device=u.device) for u in tensor]
|
||||
|
||||
out2 = [torch.ones(u.shape, dtype=u.dtype, device=u.device) for u in tensor]
|
||||
|
||||
if zero:
|
||||
if generator is not None:
|
||||
for u, v in zip(out1, out2, strict=False):
|
||||
random_num = torch.rand(1,
|
||||
generator=generator,
|
||||
device=generator.device).item()
|
||||
if random_num < p:
|
||||
u[:, 0] = torch.normal(mean=-3.5,
|
||||
std=0.5,
|
||||
size=(1, ),
|
||||
device=u.device,
|
||||
generator=generator).expand_as(
|
||||
u[:, 0]).exp()
|
||||
v[:, 0] = torch.zeros_like(v[:, 0])
|
||||
else:
|
||||
u[:, 0] = u[:, 0]
|
||||
v[:, 0] = v[:, 0]
|
||||
else:
|
||||
for u, v in zip(out1, out2, strict=False):
|
||||
u[:, 0] = torch.zeros_like(u[:, 0])
|
||||
v[:, 0] = torch.zeros_like(v[:, 0])
|
||||
|
||||
return out1, out2
|
||||
|
||||
Reference in New Issue
Block a user