Compare commits

...
12 Commits
74 changed files with 1734 additions and 1142 deletions
+3 -1
View File
@@ -1,3 +1,5 @@
from fastvideo.v1.configs.pipelines import PipelineConfig
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.entrypoints.video_generator import VideoGenerator
__all__ = ["VideoGenerator"]
__all__ = ["VideoGenerator", "PipelineConfig", "SamplingParam"]
-9
View File
@@ -1,9 +0,0 @@
from fastvideo.v1.configs.base import BaseConfig, SlidingTileAttnConfig
from fastvideo.v1.configs.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.v1.configs.registry import get_pipeline_config_cls_for_name
from fastvideo.v1.configs.wan import WanI2V480PConfig, WanT2V480PConfig
__all__ = [
"HunyuanConfig", "FastHunyuanConfig", "BaseConfig", "SlidingTileAttnConfig",
"WanT2V480PConfig", "WanI2V480PConfig", "get_pipeline_config_cls_for_name"
]
-67
View File
@@ -1,67 +0,0 @@
from dataclasses import dataclass
from typing import Optional
@dataclass
class BaseConfig:
"""Base configuration for all pipeline architectures."""
# Video parameters
height: int = 720
width: int = 1280
num_frames: int = 125
fps: int = 24
# Video generation parameters
num_inference_steps: int = 50
guidance_scale: float = 1.0
seed: int = 1024
guidance_rescale: float = 0.0
embedded_cfg_scale: float = 6.0
flow_shift: Optional[float] = None
use_cpu_offload: bool = False
disable_autocast: bool = False
# Model configuration
precision: str = "bf16"
# VAE configuration
vae_precision: str = "fp16"
vae_tiling: bool = True
vae_sp: bool = True
vae_scale_factor: Optional[int] = None
# DiT configuration
num_channels_latents: Optional[int] = None
# Image encoder configuration
image_encoder_precision: str = "fp32"
# Text encoder configuration
text_encoder_precision: str = "fp16"
text_len: int = -1
hidden_state_skip_layer: int = 0
# STA (Spatial-Temporal Attention) parameters
mask_strategy_file_path: Optional[str] = None
enable_torch_compile: bool = False
neg_prompt: Optional[str] = None
@dataclass
class SlidingTileAttnConfig(BaseConfig):
"""Configuration for sliding tile attention."""
# Override any BaseConfig defaults as needed
# Add sliding tile specific parameters
window_size: int = 16
stride: int = 8
# You can provide custom defaults for inherited fields
height: int = 576
width: int = 1024
# Additional configuration specific to sliding tile attention
pad_to_square: bool = False
use_overlap_optimization: bool = True
+6
View File
@@ -0,0 +1,6 @@
from fastvideo.v1.configs.models.base import ModelConfig
from fastvideo.v1.configs.models.dits.base import DiTConfig
from fastvideo.v1.configs.models.encoders.base import EncoderConfig
from fastvideo.v1.configs.models.vaes.base import VAEConfig
__all__ = ["ModelConfig", "VAEConfig", "DiTConfig", "EncoderConfig"]
+62
View File
@@ -0,0 +1,62 @@
from dataclasses import dataclass, fields
from typing import Any, Dict
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
# 1. ArchConfig contains all fields from diffuser's/transformer's config.json (i.e. all fields related to the architecture of the model)
# 2. ArchConfig should be inherited & overridden by each model arch_config
# 3. Any field in ArchConfig is fixed upon initialization, and should be hidden away from users
@dataclass
class ArchConfig:
pass
@dataclass
class ModelConfig:
# Every model config parameter can be categorized into either ArchConfig or everything else
# Diffuser/Transformer parameters
arch_config: ArchConfig = ArchConfig()
# FastVideo-specific parameters here
# i.e. STA, quantization, teacache
def __getattr__(self, name):
# Only called if 'name' is not found in ModelConfig directly
if hasattr(self.arch_config, name):
return getattr(self.arch_config, name)
raise AttributeError(
f"'{type(self).__name__}' object has no attribute '{name}'")
# This should be used only when loading from transformers/diffusers
def update_model_arch(self, source_model_dict: Dict[str, Any]) -> None:
arch_config = self.arch_config
valid_fields = {f.name for f in fields(arch_config)}
for key, value in source_model_dict.items():
if key in valid_fields:
setattr(arch_config, key, value)
else:
raise AttributeError(
f"{type(arch_config).__name__} has no field '{key}'")
if hasattr(arch_config, "__post_init__"):
arch_config.__post_init__()
def update_model_config(self, source_model_dict: Dict[str, Any]) -> None:
assert "arch_config" not in source_model_dict, "Source model config shouldn't contain arch_config."
valid_fields = {f.name for f in fields(self)}
for key, value in source_model_dict.items():
if key in valid_fields:
setattr(self, key, value)
else:
logger.warning("%s does not contain field '%s'!",
type(self).__name__, key)
raise AttributeError(f"Invalid field: {key}")
if hasattr(self, "__post_init__"):
self.__post_init__()
@@ -0,0 +1,4 @@
from fastvideo.v1.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
from fastvideo.v1.configs.models.dits.wanvideo import WanVideoConfig
__all__ = ["HunyuanVideoConfig", "WanVideoConfig"]
+30
View File
@@ -0,0 +1,30 @@
from dataclasses import dataclass, field
from typing import Optional, Tuple
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
from fastvideo.v1.configs.quantization import QuantizationConfig
from fastvideo.v1.platforms import _Backend
@dataclass
class DiTArchConfig(ArchConfig):
_fsdp_shard_conditions: list = field(default_factory=list)
_param_names_mapping: dict = field(default_factory=dict)
_supported_attention_backends: Tuple[_Backend,
...] = (_Backend.SLIDING_TILE_ATTN,
_Backend.SAGE_ATTN,
_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA)
hidden_size: int = 0
num_attention_heads: int = 0
num_channels_latents: int = 0
@dataclass
class DiTConfig(ModelConfig):
arch_config: DiTArchConfig = DiTArchConfig()
# FastVideoDiT-specific parameters
prefix: str = ""
quant_config: Optional[QuantizationConfig] = None
@@ -0,0 +1,169 @@
from dataclasses import dataclass, field
from typing import Optional, Tuple
import torch
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_double_block(n: str, m) -> bool:
return "double" in n and str.isdigit(n.split(".")[-1])
def is_single_block(n: str, m) -> bool:
return "single" in n and str.isdigit(n.split(".")[-1])
def is_refiner_block(n: str, m) -> bool:
return "refiner" in n and str.isdigit(n.split(".")[-1])
@dataclass
class HunyuanVideoArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(
default_factory=lambda:
[is_double_block, is_single_block, is_refiner_block])
_param_names_mapping: dict = field(
default_factory=lambda: {
# 1. context_embedder.time_text_embed submodules (specific rules, applied first):
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_1\.(.*)$":
r"txt_in.t_embedder.mlp.fc_in.\1",
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_2\.(.*)$":
r"txt_in.t_embedder.mlp.fc_out.\1",
r"^context_embedder\.proj_in\.(.*)$":
r"txt_in.input_embedder.\1",
r"^context_embedder\.time_text_embed\.text_embedder\.linear_1\.(.*)$":
r"txt_in.c_embedder.fc_in.\1",
r"^context_embedder\.time_text_embed\.text_embedder\.linear_2\.(.*)$":
r"txt_in.c_embedder.fc_out.\1",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm1\.(.*)$":
r"txt_in.refiner_blocks.\1.norm1.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm2\.(.*)$":
r"txt_in.refiner_blocks.\1.norm2.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_q\.(.*)$":
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 0, 3),
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_k\.(.*)$":
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 1, 3),
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_v\.(.*)$":
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 2, 3),
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$":
r"txt_in.refiner_blocks.\1.self_attn_proj.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$":
r"txt_in.refiner_blocks.\1.mlp.fc_in.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$":
r"txt_in.refiner_blocks.\1.mlp.fc_out.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm_out\.linear\.(.*)$":
r"txt_in.refiner_blocks.\1.adaLN_modulation.linear.\2",
# 3. x_embedder mapping:
r"^x_embedder\.proj\.(.*)$":
r"img_in.proj.\1",
# 4. Top-level time_text_embed mappings:
r"^time_text_embed\.timestep_embedder\.linear_1\.(.*)$":
r"time_in.mlp.fc_in.\1",
r"^time_text_embed\.timestep_embedder\.linear_2\.(.*)$":
r"time_in.mlp.fc_out.\1",
r"^time_text_embed\.guidance_embedder\.linear_1\.(.*)$":
r"guidance_in.mlp.fc_in.\1",
r"^time_text_embed\.guidance_embedder\.linear_2\.(.*)$":
r"guidance_in.mlp.fc_out.\1",
r"^time_text_embed\.text_embedder\.linear_1\.(.*)$":
r"vector_in.fc_in.\1",
r"^time_text_embed\.text_embedder\.linear_2\.(.*)$":
r"vector_in.fc_out.\1",
# 5. transformer_blocks mapping:
r"^transformer_blocks\.(\d+)\.norm1\.linear\.(.*)$":
r"double_blocks.\1.img_mod.linear.\2",
r"^transformer_blocks\.(\d+)\.norm1_context\.linear\.(.*)$":
r"double_blocks.\1.txt_mod.linear.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$":
r"double_blocks.\1.img_attn_q_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$":
r"double_blocks.\1.img_attn_k_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 0, 3),
r"^transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 1, 3),
r"^transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 2, 3),
r"^transformer_blocks\.(\d+)\.attn\.add_q_proj\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 0, 3),
r"^transformer_blocks\.(\d+)\.attn\.add_k_proj\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 1, 3),
r"^transformer_blocks\.(\d+)\.attn\.add_v_proj\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 2, 3),
r"^transformer_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$":
r"double_blocks.\1.img_attn_proj.\2",
# Corrected: merge attn.to_add_out into the main projection.
r"^transformer_blocks\.(\d+)\.attn\.to_add_out\.(.*)$":
r"double_blocks.\1.txt_attn_proj.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_added_q\.(.*)$":
r"double_blocks.\1.txt_attn_q_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_added_k\.(.*)$":
r"double_blocks.\1.txt_attn_k_norm.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$":
r"double_blocks.\1.img_mlp.fc_in.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$":
r"double_blocks.\1.img_mlp.fc_out.\2",
r"^transformer_blocks\.(\d+)\.ff_context\.net\.0(?:\.proj)?\.(.*)$":
r"double_blocks.\1.txt_mlp.fc_in.\2",
r"^transformer_blocks\.(\d+)\.ff_context\.net\.2(?:\.proj)?\.(.*)$":
r"double_blocks.\1.txt_mlp.fc_out.\2",
# 6. single_transformer_blocks mapping:
r"^single_transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$":
r"single_blocks.\1.q_norm.\2",
r"^single_transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$":
r"single_blocks.\1.k_norm.\2",
r"^single_transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$":
(r"single_blocks.\1.linear1.\2", 0, 4),
r"^single_transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$":
(r"single_blocks.\1.linear1.\2", 1, 4),
r"^single_transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$":
(r"single_blocks.\1.linear1.\2", 2, 4),
r"^single_transformer_blocks\.(\d+)\.proj_mlp\.(.*)$":
(r"single_blocks.\1.linear1.\2", 3, 4),
# Corrected: map proj_out to modulation.linear rather than a separate proj_out branch.
r"^single_transformer_blocks\.(\d+)\.proj_out\.(.*)$":
r"single_blocks.\1.linear2.\2",
r"^single_transformer_blocks\.(\d+)\.norm\.linear\.(.*)$":
r"single_blocks.\1.modulation.linear.\2",
# 7. Final layers mapping:
r"^norm_out\.linear\.(.*)$":
r"final_layer.adaLN_modulation.linear.\1",
r"^proj_out\.(.*)$":
r"final_layer.linear.\1",
})
patch_size: int = 2
patch_size_t: int = 1
in_channels: int = 16
out_channels: int = 16
num_attention_heads: int = 24
attention_head_dim: int = 128
mlp_ratio: float = 4.0
num_layers: int = 20
num_single_layers: int = 40
num_refiner_layers: int = 2
rope_axes_dim: Tuple[int, int, int] = (16, 56, 56)
guidance_embeds: bool = False
dtype: Optional[torch.dtype] = None
text_embed_dim: int = 4096
pooled_projection_dim: int = 768
rope_theta: int = 256
qk_norm: str = "rms_norm"
def __post_init__(self):
self.hidden_size: int = self.attention_head_dim * self.num_attention_heads
self.num_channels_latents: int = self.in_channels
@dataclass
class HunyuanVideoConfig(DiTConfig):
arch_config: DiTArchConfig = HunyuanVideoArchConfig()
prefix: str = "Hunyuan"
@@ -0,0 +1,82 @@
from dataclasses import dataclass, field
from typing import Optional, Tuple
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_blocks(n: str, m) -> bool:
return "blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass
class WanVideoArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
_param_names_mapping: dict = field(
default_factory=lambda: {
r"^patch_embedding\.(.*)$":
r"patch_embedding.proj.\1",
r"^condition_embedder\.text_embedder\.linear_1\.(.*)$":
r"condition_embedder.text_embedder.fc_in.\1",
r"^condition_embedder\.text_embedder\.linear_2\.(.*)$":
r"condition_embedder.text_embedder.fc_out.\1",
r"^condition_embedder\.time_embedder\.linear_1\.(.*)$":
r"condition_embedder.time_embedder.mlp.fc_in.\1",
r"^condition_embedder\.time_embedder\.linear_2\.(.*)$":
r"condition_embedder.time_embedder.mlp.fc_out.\1",
r"^condition_embedder\.time_proj\.(.*)$":
r"condition_embedder.time_modulation.linear.\1",
r"^condition_embedder\.image_embedder\.ff\.net\.0\.proj\.(.*)$":
r"condition_embedder.image_embedder.ff.fc_in.\1",
r"^condition_embedder\.image_embedder\.ff\.net\.2\.(.*)$":
r"condition_embedder.image_embedder.ff.fc_out.\1",
r"^blocks\.(\d+)\.attn1\.to_q\.(.*)$":
r"blocks.\1.to_q.\2",
r"^blocks\.(\d+)\.attn1\.to_k\.(.*)$":
r"blocks.\1.to_k.\2",
r"^blocks\.(\d+)\.attn1\.to_v\.(.*)$":
r"blocks.\1.to_v.\2",
r"^blocks\.(\d+)\.attn1\.to_out\.0\.(.*)$":
r"blocks.\1.to_out.\2",
r"^blocks\.(\d+)\.attn1\.norm_q\.(.*)$":
r"blocks.\1.norm_q.\2",
r"^blocks\.(\d+)\.attn1\.norm_k\.(.*)$":
r"blocks.\1.norm_k.\2",
r"^blocks\.(\d+)\.attn2\.to_out\.0\.(.*)$":
r"blocks.\1.attn2.to_out.\2",
r"^blocks\.(\d+)\.ffn\.net\.0\.proj\.(.*)$":
r"blocks.\1.ffn.fc_in.\2",
r"^blocks\.(\d+)\.ffn\.net\.2\.(.*)$":
r"blocks.\1.ffn.fc_out.\2",
r"blocks\.(\d+)\.norm2\.(.*)$":
r"blocks.\1.self_attn_residual_norm.norm.\2",
})
patch_size: Tuple[int, int, int] = (1, 2, 2)
text_len = 512
num_attention_heads: int = 40
attention_head_dim: int = 128
in_channels: int = 16
out_channels: int = 16
text_dim: int = 4096
freq_dim: int = 256
ffn_dim: int = 13824
num_layers: int = 40
cross_attn_norm: bool = True
qk_norm: str = "rms_norm_across_heads"
eps: float = 1e-6
image_dim: Optional[int] = None
added_kv_proj_dim: Optional[int] = None
rope_max_seq_len: int = 1024
def __post_init__(self):
self.out_channels = self.out_channels or self.in_channels
self.hidden_size = self.num_attention_heads * self.attention_head_dim
self.num_channels_latents = self.in_channels if self.added_kv_proj_dim is None else self.out_channels
@dataclass
class WanVideoConfig(DiTConfig):
arch_config: DiTArchConfig = WanVideoArchConfig()
prefix: str = "Wan"
@@ -0,0 +1,12 @@
from fastvideo.v1.configs.models.encoders.base import (EncoderConfig,
ImageEncoderConfig,
TextEncoderConfig)
from fastvideo.v1.configs.models.encoders.clip import (CLIPTextConfig,
CLIPVisionConfig)
from fastvideo.v1.configs.models.encoders.llama import LlamaConfig
from fastvideo.v1.configs.models.encoders.t5 import T5Config
__all__ = [
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig",
"CLIPTextConfig", "CLIPVisionConfig", "LlamaConfig", "T5Config"
]
@@ -0,0 +1,55 @@
from dataclasses import dataclass, field
from typing import Any, List, Optional, Tuple
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
from fastvideo.v1.configs.quantization import QuantizationConfig
from fastvideo.v1.platforms import _Backend
@dataclass
class EncoderArchConfig(ArchConfig):
architectures: List[str] = field(default_factory=lambda: [])
_supported_attention_backends: Tuple[_Backend, ...] = (_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA)
output_hidden_states: bool = False
use_return_dict: bool = True
@dataclass
class TextEncoderArchConfig(EncoderArchConfig):
vocab_size: int = 0
hidden_size: int = 0
num_hidden_layers: int = 0
num_attention_heads: int = 0
pad_token_id: int = 0
eos_token_id: int = 0
text_len: int = 0
hidden_state_skip_layer: int = 0
decoder_start_token_id: int = 0
output_past: bool = True
scalable_attention: bool = True
tie_word_embeddings: bool = False
@dataclass
class ImageEncoderArchConfig(EncoderArchConfig):
pass
@dataclass
class EncoderConfig(ModelConfig):
arch_config: ArchConfig = EncoderArchConfig()
prefix: str = ""
quant_config: Optional[QuantizationConfig] = None
lora_config: Optional[Any] = None
@dataclass
class TextEncoderConfig(EncoderConfig):
arch_config: ArchConfig = TextEncoderArchConfig()
@dataclass
class ImageEncoderConfig(EncoderConfig):
arch_config: ArchConfig = ImageEncoderArchConfig()
@@ -0,0 +1,64 @@
from dataclasses import dataclass
from typing import Optional
from fastvideo.v1.configs.models.encoders.base import (ImageEncoderArchConfig,
ImageEncoderConfig,
TextEncoderArchConfig,
TextEncoderConfig)
@dataclass
class CLIPTextArchConfig(TextEncoderArchConfig):
vocab_size: int = 49408
hidden_size: int = 512
intermediate_size: int = 2048
projection_dim: int = 512
num_hidden_layers: int = 12
num_attention_heads: int = 8
max_position_embeddings: int = 77
hidden_act: str = "quick_gelu"
layer_norm_eps: float = 1e-5
dropout: float = 0.0
attention_dropout: float = 0.0
initializer_range: float = 0.02
initializer_factor: float = 1.0
pad_token_id: int = 1
bos_token_id: int = 49406
eos_token_id: int = 49407
text_len: int = 77
@dataclass
class CLIPVisionArchConfig(ImageEncoderArchConfig):
hidden_size: int = 768
intermediate_size: int = 3072
projection_dim: int = 512
num_hidden_layers: int = 12
num_attention_heads: int = 12
num_channels: int = 3
image_size: int = 224
patch_size: int = 32
hidden_act: str = "quick_gelu"
layer_norm_eps: float = 1e-5
dropout: float = 0.0
attention_dropout: float = 0.0
initializer_range: float = 0.02
initializer_factor: float = 1.0
@dataclass
class CLIPTextConfig(TextEncoderConfig):
arch_config: TextEncoderArchConfig = CLIPTextArchConfig()
num_hidden_layers_override: Optional[int] = None
require_post_norm: Optional[bool] = None
prefix: str = "clip"
@dataclass
class CLIPVisionConfig(ImageEncoderConfig):
arch_config: ImageEncoderArchConfig = CLIPVisionArchConfig()
num_hidden_layers_override: Optional[int] = None
require_post_norm: Optional[bool] = None
prefix: str = "clip"
@@ -0,0 +1,40 @@
from dataclasses import dataclass
from typing import Optional
from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig,
TextEncoderConfig)
@dataclass
class LlamaArchConfig(TextEncoderArchConfig):
vocab_size: int = 32000
hidden_size: int = 4096
intermediate_size: int = 11008
num_hidden_layers: int = 32
num_attention_heads: int = 32
num_key_value_heads: Optional[int] = None
hidden_act: str = "silu"
max_position_embeddings: int = 2048
initializer_range: float = 0.02
rms_norm_eps: float = 1e-6
use_cache: bool = True
pad_token_id: int = 0
bos_token_id: int = 1
eos_token_id: int = 2
pretraining_tp: int = 1
tie_word_embeddings: bool = False
rope_theta: float = 10000.0
rope_scaling: Optional[float] = None
attention_bias: bool = False
attention_dropout: float = 0.0
mlp_bias: bool = False
head_dim: Optional[int] = None
hidden_state_skip_layer: int = 2
text_len: int = 256
@dataclass
class LlamaConfig(TextEncoderConfig):
arch_config: TextEncoderArchConfig = LlamaArchConfig()
prefix: str = "llama"
@@ -0,0 +1,45 @@
from dataclasses import dataclass
from typing import Optional
from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig,
TextEncoderConfig)
@dataclass
class T5ArchConfig(TextEncoderArchConfig):
vocab_size: int = 32128
d_model: int = 512
d_kv: int = 64
d_ff: int = 2048
num_layers: int = 6
num_decoder_layers: Optional[int] = None
num_heads: int = 8
relative_attention_num_buckets: int = 32
relative_attention_max_distance: int = 128
dropout_rate: float = 0.1
layer_norm_epsilon: float = 1e-6
initializer_factor: float = 1.0
feed_forward_proj: str = "relu"
dense_act_fn: str = ""
is_gated_act: bool = False
is_encoder_decoder: bool = True
use_cache: bool = True
pad_token_id: int = 0
eos_token_id: int = 1
classifier_dropout: float = 0.0
text_len: int = 512
# Referenced from https://github.com/huggingface/transformers/blob/main/src/transformers/models/t5/configuration_t5.py
def __post_init__(self):
act_info = self.feed_forward_proj.split("-")
self.dense_act_fn: str = act_info[-1]
self.is_gated_act: bool = act_info[0] == "gated"
if self.feed_forward_proj == "gated-gelu":
self.dense_act_fn = "gelu_new"
@dataclass
class T5Config(TextEncoderConfig):
arch_config: TextEncoderArchConfig = T5ArchConfig()
prefix: str = "t5"
@@ -0,0 +1,7 @@
from fastvideo.v1.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
from fastvideo.v1.configs.models.vaes.wanvae import WanVAEConfig
__all__ = [
"HunyuanVAEConfig",
"WanVAEConfig",
]
+38
View File
@@ -0,0 +1,38 @@
from dataclasses import dataclass
from typing import Union
import torch
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
@dataclass
class VAEArchConfig(ArchConfig):
scaling_factor: Union[float, torch.tensor] = 0
temporal_compression_ratio: int = 4
spatial_compression_ratio: int = 8
@dataclass
class VAEConfig(ModelConfig):
arch_config: VAEArchConfig = VAEArchConfig()
# FastVideoVAE-specific parameters
load_encoder: bool = True
load_decoder: bool = True
tile_sample_min_height: int = 256
tile_sample_min_width: int = 256
tile_sample_min_num_frames: int = 16
tile_sample_stride_height: int = 192
tile_sample_stride_width: int = 192
tile_sample_stride_num_frames: int = 12
blend_num_frames: int = 0
use_tiling: bool = True
use_temporal_tiling: bool = True
use_parallel_tiling: bool = True
def __post_init__(self):
self.blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
@@ -0,0 +1,40 @@
from dataclasses import dataclass
from typing import Tuple
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class HunyuanVAEArchConfig(VAEArchConfig):
in_channels: int = 3
out_channels: int = 3
latent_channels: int = 16
down_block_types: Tuple[str, ...] = (
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
)
up_block_types: Tuple[str, ...] = (
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
)
block_out_channels: Tuple[int, ...] = (128, 256, 512, 512)
layers_per_block: int = 2
act_fn: str = "silu"
norm_num_groups: int = 32
scaling_factor: float = 0.476986
spatial_compression_ratio: int = 8
temporal_compression_ratio: int = 4
mid_block_add_attention: bool = True
def __post_init__(self):
self.spatial_compression_ratio: int = 2**(len(self.block_out_channels) -
1)
@dataclass
class HunyuanVAEConfig(VAEConfig):
arch_config: VAEArchConfig = HunyuanVAEArchConfig()
@@ -0,0 +1,75 @@
from dataclasses import dataclass
from typing import Tuple
import torch
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class WanVAEArchConfig(VAEArchConfig):
base_dim: int = 96
z_dim: int = 16
dim_mult: Tuple[int, ...] = (1, 2, 4, 4)
num_res_blocks: int = 2
attn_scales: Tuple[float, ...] = ()
temperal_downsample: Tuple[bool, ...] = (False, True, True)
dropout: float = 0.0
latents_mean: Tuple[float, ...] = (
-0.7571,
-0.7089,
-0.9113,
0.1075,
-0.1745,
0.9653,
-0.1517,
1.5508,
0.4134,
-0.0715,
0.5517,
-0.3632,
-0.1922,
-0.9497,
0.2503,
-0.2921,
)
latents_std: Tuple[float, ...] = (
2.8184,
1.4541,
2.3275,
2.6558,
1.2196,
1.7708,
2.6052,
2.0743,
3.2687,
2.1526,
2.8652,
1.5579,
1.6382,
1.1253,
2.8251,
1.9160,
)
temporal_compression_ratio = 4
spatial_compression_ratio = 8
def __post_init__(self):
self.scaling_factor: torch.tensor = 1.0 / torch.tensor(
self.latents_std).view(1, self.z_dim, 1, 1, 1)
self.shift_factor: torch.tensor = torch.tensor(self.latents_mean).view(
1, self.z_dim, 1, 1, 1)
@dataclass
class WanVAEConfig(VAEConfig):
arch_config: VAEArchConfig = WanVAEArchConfig()
use_feature_cache: bool = True
use_tiling: bool = False
use_temporal_tiling: bool = False
use_parallel_tiling: bool = False
def __post_init__(self):
self.blend_num_frames = (self.tile_sample_min_num_frames -
self.tile_sample_stride_num_frames) * 2
@@ -0,0 +1,14 @@
from fastvideo.v1.configs.pipelines.base import (PipelineConfig,
SlidingTileAttnConfig)
from fastvideo.v1.configs.pipelines.hunyuan import (FastHunyuanConfig,
HunyuanConfig)
from fastvideo.v1.configs.pipelines.registry import (
get_pipeline_config_cls_for_name)
from fastvideo.v1.configs.pipelines.wan import (WanI2V480PConfig,
WanT2V480PConfig)
__all__ = [
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
"get_pipeline_config_cls_for_name"
]
+107
View File
@@ -0,0 +1,107 @@
import json
from dataclasses import asdict, dataclass, fields
from typing import Any, Dict, Optional
from fastvideo.v1.configs.models import (DiTConfig, EncoderConfig, ModelConfig,
VAEConfig)
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import shallow_asdict
logger = init_logger(__name__)
@dataclass
class PipelineConfig:
"""Base configuration for all pipeline architectures."""
# Video generation parameters
embedded_cfg_scale: float = 6.0
flow_shift: Optional[float] = None
use_cpu_offload: bool = False
disable_autocast: bool = False
# Model configuration
precision: str = "bf16"
# VAE configuration
vae_precision: str = "fp16"
vae_tiling: bool = True
vae_sp: bool = True
vae_config: VAEConfig = VAEConfig()
# DiT configuration
dit_config: DiTConfig = DiTConfig()
# Text encoder configuration
text_encoder_precision: str = "fp16"
text_encoder_config: EncoderConfig = EncoderConfig()
# STA (Spatial-Temporal Attention) parameters
mask_strategy_file_path: Optional[str] = None
enable_torch_compile: bool = False
@classmethod
def from_pretrained(cls, model_path: str) -> "PipelineConfig":
from fastvideo.v1.configs.pipelines.registry import (
get_pipeline_config_cls_for_name)
pipeline_config_cls = get_pipeline_config_cls_for_name(model_path)
if pipeline_config_cls is not None:
pipeline_config = pipeline_config_cls()
else:
logger.warning(
"Couldn't find an optimal sampling param for %s. Using the default sampling param.",
model_path)
pipeline_config = cls()
return pipeline_config
def dump_to_json(self, file_path: str):
output_dict = shallow_asdict(self)
for key, value in output_dict.items():
if isinstance(value, ModelConfig):
model_dict = asdict(value)
# Model Arch Config should be hidden away from the users
model_dict.pop("arch_config")
output_dict[key] = model_dict
with open(file_path, "w") as f:
json.dump(output_dict, f, indent=2)
def load_from_json(self, file_path: str):
with open(file_path) as f:
input_pipeline_dict = json.load(f)
self.update_pipeline_config(input_pipeline_dict)
def update_pipeline_config(self, source_pipeline_dict: Dict[str,
Any]) -> None:
for f in fields(self):
key = f.name
if key in source_pipeline_dict:
current_value = getattr(self, key)
new_value = source_pipeline_dict[key]
# If it's a nested ModelConfig, update it recursively
if isinstance(current_value, ModelConfig):
current_value.update_model_config(new_value)
else:
setattr(self, key, new_value)
if hasattr(self, "__post_init__"):
self.__post_init__()
@dataclass
class SlidingTileAttnConfig(PipelineConfig):
"""Configuration for sliding tile attention."""
# Override any BaseConfig defaults as needed
# Add sliding tile specific parameters
window_size: int = 16
stride: int = 8
# You can provide custom defaults for inherited fields
height: int = 576
width: int = 1024
# Additional configuration specific to sliding tile attention
pad_to_square: bool = False
use_overlap_optimization: bool = True
@@ -1,21 +1,27 @@
from dataclasses import dataclass
from fastvideo.v1.configs.base import BaseConfig
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.v1.configs.models.dits import HunyuanVideoConfig
from fastvideo.v1.configs.models.encoders import CLIPTextConfig, LlamaConfig
from fastvideo.v1.configs.models.vaes import HunyuanVAEConfig
from fastvideo.v1.configs.pipelines.base import PipelineConfig
@dataclass
class HunyuanConfig(BaseConfig):
class HunyuanConfig(PipelineConfig):
"""Base configuration for HunYuan pipeline architecture."""
# HunyuanConfig-specific parameters with defaults
# DiT
dit_config: DiTConfig = HunyuanVideoConfig()
# VAE
vae_config: VAEConfig = HunyuanVAEConfig()
# Denoising stage
embedded_cfg_scale: int = 6
flow_shift: int = 7
num_inference_steps: int = 50
# Text encoding stage
hidden_state_skip_layer: int = 2
text_len: int = 256
text_encoder_config: EncoderConfig = LlamaConfig()
# Precision for each component
precision: str = "bf16"
@@ -24,8 +30,12 @@ class HunyuanConfig(BaseConfig):
# HunyuanConfig-specific added parameters
# Secondary text encoder
text_encoder_config_2: EncoderConfig = CLIPTextConfig()
text_encoder_precision_2: str = "fp16"
text_len_2: int = 77
def __post_init__(self):
self.vae_config.load_encoder = False
self.vae_config.load_decoder = True
@dataclass
@@ -33,7 +43,6 @@ class FastHunyuanConfig(HunyuanConfig):
"""Configuration specifically optimized for FastHunyuan weights."""
# Override HunyuanConfig defaults
num_inference_steps: int = 6
flow_shift: int = 17
# No need to re-specify guidance_scale or embedded_cfg_scale as they
@@ -3,9 +3,11 @@
import os
from typing import Callable, Dict, Optional, Type
from fastvideo.v1.configs.base import BaseConfig
from fastvideo.v1.configs.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.v1.configs.wan import WanI2V480PConfig, WanT2V480PConfig
from fastvideo.v1.configs.pipelines.base import PipelineConfig
from fastvideo.v1.configs.pipelines.hunyuan import (FastHunyuanConfig,
HunyuanConfig)
from fastvideo.v1.configs.pipelines.wan import (WanI2V480PConfig,
WanT2V480PConfig)
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import (maybe_download_model_index,
verify_model_config_and_directory)
@@ -13,8 +15,8 @@ from fastvideo.v1.utils import (maybe_download_model_index,
logger = init_logger(__name__)
# Registry maps specific model weights to their config classes
WEIGHT_CONFIG_REGISTRY: Dict[str, Type[BaseConfig]] = {
"FastVideo/FastHunyuan-Diffusers": FastHunyuanConfig,
WEIGHT_CONFIG_REGISTRY: Dict[str, Type[PipelineConfig]] = {
"FastVideo/FastHunyuan-diffusers": FastHunyuanConfig,
"hunyuanvideo-community/HunyuanVideo": HunyuanConfig,
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V480PConfig
@@ -30,7 +32,7 @@ PIPELINE_DETECTOR: Dict[str, Callable[[str], bool]] = {
}
# Fallback configs when exact match isn't found but architecture is detected
PIPELINE_FALLBACK_CONFIG: Dict[str, Type[BaseConfig]] = {
PIPELINE_FALLBACK_CONFIG: Dict[str, Type[PipelineConfig]] = {
"hunyuan":
HunyuanConfig, # Base Hunyuan config as fallback for any Hunyuan variant
"wanpipeline":
@@ -41,7 +43,7 @@ PIPELINE_FALLBACK_CONFIG: Dict[str, Type[BaseConfig]] = {
def get_pipeline_config_cls_for_name(
pipeline_name_or_path: str) -> Optional[type[BaseConfig]]:
pipeline_name_or_path: str) -> Optional[type[PipelineConfig]]:
"""Get the appropriate config class for specific pretrained weights."""
if os.path.exists(pipeline_name_or_path):
@@ -65,7 +67,6 @@ def get_pipeline_config_cls_for_name(
# If no match, try to use the fallback config
fallback_config = None
print(pipeline_name)
# Try to determine pipeline architecture for fallback
for pipeline_type, detector in PIPELINE_DETECTOR.items():
if detector(pipeline_name.lower()):
+55
View File
@@ -0,0 +1,55 @@
from dataclasses import dataclass
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.v1.configs.models.dits import WanVideoConfig
from fastvideo.v1.configs.models.encoders import CLIPVisionConfig, T5Config
from fastvideo.v1.configs.models.vaes import WanVAEConfig
from fastvideo.v1.configs.pipelines.base import PipelineConfig
@dataclass
class WanT2V480PConfig(PipelineConfig):
"""Base configuration for Wan T2V 1.3B pipeline architecture."""
# WanConfig-specific parameters with defaults
# DiT
dit_config: DiTConfig = WanVideoConfig()
# VAE
vae_config: VAEConfig = WanVAEConfig()
vae_tiling: bool = False
vae_sp: bool = False
# Video parameters
use_cpu_offload: bool = True
# Denoising stage
flow_shift: int = 3
# Text encoding stage
text_encoder_config: EncoderConfig = T5Config()
# Precision for each component
precision: str = "bf16"
vae_precision: str = "fp16"
text_encoder_precision: str = "fp32"
# WanConfig-specific added parameters
def __post_init__(self):
self.vae_config.load_encoder = False
self.vae_config.load_decoder = True
@dataclass
class WanI2V480PConfig(WanT2V480PConfig):
"""Base configuration for Wan I2V 14B 480P pipeline architecture."""
# WanConfig-specific parameters with defaults
# Precision for each component
image_encoder_config: EncoderConfig = CLIPVisionConfig()
image_encoder_precision: str = "fp32"
def __post_init__(self):
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
@@ -0,0 +1,3 @@
from fastvideo.v1.configs.quantization.base import QuantizationConfig
__all__ = ["QuantizationConfig"]
@@ -0,0 +1,6 @@
from dataclasses import dataclass
@dataclass
class QuantizationConfig:
pass
+3
View File
@@ -0,0 +1,3 @@
from fastvideo.v1.configs.sample.base import SamplingParam
__all__ = ["SamplingParam"]
+72
View File
@@ -0,0 +1,72 @@
from dataclasses import dataclass
from typing import Any, Dict, List, Optional, Union
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
@dataclass
class SamplingParam:
# All fields below are copied from ForwardBatch
data_type: str = "video"
# Image inputs
image_path: Optional[str] = None
# Text inputs
prompt: Optional[Union[str, List[str]]] = None
negative_prompt: Optional[str] = None
prompt_path: Optional[str] = None
output_path: str = "outputs/"
# Batch info
num_videos_per_prompt: int = 1
seed: int = 1024
# Original dimensions (before VAE scaling)
num_frames: int = 125
height: int = 720
width: int = 1280
fps: int = 24
# Denoising parameters
num_inference_steps: int = 50
guidance_scale: float = 1.0
guidance_rescale: float = 0.0
# Misc
save_video: bool = True
return_frames: bool = False
def __post_init__(self) -> None:
self.data_type = "video" if self.num_frames > 1 else "image"
def check_sampling_param(self):
if self.prompt_path and not self.prompt_path.endswith(".txt"):
raise ValueError("prompt_path must be a txt file")
def update(self, source_dict: Dict[str, Any]) -> None:
for key, value in source_dict.items():
if hasattr(self, key):
setattr(self, key, value)
else:
logger.exception("%s has no attribute %s",
type(self).__name__, key)
self.__post_init__()
@classmethod
def from_pretrained(cls, model_path: str) -> "SamplingParam":
from fastvideo.v1.configs.sample.registry import (
get_sampling_param_cls_for_name)
sampling_cls = get_sampling_param_cls_for_name(model_path)
if sampling_cls is not None:
sampling_param: SamplingParam = sampling_cls()
else:
logger.warning(
"Couldn't find an optimal sampling param for %s. Using the default sampling param.",
model_path)
sampling_param = cls()
return sampling_param
+20
View File
@@ -0,0 +1,20 @@
from dataclasses import dataclass
from fastvideo.v1.configs.sample.base import SamplingParam
@dataclass
class HunyuanSamplingParam(SamplingParam):
num_inference_steps: int = 50
num_frames: int = 125
height: int = 720
width: int = 1280
fps: int = 24
guidance_scale: float = 1.0
@dataclass
class FastHunyuanSamplingParam(HunyuanSamplingParam):
num_inference_steps: int = 6
+75
View File
@@ -0,0 +1,75 @@
import os
from typing import Any, Callable, Dict, Optional
from fastvideo.v1.configs.sample.hunyuan import (FastHunyuanSamplingParam,
HunyuanSamplingParam)
from fastvideo.v1.configs.sample.wan import (WanI2V480PSamplingParam,
WanT2V480PSamplingParam)
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import (maybe_download_model_index,
verify_model_config_and_directory)
logger = init_logger(__name__)
# Registry maps specific model weights to their config classes
SAMPLING_PARAM_REGISTRY: Dict[str, Any] = {
"FastVideo/FastHunyuan-diffusers": FastHunyuanSamplingParam,
"hunyuanvideo-community/HunyuanVideo": HunyuanSamplingParam,
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PSamplingParam,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V480PSamplingParam
# Add other specific weight variants
}
# For determining pipeline type from model ID
SAMPLING_PARAM_DETECTOR: Dict[str, Callable[[str], bool]] = {
"hunyuan": lambda id: "hunyuan" in id.lower(),
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
# Add other pipeline architecture detectors
}
# Fallback configs when exact match isn't found but architecture is detected
SAMPLING_FALLBACK_PARAM: Dict[str, Any] = {
"hunyuan":
HunyuanSamplingParam, # Base Hunyuan config as fallback for any Hunyuan variant
"wanpipeline":
WanT2V480PSamplingParam, # Base Wan config as fallback for any Wan variant
"wanimagetovideo": WanI2V480PSamplingParam,
# Other fallbacks by architecture
}
def get_sampling_param_cls_for_name(
pipeline_name_or_path: str) -> Optional[Any]:
"""Get the appropriate sampling param for specific pretrained weights."""
if os.path.exists(pipeline_name_or_path):
config = verify_model_config_and_directory(pipeline_name_or_path)
logger.warning(
"FastVideo may not correctly identify the optimal sampling param for this model, as the local directory may have been renamed."
)
else:
config = maybe_download_model_index(pipeline_name_or_path)
pipeline_name = config["_class_name"]
# First try exact match for specific weights
if pipeline_name_or_path in SAMPLING_PARAM_REGISTRY:
return SAMPLING_PARAM_REGISTRY[pipeline_name_or_path]
# Try partial matches (for local paths that might include the weight ID)
for registered_id, config_class in SAMPLING_PARAM_REGISTRY.items():
if registered_id in pipeline_name_or_path:
return config_class
# If no match, try to use the fallback config
fallback_config = None
# Try to determine pipeline architecture for fallback
for pipeline_type, detector in SAMPLING_PARAM_DETECTOR.items():
if detector(pipeline_name.lower()):
fallback_config = SAMPLING_FALLBACK_PARAM.get(pipeline_type)
break
logger.warning(
"No match found for pipeline %s, using fallback sampling param %s.",
pipeline_name_or_path, fallback_config)
return fallback_config
+24
View File
@@ -0,0 +1,24 @@
from dataclasses import dataclass
from fastvideo.v1.configs.sample.base import SamplingParam
@dataclass
class WanT2V480PSamplingParam(SamplingParam):
# Video parameters
height: int = 480
width: int = 832
num_frames: int = 81
fps: int = 16
# Denoising stage
guidance_scale: float = 3.0
negative_prompt: str = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
num_inference_steps: int = 50
@dataclass
class WanI2V480PSamplingParam(WanT2V480PSamplingParam):
# Denoising stage
guidance_scale: float = 5.0
num_inference_steps: int = 40
-45
View File
@@ -1,45 +0,0 @@
from dataclasses import dataclass
from fastvideo.v1.configs.base import BaseConfig
@dataclass
class WanT2V480PConfig(BaseConfig):
"""Base configuration for Wan T2V 1.3B pipeline architecture."""
# WanConfig-specific parameters with defaults
# Video parameters
height: int = 480
width: int = 832
num_frames: int = 81
fps: int = 16
use_cpu_offload: bool = True
# Denoising stage
guidance_scale: float = 3.0
neg_prompt: str = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
flow_shift: int = 3
num_inference_steps: int = 50
# Text encoding stage
text_len: int = 512
# Precision for each component
precision: str = "bf16"
vae_precision: str = "fp16"
text_encoder_precision: str = "fp32"
# WanConfig-specific added parameters
@dataclass
class WanI2V480PConfig(WanT2V480PConfig):
"""Base configuration for Wan I2V 14B 480P pipeline architecture."""
# WanConfig-specific parameters with defaults
# Denoising stage
guidance_scale: float = 5.0
num_inference_steps: int = 40
# Precision for each component
image_encoder_precision: str = "fp32"
@@ -0,0 +1,37 @@
{
"embedded_cfg_scale": 6.0,
"flow_shift": 3,
"use_cpu_offload": true,
"disable_autocast": false,
"precision": "bf16",
"vae_precision": "fp16",
"vae_tiling": false,
"vae_sp": false,
"vae_config": {
"load_encoder": false,
"load_decoder": true,
"tile_sample_min_height": 256,
"tile_sample_min_width": 256,
"tile_sample_min_num_frames": 16,
"tile_sample_stride_height": 192,
"tile_sample_stride_width": 192,
"tile_sample_stride_num_frames": 12,
"blend_num_frames": 8,
"use_tiling": false,
"use_temporal_tiling": false,
"use_parallel_tiling": false,
"use_feature_cache": true
},
"dit_config": {
"prefix": "Wan",
"quant_config": null
},
"text_encoder_precision": "fp32",
"text_encoder_config": {
"prefix": "t5",
"quant_config": null,
"lora_config": null
},
"mask_strategy_file_path": null,
"enable_torch_compile": false
}
@@ -0,0 +1,45 @@
{
"embedded_cfg_scale": 6.0,
"flow_shift": 3,
"use_cpu_offload": true,
"disable_autocast": false,
"precision": "bf16",
"vae_precision": "fp16",
"vae_tiling": false,
"vae_sp": false,
"vae_config": {
"load_encoder": true,
"load_decoder": true,
"tile_sample_min_height": 256,
"tile_sample_min_width": 256,
"tile_sample_min_num_frames": 16,
"tile_sample_stride_height": 192,
"tile_sample_stride_width": 192,
"tile_sample_stride_num_frames": 12,
"blend_num_frames": 8,
"use_tiling": false,
"use_temporal_tiling": false,
"use_parallel_tiling": false,
"use_feature_cache": true
},
"dit_config": {
"prefix": "Wan",
"quant_config": null
},
"text_encoder_precision": "fp32",
"text_encoder_config": {
"prefix": "t5",
"quant_config": null,
"lora_config": null
},
"mask_strategy_file_path": null,
"enable_torch_compile": false,
"image_encoder_config": {
"prefix": "clip",
"quant_config": null,
"lora_config": null,
"num_hidden_layers_override": null,
"require_post_norm": null
},
"image_encoder_precision": "fp32"
}
+59 -77
View File
@@ -9,7 +9,7 @@ diffusion models.
import os
import time
from dataclasses import asdict
from typing import Any, Callable, Dict, List, Optional, Union
from typing import Any, Dict, List, Optional, Union
import imageio
import numpy as np
@@ -17,11 +17,13 @@ import torch
import torchvision
from einops import rearrange
from fastvideo.v1.configs import get_pipeline_config_cls_for_name
from fastvideo.v1.configs.pipelines import (PipelineConfig,
get_pipeline_config_cls_for_name)
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines import ForwardBatch
from fastvideo.v1.utils import align_to
from fastvideo.v1.utils import align_to, shallow_asdict
from fastvideo.v1.worker.executor import Executor
logger = init_logger(__name__)
@@ -52,6 +54,9 @@ class VideoGenerator:
model_path: str,
device: Optional[str] = None,
torch_dtype: Optional[torch.dtype] = None,
pipeline_config: Optional[
Union[str
| PipelineConfig]] = None,
**kwargs) -> "VideoGenerator":
"""
Create a video generator from a pretrained model.
@@ -64,22 +69,30 @@ class VideoGenerator:
Returns:
The created video generator
Priority level: Default pipeline config < User's pipeline config < User's kwargs
"""
config = None
config_cls = get_pipeline_config_cls_for_name(model_path)
if config_cls is not None:
config = config_cls()
# 1. If users provide a pipeline config, it will override the default pipeline config
if isinstance(pipeline_config, PipelineConfig):
config = pipeline_config
else:
config_cls = get_pipeline_config_cls_for_name(model_path)
if config_cls is not None:
config = config_cls()
if isinstance(pipeline_config, str):
config.load_from_json(pipeline_config)
# 2. If users also provide some kwargs, it will override the pipeline config.
# The user kwargs shouldn't contain model config parameters!
if config is None:
logger.warning("No config found for model %s, using default config",
model_path)
config_args = {}
config_args = kwargs
else:
config_args = asdict(config)
# override config_args with kwargs
config_args.update(kwargs)
config_args = shallow_asdict(config)
config_args.update(kwargs)
fastvideo_args = FastVideoArgs(
model_path=model_path,
@@ -115,19 +128,8 @@ class VideoGenerator:
def generate_video(
self,
prompt: str,
negative_prompt: Optional[str] = None,
output_path: Optional[str] = None,
save_video: bool = True,
return_frames: bool = False,
num_inference_steps: Optional[int] = None,
guidance_scale: Optional[float] = None,
num_frames: Optional[int] = None,
height: Optional[int] = None,
width: Optional[int] = None,
fps: Optional[int] = None,
seed: Optional[int] = None,
callback: Optional[Callable[[int, int, torch.Tensor], None]] = None,
callback_steps: int = 1,
sampling_param: Optional[SamplingParam] = None,
**kwargs,
) -> Union[Dict[str, Any], List[np.ndarray]]:
"""
Generate a video based on the given prompt.
@@ -154,87 +156,68 @@ class VideoGenerator:
# Create a copy of inference args to avoid modifying the original
fastvideo_args = self.fastvideo_args
# Override parameters if provided
if negative_prompt is not None:
fastvideo_args.neg_prompt = negative_prompt
if num_inference_steps is not None:
fastvideo_args.num_inference_steps = num_inference_steps
if guidance_scale is not None:
fastvideo_args.guidance_scale = guidance_scale
if num_frames is not None:
fastvideo_args.num_frames = num_frames
if height is not None:
fastvideo_args.height = height
if width is not None:
fastvideo_args.width = width
if fps is not None:
fastvideo_args.fps = fps
if seed is not None:
fastvideo_args.seed = seed
# Validate inputs
if not isinstance(prompt, str):
raise TypeError(
f"`prompt` must be a string, but got {type(prompt)}")
prompt = prompt.strip()
if sampling_param is None:
sampling_param = SamplingParam.from_pretrained(
fastvideo_args.model_path)
kwargs["prompt"] = prompt
sampling_param.update(kwargs)
# Process negative prompt
if fastvideo_args.neg_prompt is not None:
fastvideo_args.neg_prompt = fastvideo_args.neg_prompt.strip()
if sampling_param.negative_prompt is not None:
sampling_param.negative_prompt = sampling_param.negative_prompt.strip(
)
# Validate dimensions
if (fastvideo_args.height <= 0 or fastvideo_args.width <= 0
or fastvideo_args.num_frames <= 0):
if (sampling_param.height <= 0 or sampling_param.width <= 0
or sampling_param.num_frames <= 0):
raise ValueError(
f"Height, width, and num_frames must be positive integers, got "
f"height={fastvideo_args.height}, width={fastvideo_args.width}, "
f"num_frames={fastvideo_args.num_frames}")
f"height={sampling_param.height}, width={sampling_param.width}, "
f"num_frames={sampling_param.num_frames}")
if (fastvideo_args.num_frames - 1) % 4 != 0:
if (
sampling_param.num_frames - 1
) % fastvideo_args.vae_config.arch_config.temporal_compression_ratio != 0:
raise ValueError(
f"num_frames-1 must be a multiple of 4, got {fastvideo_args.num_frames}"
f"num_frames-1 must be a multiple of {fastvideo_args.vae_config.arch_config.temporal_compression_ratio}, got {sampling_param.num_frames}"
)
# Calculate sizes
target_height = align_to(fastvideo_args.height, 16)
target_width = align_to(fastvideo_args.width, 16)
target_height = align_to(sampling_param.height, 16)
target_width = align_to(sampling_param.width, 16)
# Calculate latent sizes
latents_size = [(fastvideo_args.num_frames - 1) // 4 + 1,
fastvideo_args.height // 8, fastvideo_args.width // 8]
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8, sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
# Log parameters
debug_str = f"""
height: {target_height}
width: {target_width}
video_length: {fastvideo_args.num_frames}
video_length: {sampling_param.num_frames}
prompt: {prompt}
neg_prompt: {fastvideo_args.neg_prompt}
seed: {fastvideo_args.seed}
infer_steps: {fastvideo_args.num_inference_steps}
num_videos_per_prompt: {fastvideo_args.num_videos}
guidance_scale: {fastvideo_args.guidance_scale}
neg_prompt: {sampling_param.negative_prompt}
seed: {sampling_param.seed}
infer_steps: {sampling_param.num_inference_steps}
num_videos_per_prompt: {sampling_param.num_videos_per_prompt}
guidance_scale: {sampling_param.guidance_scale}
n_tokens: {n_tokens}
flow_shift: {fastvideo_args.flow_shift}
embedded_guidance_scale: {fastvideo_args.embedded_cfg_scale}"""
logger.info(debug_str)
# Prepare batch
device = torch.device(fastvideo_args.device_str)
batch = ForwardBatch(
prompt=prompt,
negative_prompt=fastvideo_args.neg_prompt,
num_videos_per_prompt=fastvideo_args.num_videos,
height=fastvideo_args.height,
width=fastvideo_args.width,
num_frames=fastvideo_args.num_frames,
num_inference_steps=fastvideo_args.num_inference_steps,
guidance_scale=fastvideo_args.guidance_scale,
**asdict(sampling_param),
eta=0.0,
n_tokens=n_tokens,
data_type="video" if fastvideo_args.num_frames > 1 else "image",
device=device,
extra={},
)
@@ -255,23 +238,22 @@ class VideoGenerator:
frames.append((x * 255).numpy().astype(np.uint8))
# Save video if requested
if save_video:
save_path = output_path or fastvideo_args.output_path
if batch.save_video:
save_path = batch.output_path
if save_path:
os.makedirs(os.path.dirname(save_path), exist_ok=True)
video_path = os.path.join(save_path, f"{prompt[:100]}.mp4")
imageio.mimsave(video_path, frames, fps=fastvideo_args.fps)
imageio.mimsave(video_path, frames, fps=batch.fps)
logger.info("Saved video to %s", video_path)
else:
logger.warning("No output path provided, video not saved")
if return_frames:
if batch.return_frames:
return frames
else:
return {
"samples": samples,
"prompts": prompt,
"size":
(target_height, target_width, fastvideo_args.num_frames),
"size": (target_height, target_width, batch.num_frames),
"generation_time": gen_time
}
@@ -1,4 +1,5 @@
from fastvideo import VideoGenerator
# from fastvideo.v1.configs.sample import SamplingParam
def main():
# FastVideo will automatically use the optimal default arguments for the
@@ -11,9 +12,13 @@ def main():
num_gpus=4,
)
# sampling_param = SamplingParam.from_pretrained("/workspace/data/Wan-AI/Wan2.1-I2V-14B-480P-Diffusers")
# sampling_param.num_frames = 45
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
# Generate videos with the same simple API, regardless of GPU count
prompt = "A beautiful woman in a red dress walking down a street"
video = generator.generate_video(prompt)
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
# Generate another video with a different prompt, without reloading the
# model!
+10 -153
View File
@@ -7,6 +7,7 @@ import dataclasses
from contextlib import contextmanager
from typing import List, Optional
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import FlexibleArgumentParser
@@ -34,55 +35,38 @@ class FastVideoArgs:
dist_timeout: Optional[int] = None # timeout for torch.distributed
# Video generation parameters
height: int = 720
width: int = 1280
num_frames: int = 117
num_inference_steps: int = 50
guidance_scale: float = 1.0
guidance_rescale: float = 0.0
embedded_cfg_scale: float = 6.0
flow_shift: Optional[float] = None
output_type: str = "pil"
# Model configuration
# DiT configuration
dit_config: DiTConfig = DiTConfig()
precision: str = "bf16"
# VAE configuration
vae_precision: str = "fp16"
vae_tiling: bool = True
vae_sp: bool = False
vae_scale_factor: Optional[int] = None
# DiT configuration
num_channels_latents: Optional[int] = None
vae_tiling: bool = True # Might change in between forward passes
vae_sp: bool = False # Might change in between forward passes
# vae_scale_factor: Optional[int] = None # Deprecated
vae_config: VAEConfig = VAEConfig()
# Image encoder configuration
image_encoder_precision: str = "fp32"
image_encoder_config: EncoderConfig = EncoderConfig()
# Text encoder configuration
text_encoder_precision: str = "fp16"
text_len: int = 256
hidden_state_skip_layer: int = 2
text_encoder_config: EncoderConfig = EncoderConfig()
# Secondary text encoder
text_encoder_config_2: EncoderConfig = EncoderConfig()
text_encoder_precision_2: str = "fp16"
text_len_2: int = 77
# Flow Matching parameters
flow_solver: str = "euler"
denoise_type: str = "flow" # Deprecated. Will use scheduler_config.json
# STA (Spatial-Temporal Attention) parameters
mask_strategy_file_path: Optional[str] = None
enable_torch_compile: bool = False
# Scheduler options
scheduler_type: str = "euler" # Deprecated. Will use the param in scheduler_config.json
neg_prompt: Optional[str] = None
num_videos: int = 1
fps: int = 24
use_cpu_offload: bool = False
disable_autocast: bool = False
@@ -90,11 +74,6 @@ class FastVideoArgs:
log_level: str = "info"
# Inference parameters
image_path: Optional[str] = None
prompt: Optional[str] = None
prompt_path: Optional[str] = None
output_path: str = "outputs/"
seed: int = 1024
device_str: Optional[str] = None
device = None
@@ -174,43 +153,6 @@ class FastVideoArgs:
help="Set timeout for torch.distributed initialization.",
)
# Video generation parameters
parser.add_argument(
"--height",
type=int,
default=FastVideoArgs.height,
help="Height of generated video",
)
parser.add_argument(
"--width",
type=int,
default=FastVideoArgs.width,
help="Width of generated video",
)
parser.add_argument(
"--num-frames",
type=int,
default=FastVideoArgs.num_frames,
help="Number of frames to generate",
)
parser.add_argument(
"--num-inference-steps",
type=int,
default=FastVideoArgs.num_inference_steps,
help="Number of inference steps",
)
parser.add_argument(
"--guidance-scale",
type=float,
default=FastVideoArgs.guidance_scale,
help="Guidance scale for classifier-free guidance",
)
parser.add_argument(
"--guidance-rescale",
type=float,
default=FastVideoArgs.guidance_rescale,
help="Guidance rescale for classifier-free guidance",
)
parser.add_argument(
"--embedded-cfg-scale",
type=float,
@@ -267,12 +209,6 @@ class FastVideoArgs:
choices=["fp32", "fp16", "bf16"],
help="Precision for text encoder",
)
parser.add_argument(
"--text-len",
type=int,
default=FastVideoArgs.text_len,
help="Maximum text length",
)
# Image encoder config
parser.add_argument(
@@ -292,26 +228,6 @@ class FastVideoArgs:
choices=["fp32", "fp16", "bf16"],
help="Precision for secondary text encoder",
)
parser.add_argument(
"--text-len-2",
type=int,
default=FastVideoArgs.text_len_2,
help="Maximum secondary text length",
)
# Flow Matching parameters
parser.add_argument(
"--flow-solver",
type=str,
default=FastVideoArgs.flow_solver,
help="Solver for flow matching",
)
parser.add_argument(
"--denoise-type",
type=str,
default=FastVideoArgs.denoise_type,
help="Denoise type for noised inputs",
)
# STA (Spatial-Temporal Attention) parameters
parser.add_argument(
@@ -326,33 +242,6 @@ class FastVideoArgs:
"Use torch.compile for speeding up STA inference without teacache",
)
# Scheduler options
parser.add_argument(
"--scheduler-type",
type=str,
default=FastVideoArgs.scheduler_type,
help="Type of scheduler to use",
)
# HunYuan specific parameters
parser.add_argument(
"--neg-prompt",
type=str,
default=FastVideoArgs.neg_prompt,
help="Negative prompt for sampling",
)
parser.add_argument(
"--num-videos",
type=int,
default=FastVideoArgs.num_videos,
help="Number of videos to generate per prompt",
)
parser.add_argument(
"--fps",
type=int,
default=FastVideoArgs.fps,
help="Frames per second for output video",
)
parser.add_argument(
"--use-cpu-offload",
action="store_true",
@@ -373,36 +262,6 @@ class FastVideoArgs:
help="The logging level of all loggers.",
)
# Inference parameters
prompt_group = parser.add_mutually_exclusive_group(required=True)
prompt_group.add_argument(
"--prompt",
type=str,
help="Text prompt for video generation",
)
prompt_group.add_argument(
"--prompt-path",
type=str,
help="Path to a text file containing the prompt",
)
parser.add_argument("--image-path",
type=str,
help="Path to the image for I2V generation")
parser.add_argument(
"--output-path",
type=str,
default=FastVideoArgs.output_path,
help="Directory to save generated videos",
)
parser.add_argument(
"--seed",
type=int,
default=FastVideoArgs.seed,
help="Random seed for reproducibility",
)
return parser
@classmethod
@@ -451,8 +310,6 @@ class FastVideoArgs:
raise ValueError(
"Currently enabling vae_sp requires enabling vae_tiling, please set --vae-tiling to True."
)
if self.prompt_path and not self.prompt_path.endswith(".txt"):
raise ValueError("prompt_path must be a text file")
_current_fastvideo_args = None
+2
View File
@@ -1,3 +1,5 @@
# type: ignore
# SPDX-License-Identifier: Apache-2.0
"""
Inference module for diffusion models.
+9 -5
View File
@@ -5,19 +5,20 @@ from typing import List, Optional, Tuple, Union
import torch
from torch import nn
from fastvideo.v1.configs.models import DiTConfig
from fastvideo.v1.platforms import _Backend
# TODO
class BaseDiT(nn.Module, ABC):
_fsdp_shard_conditions: list = []
attention_head_dim: int | None = None
_param_names_mapping: dict
hidden_size: int
num_attention_heads: int
num_channels_latents: int
# always supports torch_sdpa
_supported_attention_backends: Tuple[_Backend,
...] = (_Backend.TORCH_SDPA, )
_supported_attention_backends: Tuple[
_Backend, ...] = DiTConfig()._supported_attention_backends
def __init_subclass__(cls) -> None:
required_class_attrs = [
@@ -30,8 +31,9 @@ class BaseDiT(nn.Module, ABC):
f"Subclasses of BaseDiT must define '{attr}' class variable"
)
def __init__(self, *args, **kwargs) -> None:
def __init__(self, config: DiTConfig, **kwargs) -> None:
super().__init__()
self.config = config
if not self.supported_attention_backends:
raise ValueError(
f"Subclass {self.__class__.__name__} must define _supported_attention_backends"
@@ -49,7 +51,9 @@ class BaseDiT(nn.Module, ABC):
pass
def __post_init__(self) -> None:
required_attrs = ["hidden_size", "num_attention_heads"]
required_attrs = [
"hidden_size", "num_attention_heads", "num_channels_latents"
]
for attr in required_attrs:
if not hasattr(self, attr):
raise AttributeError(
+56 -190
View File
@@ -6,6 +6,7 @@ import torch
import torch.nn as nn
from fastvideo.v1.attention import DistributedAttention, LocalAttention
from fastvideo.v1.configs.models.dits import HunyuanVideoConfig
from fastvideo.v1.distributed.parallel_state import (
get_sequence_model_parallel_world_size)
from fastvideo.v1.layers.layernorm import (LayerNormScaleShift, ScaleResidual,
@@ -431,239 +432,104 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
# PY: we make the input args the same as HF config
# shard single stream, double stream blocks, and refiner_blocks
_fsdp_shard_conditions = [
lambda n, m: "double" in n and str.isdigit(n.split(".")[-1]),
lambda n, m: "single" in n and str.isdigit(n.split(".")[-1]),
lambda n, m: "refiner" in n and str.isdigit(n.split(".")[-1]),
]
_supported_attention_backends = (_Backend.SLIDING_TILE_ATTN,
_Backend.SAGE_ATTN, _Backend.FLASH_ATTN,
_Backend.TORCH_SDPA)
_param_names_mapping = {
# 1. context_embedder.time_text_embed submodules (specific rules, applied first):
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_1\.(.*)$":
r"txt_in.t_embedder.mlp.fc_in.\1",
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_2\.(.*)$":
r"txt_in.t_embedder.mlp.fc_out.\1",
r"^context_embedder\.proj_in\.(.*)$":
r"txt_in.input_embedder.\1",
r"^context_embedder\.time_text_embed\.text_embedder\.linear_1\.(.*)$":
r"txt_in.c_embedder.fc_in.\1",
r"^context_embedder\.time_text_embed\.text_embedder\.linear_2\.(.*)$":
r"txt_in.c_embedder.fc_out.\1",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm1\.(.*)$":
r"txt_in.refiner_blocks.\1.norm1.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm2\.(.*)$":
r"txt_in.refiner_blocks.\1.norm2.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_q\.(.*)$":
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 0, 3),
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_k\.(.*)$":
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 1, 3),
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_v\.(.*)$":
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 2, 3),
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$":
r"txt_in.refiner_blocks.\1.self_attn_proj.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$":
r"txt_in.refiner_blocks.\1.mlp.fc_in.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$":
r"txt_in.refiner_blocks.\1.mlp.fc_out.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm_out\.linear\.(.*)$":
r"txt_in.refiner_blocks.\1.adaLN_modulation.linear.\2",
_fsdp_shard_conditions = HunyuanVideoConfig()._fsdp_shard_conditions
_supported_attention_backends = HunyuanVideoConfig(
)._supported_attention_backends
_param_names_mapping = HunyuanVideoConfig()._param_names_mapping
# 3. x_embedder mapping:
r"^x_embedder\.proj\.(.*)$":
r"img_in.proj.\1",
def __init__(self, config: HunyuanVideoConfig):
super().__init__(config=config)
# 4. Top-level time_text_embed mappings:
r"^time_text_embed\.timestep_embedder\.linear_1\.(.*)$":
r"time_in.mlp.fc_in.\1",
r"^time_text_embed\.timestep_embedder\.linear_2\.(.*)$":
r"time_in.mlp.fc_out.\1",
r"^time_text_embed\.guidance_embedder\.linear_1\.(.*)$":
r"guidance_in.mlp.fc_in.\1",
r"^time_text_embed\.guidance_embedder\.linear_2\.(.*)$":
r"guidance_in.mlp.fc_out.\1",
r"^time_text_embed\.text_embedder\.linear_1\.(.*)$":
r"vector_in.fc_in.\1",
r"^time_text_embed\.text_embedder\.linear_2\.(.*)$":
r"vector_in.fc_out.\1",
# 5. transformer_blocks mapping:
r"^transformer_blocks\.(\d+)\.norm1\.linear\.(.*)$":
r"double_blocks.\1.img_mod.linear.\2",
r"^transformer_blocks\.(\d+)\.norm1_context\.linear\.(.*)$":
r"double_blocks.\1.txt_mod.linear.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$":
r"double_blocks.\1.img_attn_q_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$":
r"double_blocks.\1.img_attn_k_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 0, 3),
r"^transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 1, 3),
r"^transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 2, 3),
r"^transformer_blocks\.(\d+)\.attn\.add_q_proj\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 0, 3),
r"^transformer_blocks\.(\d+)\.attn\.add_k_proj\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 1, 3),
r"^transformer_blocks\.(\d+)\.attn\.add_v_proj\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 2, 3),
r"^transformer_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$":
r"double_blocks.\1.img_attn_proj.\2",
# Corrected: merge attn.to_add_out into the main projection.
r"^transformer_blocks\.(\d+)\.attn\.to_add_out\.(.*)$":
r"double_blocks.\1.txt_attn_proj.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_added_q\.(.*)$":
r"double_blocks.\1.txt_attn_q_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_added_k\.(.*)$":
r"double_blocks.\1.txt_attn_k_norm.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$":
r"double_blocks.\1.img_mlp.fc_in.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$":
r"double_blocks.\1.img_mlp.fc_out.\2",
r"^transformer_blocks\.(\d+)\.ff_context\.net\.0(?:\.proj)?\.(.*)$":
r"double_blocks.\1.txt_mlp.fc_in.\2",
r"^transformer_blocks\.(\d+)\.ff_context\.net\.2(?:\.proj)?\.(.*)$":
r"double_blocks.\1.txt_mlp.fc_out.\2",
# 6. single_transformer_blocks mapping:
r"^single_transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$":
r"single_blocks.\1.q_norm.\2",
r"^single_transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$":
r"single_blocks.\1.k_norm.\2",
r"^single_transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$":
(r"single_blocks.\1.linear1.\2", 0, 4),
r"^single_transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$":
(r"single_blocks.\1.linear1.\2", 1, 4),
r"^single_transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$":
(r"single_blocks.\1.linear1.\2", 2, 4),
r"^single_transformer_blocks\.(\d+)\.proj_mlp\.(.*)$":
(r"single_blocks.\1.linear1.\2", 3, 4),
# Corrected: map proj_out to modulation.linear rather than a separate proj_out branch.
r"^single_transformer_blocks\.(\d+)\.proj_out\.(.*)$":
r"single_blocks.\1.linear2.\2",
r"^single_transformer_blocks\.(\d+)\.norm\.linear\.(.*)$":
r"single_blocks.\1.modulation.linear.\2",
# 7. Final layers mapping:
r"^norm_out\.linear\.(.*)$":
r"final_layer.adaLN_modulation.linear.\1",
r"^proj_out\.(.*)$":
r"final_layer.linear.\1",
}
def __init__(
self,
patch_size: int = 2,
patch_size_t: int = 1,
in_channels: int = 16,
out_channels: int = 16,
num_attention_heads: int = 24,
attention_head_dim: int = 128,
mlp_ratio: float = 4.0,
num_layers: int = 20,
num_single_layers: int = 40,
num_refiner_layers: int = 2,
rope_axes_dim: Tuple[int, int, int] = (16, 56, 56),
guidance_embeds: bool = False,
dtype: Optional[torch.dtype] = None,
text_embed_dim: int = 4096,
pooled_projection_dim: int = 768,
rope_theta: int = 256,
qk_norm: str = "rms_norm", #TODO(PY)
prefix="Hunyuan",
):
super().__init__()
hidden_size = attention_head_dim * num_attention_heads
self.patch_size = [patch_size_t, patch_size, patch_size]
self.in_channels = in_channels
self.out_channels = in_channels if out_channels is None else out_channels
self.patch_size = [
config.patch_size_t, config.patch_size, config.patch_size
]
self.in_channels = config.in_channels
self.num_channels_latents = config.num_channels_latents
self.out_channels = config.in_channels if config.out_channels is None else config.out_channels
self.unpatchify_channels = self.out_channels
self.guidance_embeds = guidance_embeds
self.rope_dim_list = list(rope_axes_dim)
self.rope_theta = rope_theta
self.text_states_dim = text_embed_dim
self.text_states_dim_2 = pooled_projection_dim
self.guidance_embeds = config.guidance_embeds
self.rope_dim_list = list(config.rope_axes_dim)
self.rope_theta = config.rope_theta
self.text_states_dim = config.text_embed_dim
self.text_states_dim_2 = config.pooled_projection_dim
# TODO(will): hack?
self.dtype = dtype
self.dtype = config.dtype
if hidden_size % num_attention_heads != 0:
pe_dim = config.hidden_size // config.num_attention_heads
if sum(config.rope_axes_dim) != pe_dim:
raise ValueError(
f"Hidden size {hidden_size} must be divisible by num_attention_heads {num_attention_heads}"
f"Got {config.rope_axes_dim} but expected positional dim {pe_dim}"
)
pe_dim = hidden_size // num_attention_heads
if sum(rope_axes_dim) != pe_dim:
raise ValueError(
f"Got {rope_axes_dim} but expected positional dim {pe_dim}")
self.hidden_size = hidden_size
self.num_attention_heads = num_attention_heads
self.hidden_size = config.hidden_size
self.num_attention_heads = config.num_attention_heads
self.num_channels_latents = config.num_channels_latents
# Image projection
self.img_in = PatchEmbed(self.patch_size,
self.in_channels,
self.hidden_size,
dtype=dtype,
prefix=f"{prefix}.img_in")
dtype=config.dtype,
prefix=f"{config.prefix}.img_in")
self.txt_in = SingleTokenRefiner(self.text_states_dim,
hidden_size,
num_attention_heads,
depth=num_refiner_layers,
dtype=dtype,
prefix=f"{prefix}.txt_in")
config.hidden_size,
config.num_attention_heads,
depth=config.num_refiner_layers,
dtype=config.dtype,
prefix=f"{config.prefix}.txt_in")
# Time modulation
self.time_in = TimestepEmbedder(self.hidden_size,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.time_in")
dtype=config.dtype,
prefix=f"{config.prefix}.time_in")
# Text modulation
self.vector_in = MLP(self.text_states_dim_2,
self.hidden_size,
self.hidden_size,
act_type="silu",
dtype=dtype,
prefix=f"{prefix}.vector_in")
dtype=config.dtype,
prefix=f"{config.prefix}.vector_in")
# Guidance modulation
self.guidance_in = (TimestepEmbedder(self.hidden_size,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.guidance_in")
self.guidance_in = (TimestepEmbedder(
self.hidden_size,
act_layer="silu",
dtype=config.dtype,
prefix=f"{config.prefix}.guidance_in")
if self.guidance_embeds else None)
# Double blocks
self.double_blocks = nn.ModuleList([
MMDoubleStreamBlock(
hidden_size,
num_attention_heads,
mlp_ratio=mlp_ratio,
dtype=dtype,
config.hidden_size,
config.num_attention_heads,
mlp_ratio=config.mlp_ratio,
dtype=config.dtype,
supported_attention_backends=self._supported_attention_backends,
prefix=f"{prefix}.double_blocks.{i}") for i in range(num_layers)
prefix=f"{config.prefix}.double_blocks.{i}")
for i in range(config.num_layers)
])
# Single blocks
self.single_blocks = nn.ModuleList([
MMSingleStreamBlock(
hidden_size,
num_attention_heads,
mlp_ratio=mlp_ratio,
dtype=dtype,
config.hidden_size,
config.num_attention_heads,
mlp_ratio=config.mlp_ratio,
dtype=config.dtype,
supported_attention_backends=self._supported_attention_backends,
prefix=f"{prefix}.single_blocks.{i+num_layers}")
for i in range(num_single_layers)
prefix=f"{config.prefix}.single_blocks.{i+config.num_layers}")
for i in range(config.num_single_layers)
])
self.final_layer = FinalLayer(hidden_size,
self.final_layer = FinalLayer(config.hidden_size,
self.patch_size,
self.out_channels,
dtype=dtype,
prefix=f"{prefix}.final_layer")
dtype=config.dtype,
prefix=f"{config.prefix}.final_layer")
self.__post_init__()
+31 -86
View File
@@ -7,6 +7,7 @@ import torch
import torch.nn as nn
from fastvideo.v1.attention import DistributedAttention, LocalAttention
from fastvideo.v1.configs.models.dits import WanVideoConfig
from fastvideo.v1.distributed.parallel_state import (
get_sequence_model_parallel_world_size)
from fastvideo.v1.layers.layernorm import (LayerNormScaleShift, RMSNorm,
@@ -350,115 +351,59 @@ class WanTransformerBlock(nn.Module):
class WanTransformer3DModel(BaseDiT):
_fsdp_shard_conditions = [
lambda n, m: "blocks" in n and str.isdigit(n.split(".")[-1]),
]
_supported_attention_backends = (_Backend.SLIDING_TILE_ATTN,
_Backend.SAGE_ATTN, _Backend.FLASH_ATTN,
_Backend.TORCH_SDPA)
_param_names_mapping = {
r"^patch_embedding\.(.*)$":
r"patch_embedding.proj.\1",
r"^condition_embedder\.text_embedder\.linear_1\.(.*)$":
r"condition_embedder.text_embedder.fc_in.\1",
r"^condition_embedder\.text_embedder\.linear_2\.(.*)$":
r"condition_embedder.text_embedder.fc_out.\1",
r"^condition_embedder\.time_embedder\.linear_1\.(.*)$":
r"condition_embedder.time_embedder.mlp.fc_in.\1",
r"^condition_embedder\.time_embedder\.linear_2\.(.*)$":
r"condition_embedder.time_embedder.mlp.fc_out.\1",
r"^condition_embedder\.time_proj\.(.*)$":
r"condition_embedder.time_modulation.linear.\1",
r"^condition_embedder\.image_embedder\.ff\.net\.0\.proj\.(.*)$":
r"condition_embedder.image_embedder.ff.fc_in.\1",
r"^condition_embedder\.image_embedder\.ff\.net\.2\.(.*)$":
r"condition_embedder.image_embedder.ff.fc_out.\1",
r"^blocks\.(\d+)\.attn1\.to_q\.(.*)$":
r"blocks.\1.to_q.\2",
r"^blocks\.(\d+)\.attn1\.to_k\.(.*)$":
r"blocks.\1.to_k.\2",
r"^blocks\.(\d+)\.attn1\.to_v\.(.*)$":
r"blocks.\1.to_v.\2",
r"^blocks\.(\d+)\.attn1\.to_out\.0\.(.*)$":
r"blocks.\1.to_out.\2",
r"^blocks\.(\d+)\.attn1\.norm_q\.(.*)$":
r"blocks.\1.norm_q.\2",
r"^blocks\.(\d+)\.attn1\.norm_k\.(.*)$":
r"blocks.\1.norm_k.\2",
r"^blocks\.(\d+)\.attn2\.to_out\.0\.(.*)$":
r"blocks.\1.attn2.to_out.\2",
r"^blocks\.(\d+)\.ffn\.net\.0\.proj\.(.*)$":
r"blocks.\1.ffn.fc_in.\2",
r"^blocks\.(\d+)\.ffn\.net\.2\.(.*)$":
r"blocks.\1.ffn.fc_out.\2",
r"blocks\.(\d+)\.norm2\.(.*)$":
r"blocks.\1.self_attn_residual_norm.norm.\2",
}
_fsdp_shard_conditions = WanVideoConfig()._fsdp_shard_conditions
_supported_attention_backends = WanVideoConfig(
)._supported_attention_backends
_param_names_mapping = WanVideoConfig()._param_names_mapping
def __init__(self,
patch_size: Tuple[int, int, int] = (1, 2, 2),
text_len=512,
num_attention_heads: int = 40,
attention_head_dim: int = 128,
in_channels: int = 16,
out_channels: int = 16,
text_dim: int = 4096,
freq_dim: int = 256,
ffn_dim: int = 13824,
num_layers: int = 40,
cross_attn_norm: bool = True,
qk_norm: str = "rms_norm_across_heads",
eps: float = 1e-6,
image_dim: Optional[int] = None,
added_kv_proj_dim: Optional[int] = None,
rope_max_seq_len: int = 1024,
prefix="Wan") -> None:
super().__init__()
def __init__(self, config: WanVideoConfig) -> None:
super().__init__(config=config)
inner_dim = num_attention_heads * attention_head_dim
self.hidden_size = inner_dim
self.num_attention_heads = num_attention_heads
self.in_channels = in_channels
self.out_channels = out_channels or in_channels
self.patch_size = patch_size
self.text_len = text_len
inner_dim = config.num_attention_heads * config.attention_head_dim
self.hidden_size = config.hidden_size
self.num_attention_heads = config.num_attention_heads
self.in_channels = config.in_channels
self.out_channels = config.out_channels
self.num_channels_latents = config.num_channels_latents
self.patch_size = config.patch_size
self.text_len = config.text_len
# 1. Patch & position embedding
self.patch_embedding = PatchEmbed(in_chans=in_channels,
self.patch_embedding = PatchEmbed(in_chans=config.in_channels,
embed_dim=inner_dim,
patch_size=patch_size,
patch_size=config.patch_size,
flatten=False)
# 2. Condition embeddings
self.condition_embedder = WanTimeTextImageEmbedding(
dim=inner_dim,
time_freq_dim=freq_dim,
text_embed_dim=text_dim,
image_embed_dim=image_dim,
time_freq_dim=config.freq_dim,
text_embed_dim=config.text_dim,
image_embed_dim=config.image_dim,
)
# 3. Transformer blocks
self.blocks = nn.ModuleList([
WanTransformerBlock(inner_dim,
ffn_dim,
num_attention_heads,
qk_norm,
cross_attn_norm,
eps,
added_kv_proj_dim,
config.ffn_dim,
config.num_attention_heads,
config.qk_norm,
config.cross_attn_norm,
config.eps,
config.added_kv_proj_dim,
self._supported_attention_backends,
prefix=f"{prefix}.blocks.{i}")
for i in range(num_layers)
prefix=f"{config.prefix}.blocks.{i}")
for i in range(config.num_layers)
])
# 4. Output norm & projection
self.norm_out = LayerNormScaleShift(inner_dim,
norm_type="layer",
eps=eps,
eps=config.eps,
elementwise_affine=False,
dtype=torch.float32)
self.proj_out = nn.Linear(inner_dim,
out_channels * math.prod(patch_size))
self.proj_out = nn.Linear(
inner_dim, config.out_channels * math.prod(config.patch_size))
self.scale_shift_table = nn.Parameter(
torch.randn(1, 2, inner_dim) / inner_dim**0.5)
+5 -3
View File
@@ -2,15 +2,17 @@ from typing import Tuple
from torch import nn
from fastvideo.v1.configs.models.encoders import EncoderConfig
from fastvideo.v1.platforms import _Backend
class BaseEncoder(nn.Module):
_supported_attention_backends: Tuple[_Backend,
...] = (_Backend.TORCH_SDPA, )
_supported_attention_backends: Tuple[
_Backend, ...] = EncoderConfig()._supported_attention_backends
def __init__(self, *args, **kwargs) -> None:
def __init__(self, config: EncoderConfig) -> None:
super().__init__()
self.config = config
if not self.supported_attention_backends:
raise ValueError(
f"Subclass {self.__class__.__name__} must define _supported_attention_backends"
+26 -81
View File
@@ -3,16 +3,17 @@
# Adapted from transformers: https://github.com/huggingface/transformers/blob/v4.39.0/src/transformers/models/clip/modeling_clip.py
"""Minimal implementation of CLIPVisionModel intended to be only used
within a vision language model."""
from typing import Iterable, Optional, Set, Tuple, Union, cast
from typing import Iterable, Optional, Set, Tuple, Union
import torch
import torch.nn as nn
from transformers import CLIPTextConfig, CLIPVisionConfig
from transformers.modeling_outputs import BaseModelOutputWithPooling
from vllm.model_executor.models.interfaces import SupportsQuant
# from transformers.modeling_attn_mask_utils import _create_4d_causal_attention_mask, _prepare_4d_attention_mask
from fastvideo.v1.attention import LocalAttention
from fastvideo.v1.configs.models.encoders import (CLIPTextConfig,
CLIPVisionConfig)
from fastvideo.v1.configs.quantization import QuantizationConfig
from fastvideo.v1.distributed import (divide,
get_tensor_model_parallel_world_size)
from fastvideo.v1.layers.activation import get_act_fn
@@ -20,45 +21,14 @@ from fastvideo.v1.layers.linear import (ColumnParallelLinear, QKVParallelLinear,
RowParallelLinear)
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.encoders.base import BaseEncoder
from fastvideo.v1.models.encoders.vision import (VisionEncoderInfo,
resolve_visual_encoder_outputs)
from fastvideo.v1.models.encoders.vision import resolve_visual_encoder_outputs
# TODO: support quantization
# from vllm.model_executor.layers.quantization import QuantizationConfig
from fastvideo.v1.models.loader.weight_utils import default_weight_loader
from fastvideo.v1.platforms import _Backend
logger = init_logger(__name__)
class QuantizationConfig:
pass
class CLIPEncoderInfo(VisionEncoderInfo[CLIPVisionConfig]):
def get_num_image_tokens(
self,
*,
image_width: int,
image_height: int,
) -> int:
return self.get_patch_grid_length()**2 + 1
def get_max_image_tokens(self) -> int:
return self.get_patch_grid_length()**2 + 1
def get_image_size(self) -> int:
return cast(int, self.vision_config.image_size)
def get_patch_size(self) -> int:
return cast(int, self.vision_config.patch_size)
def get_patch_grid_length(self) -> int:
image_size, patch_size = self.get_image_size(), self.get_patch_size()
assert image_size % patch_size == 0
return image_size // patch_size
# Adapted from https://github.com/huggingface/transformers/blob/v4.39.0/src/transformers/models/clip/modeling_clip.py#L164 # noqa
class CLIPVisionEmbeddings(nn.Module):
@@ -158,7 +128,7 @@ class CLIPAttention(nn.Module):
def __init__(
self,
config: CLIPVisionConfig,
config: Union[CLIPVisionConfig, CLIPTextConfig],
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
):
@@ -193,13 +163,13 @@ class CLIPAttention(nn.Module):
self.tp_size = get_tensor_model_parallel_world_size()
self.num_heads_per_partition = divide(self.num_heads, self.tp_size)
self.attn = LocalAttention(self.num_heads_per_partition,
self.head_dim,
self.num_heads_per_partition,
softmax_scale=self.scale,
causal=True,
supported_attention_backends=self.config.
supported_attention_backends)
self.attn = LocalAttention(
self.num_heads_per_partition,
self.head_dim,
self.num_heads_per_partition,
softmax_scale=self.scale,
causal=True,
supported_attention_backends=config._supported_attention_backends)
def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):
return tensor.view(bsz, seq_len, self.num_heads,
@@ -239,7 +209,7 @@ class CLIPMLP(nn.Module):
def __init__(
self,
config: CLIPVisionConfig,
config: Union[CLIPVisionConfig, CLIPTextConfig],
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
) -> None:
@@ -269,7 +239,7 @@ class CLIPEncoderLayer(nn.Module):
def __init__(
self,
config: CLIPTextConfig,
config: Union[CLIPTextConfig, CLIPVisionConfig],
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
) -> None:
@@ -314,7 +284,7 @@ class CLIPEncoder(nn.Module):
def __init__(
self,
config: CLIPVisionConfig,
config: Union[CLIPVisionConfig, CLIPTextConfig],
quant_config: Optional[QuantizationConfig] = None,
num_hidden_layers_override: Optional[int] = None,
prefix: str = "",
@@ -356,7 +326,6 @@ class CLIPTextTransformer(nn.Module):
def __init__(self,
config: CLIPTextConfig,
quant_config: Optional[QuantizationConfig] = None,
*,
num_hidden_layers_override: Optional[int] = None,
prefix: str = ""):
super().__init__()
@@ -377,9 +346,6 @@ class CLIPTextTransformer(nn.Module):
# For `pooled_output` computation
self.eos_token_id = config.eos_token_id
# For attention mask, it differs between `flash_attention_2` and other attention implementations
self._use_flash_attention_2 = config._attn_implementation == "flash_attention_2"
def forward(
self,
input_ids: Optional[torch.Tensor] = None,
@@ -393,7 +359,6 @@ class CLIPTextTransformer(nn.Module):
Returns:
"""
output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
output_hidden_states = (output_hidden_states
if output_hidden_states is not None else
self.config.output_hidden_states)
@@ -469,21 +434,15 @@ class CLIPTextTransformer(nn.Module):
class CLIPTextModel(BaseEncoder):
_supported_attention_backends = (_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
def __init__(
self,
config: CLIPTextConfig,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
) -> None:
super().__init__()
self.config = config
self.config.supported_attention_backends = self._supported_attention_backends
super().__init__(config)
self.text_model = CLIPTextTransformer(config=config,
quant_config=quant_config,
prefix=prefix)
quant_config=config.quant_config,
prefix=config.prefix)
def forward(
self,
@@ -492,9 +451,7 @@ class CLIPTextModel(BaseEncoder):
position_ids: Optional[torch.Tensor] = None,
output_attentions: Optional[bool] = None,
output_hidden_states: Optional[bool] = None,
return_dict: Optional[bool] = None,
) -> Union[Tuple, BaseModelOutputWithPooling]:
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
return self.text_model(
input_ids=input_ids,
@@ -548,7 +505,6 @@ class CLIPVisionTransformer(nn.Module):
self,
config: CLIPVisionConfig,
quant_config: Optional[QuantizationConfig] = None,
*,
num_hidden_layers_override: Optional[int] = None,
require_post_norm: Optional[bool] = None,
prefix: str = "",
@@ -615,30 +571,19 @@ class CLIPVisionTransformer(nn.Module):
return encoder_outputs
class CLIPVisionModel(BaseEncoder, SupportsQuant):
class CLIPVisionModel(BaseEncoder):
config_class = CLIPVisionConfig
main_input_name = "pixel_values"
packed_modules_mapping = {"qkv_proj": ["q_proj", "k_proj", "v_proj"]}
_supported_attention_backends = (_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
def __init__(
self,
config: CLIPVisionConfig,
quant_config: Optional[QuantizationConfig] = None,
*,
num_hidden_layers_override: Optional[int] = None,
require_post_norm: Optional[bool] = None,
prefix: str = "",
) -> None:
super().__init__()
self.config = config
self.config.supported_attention_backends = self._supported_attention_backends
def __init__(self, config: CLIPVisionConfig) -> None:
super().__init__(config)
self.vision_model = CLIPVisionTransformer(
config=config,
quant_config=quant_config,
num_hidden_layers_override=num_hidden_layers_override,
require_post_norm=require_post_norm,
prefix=f"{prefix}.vision_model")
quant_config=config.quant_config,
num_hidden_layers_override=config.num_hidden_layers_override,
require_post_norm=config.require_post_norm,
prefix=f"{config.prefix}.vision_model")
def forward(
self,
+31 -40
View File
@@ -23,15 +23,17 @@
# See the License for the specific language governing permissions and
# limitations under the License.
"""Inference-only LLaMA model compatible with HuggingFace weights."""
from typing import Any, Dict, Iterable, Optional, Set, Tuple, Type
from typing import Any, Dict, Iterable, Optional, Set, Tuple
import torch
from torch import nn
from transformers import LlamaConfig
from transformers.modeling_outputs import BaseModelOutputWithPast
# from vllm.model_executor.layers.quantization import QuantizationConfig
from fastvideo.v1.attention import LocalAttention
# from ..utils import (extract_layer_index)
from fastvideo.v1.configs.models.encoders import LlamaConfig
from fastvideo.v1.configs.quantization import QuantizationConfig
from fastvideo.v1.distributed import get_tensor_model_parallel_world_size
from fastvideo.v1.layers.activation import SiluAndMul
from fastvideo.v1.layers.layernorm import RMSNorm
@@ -42,12 +44,6 @@ from fastvideo.v1.layers.vocab_parallel_embedding import VocabParallelEmbedding
from fastvideo.v1.models.encoders.base import BaseEncoder
from fastvideo.v1.models.loader.weight_utils import (default_weight_loader,
maybe_remap_kv_scale_name)
# from ..utils import (extract_layer_index)
from fastvideo.v1.platforms import _Backend
class QuantizationConfig:
pass
class LlamaMLP(nn.Module):
@@ -171,7 +167,7 @@ class LlamaAttention(nn.Module):
self.num_kv_heads,
softmax_scale=self.scaling,
causal=True,
supported_attention_backends=config.supported_attention_backends)
supported_attention_backends=config._supported_attention_backends)
def forward(
self,
@@ -280,27 +276,22 @@ class LlamaDecoderLayer(nn.Module):
class LlamaModel(BaseEncoder):
_supported_attention_backends = (_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
def __init__(self,
config: LlamaConfig,
prefix: str = "",
layer_type: Type[LlamaDecoderLayer] = LlamaDecoderLayer):
super().__init__()
quant_config = None
lora_config = None
def __init__(
self,
config: LlamaConfig,
):
super().__init__(config)
self.config = config
self.config.supported_attention_backends = self._supported_attention_backends
self.quant_config = quant_config
if lora_config is not None:
self.quant_config = self.config.quant_config
if config.lora_config is not None:
max_loras = 1
lora_vocab_size = 1
if hasattr(lora_config, "max_loras"):
max_loras = lora_config.max_loras
if hasattr(lora_config, "lora_extra_vocab_size"):
lora_vocab_size = lora_config.lora_extra_vocab_size
if hasattr(config.lora_config, "max_loras"):
max_loras = config.lora_config.max_loras
if hasattr(config.lora_config, "lora_extra_vocab_size"):
lora_vocab_size = config.lora_config.lora_extra_vocab_size
lora_vocab = lora_vocab_size * max_loras
else:
lora_vocab = 0
@@ -311,13 +302,13 @@ class LlamaModel(BaseEncoder):
self.vocab_size,
config.hidden_size,
org_num_embeddings=config.vocab_size,
quant_config=quant_config,
quant_config=config.quant_config,
)
self.layers = nn.ModuleList([
layer_type(config=config,
quant_config=quant_config,
prefix=f"{prefix}.layers.{i}")
LlamaDecoderLayer(config=config,
quant_config=config.quant_config,
prefix=f"{config.prefix}.layers.{i}")
for i in range(config.num_hidden_layers)
])
@@ -395,17 +386,17 @@ class LlamaModel(BaseEncoder):
# Models trained using ColossalAI may include these tensors in
# the checkpoint. Skip them.
continue
if (self.quant_config is not None and
(scale_name := self.quant_config.get_cache_scale(name))):
# Loading kv cache quantization scales
param = params_dict[scale_name]
weight_loader = getattr(param, "weight_loader",
default_weight_loader)
loaded_weight = (loaded_weight if loaded_weight.dim() == 0 else
loaded_weight[0])
weight_loader(param, loaded_weight)
loaded_params.add(scale_name)
continue
# if (self.quant_config is not None and
# (scale_name := self.quant_config.get_cache_scale(name))):
# # Loading kv cache quantization scales
# param = params_dict[scale_name]
# weight_loader = getattr(param, "weight_loader",
# default_weight_loader)
# loaded_weight = (loaded_weight if loaded_weight.dim() == 0 else
# loaded_weight[0])
# weight_loader(param, loaded_weight)
# loaded_params.add(scale_name)
# continue
if "scale" in name:
# Remapping the name of FP8 kv-scale.
kv_scale_name: Optional[str] = maybe_remap_kv_scale_name(
+7 -9
View File
@@ -26,8 +26,9 @@ from typing import Iterable, Optional, Set, Tuple
import torch
import torch.nn.functional as F
from torch import nn
from transformers import T5Config
from fastvideo.v1.configs.models.encoders import T5Config
from fastvideo.v1.configs.quantization import QuantizationConfig
from fastvideo.v1.distributed import (get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size)
from fastvideo.v1.layers.activation import get_act_fn
@@ -35,13 +36,10 @@ from fastvideo.v1.layers.layernorm import RMSNorm
from fastvideo.v1.layers.linear import (MergedColumnParallelLinear,
QKVParallelLinear, RowParallelLinear)
from fastvideo.v1.layers.vocab_parallel_embedding import VocabParallelEmbedding
from fastvideo.v1.models.encoders.base import BaseEncoder
from fastvideo.v1.models.loader.weight_utils import default_weight_loader
class QuantizationConfig:
pass
class AttentionType:
"""
Attention type.
@@ -501,10 +499,10 @@ class T5Stack(nn.Module):
return hidden_states
class T5EncoderModel(nn.Module):
class T5EncoderModel(BaseEncoder):
def __init__(self, config: T5Config, prefix: str = ""):
super().__init__()
super().__init__(config)
quant_config = None
@@ -589,10 +587,10 @@ class T5EncoderModel(nn.Module):
return loaded_params
class UMT5EncoderModel(nn.Module):
class UMT5EncoderModel(BaseEncoder):
def __init__(self, config: T5Config, prefix: str = ""):
super().__init__()
super().__init__(config)
quant_config = None
+56 -24
View File
@@ -2,6 +2,7 @@
import dataclasses
import glob
import json
import os
import time
from abc import ABC, abstractmethod
@@ -10,13 +11,12 @@ from typing import Any, Generator, Iterable, List, Optional, Tuple, cast
import torch
import torch.nn as nn
from safetensors.torch import load_file as safetensors_load_file
from transformers import AutoImageProcessor, AutoTokenizer, PretrainedConfig
from transformers import AutoImageProcessor, AutoTokenizer
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.hf_transformer_utils import (get_diffusers_config,
get_hf_config)
from fastvideo.v1.models.hf_transformer_utils import get_diffusers_config
from fastvideo.v1.models.loader.fsdp_load import load_fsdp_model
from fastvideo.v1.models.loader.utils import set_default_torch_dtype
from fastvideo.v1.models.loader.weight_utils import (
@@ -201,18 +201,34 @@ class TextEncoderLoader(ComponentLoader):
def load(self, model_path: str, architecture: str,
fastvideo_args: FastVideoArgs):
"""Load the text encoders based on the model path, architecture, and inference args."""
model_config: PretrainedConfig = get_hf_config(
model=model_path,
trust_remote_code=fastvideo_args.trust_remote_code,
revision=fastvideo_args.revision,
model_override_args=None,
)
# model_config: PretrainedConfig = get_hf_config(
# model=model_path,
# trust_remote_code=fastvideo_args.trust_remote_code,
# revision=fastvideo_args.revision,
# model_override_args=None,
# )
with open(os.path.join(model_path, "config.json")) as f:
model_config = json.load(f)
model_config.pop("_name_or_path", None)
model_config.pop("transformers_version", None)
model_config.pop("model_type", None)
model_config.pop("tokenizer_class", None)
model_config.pop("torch_dtype", None)
logger.info("HF Model config: %s", model_config)
try:
encoder_config = fastvideo_args.text_encoder_config
encoder_config.update_model_arch(model_config)
encoder_precision = fastvideo_args.text_encoder_precision
except Exception:
encoder_config = fastvideo_args.text_encoder_config_2
encoder_config.update_model_arch(model_config)
encoder_precision = fastvideo_args.text_encoder_precision_2
target_device = torch.device(fastvideo_args.device_str)
# TODO(will): add support for other dtypes
return self.load_model(model_path, model_config, target_device,
fastvideo_args.text_encoder_precision)
return self.load_model(model_path, encoder_config, target_device,
encoder_precision)
def load_model(self,
model_path: str,
@@ -251,17 +267,26 @@ class ImageEncoderLoader(TextEncoderLoader):
def load(self, model_path: str, architecture: str,
fastvideo_args: FastVideoArgs):
"""Load the text encoders based on the model path, architecture, and inference args."""
model_config: PretrainedConfig = get_hf_config(
model=model_path,
trust_remote_code=fastvideo_args.trust_remote_code,
revision=fastvideo_args.revision,
model_override_args=None,
)
# model_config: PretrainedConfig = get_hf_config(
# model=model_path,
# trust_remote_code=fastvideo_args.trust_remote_code,
# revision=fastvideo_args.revision,
# model_override_args=None,
# )
with open(os.path.join(model_path, "config.json")) as f:
model_config = json.load(f)
model_config.pop("_name_or_path", None)
model_config.pop("transformers_version", None)
model_config.pop("torch_dtype", None)
model_config.pop("model_type", None)
logger.info("HF Model config: %s", model_config)
encoder_config = fastvideo_args.image_encoder_config
encoder_config.update_model_arch(model_config)
target_device = torch.device(fastvideo_args.device_str)
# TODO(will): add support for other dtypes
return self.load_model(model_path, model_config, target_device,
return self.load_model(model_path, encoder_config, target_device,
fastvideo_args.image_encoder_precision)
@@ -311,9 +336,11 @@ class VAELoader(ComponentLoader):
assert class_name is not None, "Model config does not contain a _class_name attribute. Only diffusers format is supported."
config.pop("_diffusers_version")
vae_config = fastvideo_args.vae_config
vae_config.update_model_arch(config)
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
vae = vae_cls(**config).to(fastvideo_args.device)
vae = vae_cls(vae_config).to(fastvideo_args.device)
# Find all safetensors files
safetensors_list = glob.glob(
@@ -323,7 +350,8 @@ class VAELoader(ComponentLoader):
safetensors_list
) == 1, f"Found {len(safetensors_list)} safetensors files in {model_path}"
loaded = safetensors_load_file(safetensors_list[0])
vae.load_state_dict(loaded)
vae.load_state_dict(
loaded, strict=False) # We might only load encoder or decoder
dtype = PRECISION_TO_TYPE[fastvideo_args.vae_precision]
vae = vae.eval().to(dtype)
@@ -336,13 +364,17 @@ class TransformerLoader(ComponentLoader):
def load(self, model_path: str, architecture: str,
fastvideo_args: FastVideoArgs):
"""Load the transformer based on the model path, architecture, and inference args."""
model_config = get_diffusers_config(model=model_path)
cls_name = model_config.pop("_class_name")
config = get_diffusers_config(model=model_path)
cls_name = config.pop("_class_name")
if cls_name is None:
raise ValueError(
"Model config does not contain a _class_name attribute. "
"Only diffusers format is supported.")
model_config.pop("_diffusers_version")
config.pop("_diffusers_version")
# Config from Diffusers supersedes fastvideo's model config
dit_config = fastvideo_args.dit_config
dit_config.update_model_arch(config)
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
@@ -361,7 +393,7 @@ class TransformerLoader(ComponentLoader):
# Load the model using FSDP loader
logger.info("Loading model from %s", cls_name)
model = load_fsdp_model(model_cls=model_cls,
init_params=model_config,
init_params={"config": dit_config},
weight_dir_list=safetensors_list,
device=fastvideo_args.device,
cpu_offload=fastvideo_args.use_cpu_offload,
+37 -20
View File
@@ -2,13 +2,14 @@
from abc import ABC, abstractmethod
from math import prod
from typing import Iterator, Optional, Tuple, Union
from typing import Iterator, Optional, Tuple, Union, cast
import numpy as np
import torch
import torch.distributed as dist
from diffusers.utils.torch_utils import randn_tensor
from fastvideo.v1.configs.models import VAEConfig
from fastvideo.v1.distributed import (get_sequence_model_parallel_rank,
get_sequence_model_parallel_world_size)
@@ -20,29 +21,35 @@ class ParallelTiledVAE(ABC):
tile_sample_stride_height: int
tile_sample_stride_width: int
tile_sample_stride_num_frames: int
blend_num_frames: int
use_tiling: bool
use_temporal_tiling: bool
use_parallel_tiling: bool
temporal_compression_ratio: int
spatial_compression_ratio: int
scaling_factor: Union[float, torch.tensor]
def __init__(self, *args, **kwargs) -> None:
# Check if subclass has defined all required properties
required_attributes = [
'tile_sample_min_height', 'tile_sample_min_width',
'tile_sample_min_num_frames', 'tile_sample_stride_height',
'tile_sample_stride_width', 'tile_sample_stride_num_frames',
'spatial_compression_ratio', 'temporal_compression_ratio',
'use_tiling', 'use_temporal_tiling', 'use_parallel_tiling',
'scaling_factor'
]
def __init__(self, config: VAEConfig, **kwargs) -> None:
self.config = config
self.tile_sample_min_height = config.tile_sample_min_height
self.tile_sample_min_width = config.tile_sample_min_width
self.tile_sample_min_num_frames = config.tile_sample_min_num_frames
self.tile_sample_stride_height = config.tile_sample_stride_height
self.tile_sample_stride_width = config.tile_sample_stride_width
self.tile_sample_stride_num_frames = config.tile_sample_stride_num_frames
self.blend_num_frames = config.blend_num_frames
self.use_tiling = config.use_tiling
self.use_temporal_tiling = config.use_temporal_tiling
self.use_parallel_tiling = config.use_parallel_tiling
for attr in required_attributes:
if not hasattr(self, attr):
raise AttributeError(
f"Subclasses of ParallelVAE must define '{attr}' property")
self.blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
@property
def temporal_compression_ratio(self) -> int:
return cast(int, self.config.temporal_compression_ratio)
@property
def spatial_compression_ratio(self) -> int:
return cast(int, self.config.spatial_compression_ratio)
@property
def scaling_factor(self) -> Union[float, torch.tensor]:
return cast(Union[float, torch.tensor], self.config.scaling_factor)
@abstractmethod
def _encode(self, *args, **kwargs) -> torch.Tensor:
@@ -408,6 +415,10 @@ class ParallelTiledVAE(ABC):
tile_sample_stride_height: Optional[int] = None,
tile_sample_stride_width: Optional[int] = None,
tile_sample_stride_num_frames: Optional[int] = None,
blend_num_frames: Optional[int] = None,
use_tiling: Optional[bool] = None,
use_temporal_tiling: Optional[bool] = None,
use_parallel_tiling: Optional[bool] = None,
) -> None:
r"""
Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to
@@ -439,7 +450,13 @@ class ParallelTiledVAE(ABC):
self.tile_sample_stride_height = tile_sample_stride_height or self.tile_sample_stride_height
self.tile_sample_stride_width = tile_sample_stride_width or self.tile_sample_stride_width
self.tile_sample_stride_num_frames = tile_sample_stride_num_frames or self.tile_sample_stride_num_frames
self.blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
if blend_num_frames is not None:
self.blend_num_frames = blend_num_frames
else:
self.blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
self.use_tiling = use_tiling or self.use_tiling
self.use_temporal_tiling = use_temporal_tiling or self.use_temporal_tiling
self.use_parallel_tiling = use_parallel_tiling or self.use_parallel_tiling
def disable_tiling(self) -> None:
r"""
+31 -75
View File
@@ -21,10 +21,9 @@ import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.utils.checkpoint
from fastvideo.v1.configs.models.vaes import HunyuanVAEConfig
from fastvideo.v1.layers.activation import get_act_fn
from fastvideo.v1.models.utils import auto_attributes
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
@@ -773,95 +772,52 @@ class AutoencoderKLHunyuanVideo(nn.Module, ParallelTiledVAE):
_supports_gradient_checkpointing = True
@auto_attributes
def __init__(
self,
in_channels: int = 3,
out_channels: int = 3,
latent_channels: int = 16,
down_block_types: Tuple[str, ...] = (
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
),
up_block_types: Tuple[str, ...] = (
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
),
block_out_channels: Tuple[int, ...] = (128, 256, 512, 512),
layers_per_block: int = 2,
act_fn: str = "silu",
norm_num_groups: int = 32,
scaling_factor: float = 0.476986,
spatial_compression_ratio: int = 8,
temporal_compression_ratio: int = 4,
mid_block_add_attention: bool = True,
load_encoder: bool = True,
load_decoder: bool = True,
config: HunyuanVAEConfig,
) -> None:
super().__init__()
nn.Module.__init__(self)
ParallelTiledVAE.__init__(self, config)
# TODO(will): only pass in config. We do this by manually defining a
# config for hunyuan vae
self.block_out_channels = block_out_channels
self.block_out_channels = config.block_out_channels
if load_encoder:
if config.load_encoder:
self.encoder = HunyuanVideoEncoder3D(
in_channels=in_channels,
out_channels=latent_channels,
down_block_types=down_block_types,
block_out_channels=block_out_channels,
layers_per_block=layers_per_block,
norm_num_groups=norm_num_groups,
act_fn=act_fn,
in_channels=config.in_channels,
out_channels=config.latent_channels,
down_block_types=config.down_block_types,
block_out_channels=config.block_out_channels,
layers_per_block=config.layers_per_block,
norm_num_groups=config.norm_num_groups,
act_fn=config.act_fn,
double_z=True,
mid_block_add_attention=mid_block_add_attention,
temporal_compression_ratio=temporal_compression_ratio,
spatial_compression_ratio=spatial_compression_ratio,
mid_block_add_attention=config.mid_block_add_attention,
temporal_compression_ratio=config.temporal_compression_ratio,
spatial_compression_ratio=config.spatial_compression_ratio,
)
self.quant_conv = nn.Conv3d(2 * latent_channels,
2 * latent_channels,
self.quant_conv = nn.Conv3d(2 * config.latent_channels,
2 * config.latent_channels,
kernel_size=1)
if load_decoder:
if config.load_decoder:
self.decoder = HunyuanVideoDecoder3D(
in_channels=latent_channels,
out_channels=out_channels,
up_block_types=up_block_types,
block_out_channels=block_out_channels,
layers_per_block=layers_per_block,
norm_num_groups=norm_num_groups,
act_fn=act_fn,
time_compression_ratio=temporal_compression_ratio,
spatial_compression_ratio=spatial_compression_ratio,
mid_block_add_attention=mid_block_add_attention,
in_channels=config.latent_channels,
out_channels=config.out_channels,
up_block_types=config.up_block_types,
block_out_channels=config.block_out_channels,
layers_per_block=config.layers_per_block,
norm_num_groups=config.norm_num_groups,
act_fn=config.act_fn,
time_compression_ratio=config.temporal_compression_ratio,
spatial_compression_ratio=config.spatial_compression_ratio,
mid_block_add_attention=config.mid_block_add_attention,
)
self.post_quant_conv = nn.Conv3d(latent_channels,
latent_channels,
self.post_quant_conv = nn.Conv3d(config.latent_channels,
config.latent_channels,
kernel_size=1)
# When decoding spatially large video latents, the memory requirement is very high. By breaking the video latent
# frames spatially into smaller tiles and performing multiple forward passes for decoding, and then blending the
# intermediate tiles together, the memory requirement can be lowered.
self.use_tiling = True
self.use_temporal_tiling = True
self.use_parallel_tiling = True
self.scaling_factor = scaling_factor
# The minimal tile height and width for spatial tiling to be used
self.tile_sample_min_height = 256
self.tile_sample_min_width = 256
self.tile_sample_min_num_frames = 16
# The minimal distance between two spatial tiles
self.tile_sample_stride_height = 192
self.tile_sample_stride_width = 192
self.tile_sample_stride_num_frames = 12
ParallelTiledVAE.__init__(self)
def _encode(self, x: torch.Tensor) -> torch.Tensor:
x = self.encoder(x)
enc = self.quant_conv(x)
+35 -93
View File
@@ -22,8 +22,8 @@ import torch
import torch.nn as nn
import torch.nn.functional as F
from fastvideo.v1.configs.models.vaes import WanVAEConfig
from fastvideo.v1.layers.activation import get_act_fn
from fastvideo.v1.models.utils import auto_attributes
from fastvideo.v1.models.vaes.common import (DiagonalGaussianDistribution,
ParallelTiledVAE)
@@ -781,96 +781,36 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
_supports_gradient_checkpointing = False
@auto_attributes
def __init__(self,
base_dim: int = 96,
z_dim: int = 16,
dim_mult: Tuple[int, ...] = (1, 2, 4, 4),
num_res_blocks: int = 2,
attn_scales: Tuple[float, ...] = (),
temperal_downsample: Tuple[bool, ...] = (False, True, True),
dropout: float = 0.0,
latents_mean: Tuple[float, ...] = (
-0.7571,
-0.7089,
-0.9113,
0.1075,
-0.1745,
0.9653,
-0.1517,
1.5508,
0.4134,
-0.0715,
0.5517,
-0.3632,
-0.1922,
-0.9497,
0.2503,
-0.2921,
),
latents_std: Tuple[float, ...] = (
2.8184,
1.4541,
2.3275,
2.6558,
1.2196,
1.7708,
2.6052,
2.0743,
3.2687,
2.1526,
2.8652,
1.5579,
1.6382,
1.1253,
2.8251,
1.9160,
),
load_encoder: bool = True,
load_decoder: bool = True) -> None:
super().__init__()
def __init__(
self,
config: WanVAEConfig,
) -> None:
nn.Module.__init__(self)
ParallelTiledVAE.__init__(self, config)
self.z_dim = z_dim
self.temperal_downsample = list(temperal_downsample)
self.temperal_upsample = list(temperal_downsample)[::-1]
self.latents_mean = list(latents_mean)
self.latents_std = list(latents_std)
self.scaling_factor = 1.0 / torch.tensor(self.config.latents_std).view(
1, self.config.z_dim, 1, 1, 1)
self.shift_factor = torch.tensor(self.config.latents_mean).view(
1, self.config.z_dim, 1, 1, 1)
self.z_dim = config.z_dim
self.temperal_downsample = list(config.temperal_downsample)
self.temperal_upsample = list(config.temperal_downsample)[::-1]
self.latents_mean = list(config.latents_mean)
self.latents_std = list(config.latents_std)
self.shift_factor = config.shift_factor
if load_encoder:
self.encoder = WanEncoder3d(base_dim, z_dim * 2, dim_mult,
num_res_blocks, attn_scales,
self.temperal_downsample, dropout)
self.quant_conv = WanCausalConv3d(z_dim * 2, z_dim * 2, 1)
self.post_quant_conv = WanCausalConv3d(z_dim, z_dim, 1)
if config.load_encoder:
self.encoder = WanEncoder3d(config.base_dim, self.z_dim * 2,
config.dim_mult, config.num_res_blocks,
config.attn_scales,
self.temperal_downsample,
config.dropout)
self.quant_conv = WanCausalConv3d(self.z_dim * 2, self.z_dim * 2, 1)
self.post_quant_conv = WanCausalConv3d(self.z_dim, self.z_dim, 1)
if load_decoder:
self.decoder = WanDecoder3d(base_dim, z_dim, dim_mult,
num_res_blocks, attn_scales,
self.temperal_upsample, dropout)
if config.load_decoder:
self.decoder = WanDecoder3d(config.base_dim, self.z_dim,
config.dim_mult, config.num_res_blocks,
config.attn_scales,
self.temperal_upsample, config.dropout)
self.use_tiling = True
self.use_temporal_tiling = False
self.use_parallel_tiling = False
self.spatial_compression_ratio = 8
self.temporal_compression_ratio = 4
# The minimal tile height and width for spatial tiling to be used
self.tile_sample_min_height = 256
self.tile_sample_min_width = 256
self.tile_sample_min_num_frames = 16
# The minimal distance between two spatial tiles
self.tile_sample_stride_height = 192
self.tile_sample_stride_width = 192
self.tile_sample_stride_num_frames = 12
# Whether to use the feature cache algorithm used by diffusers and Wan2.1
self.use_feature_cache = True # default to True for best performance
ParallelTiledVAE.__init__(self)
self.use_feature_cache = config.use_feature_cache
def clear_cache(self) -> None:
@@ -881,13 +821,15 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
count += 1
return count
self._conv_num = _count_conv3d(self.decoder)
self._conv_idx = 0
self._feat_map = [None] * self._conv_num
if self.config.load_decoder:
self._conv_num = _count_conv3d(self.decoder)
self._conv_idx = 0
self._feat_map = [None] * self._conv_num
# cache encode
self._enc_conv_num = _count_conv3d(self.encoder)
self._enc_conv_idx = 0
self._enc_feat_map = [None] * self._enc_conv_num
if self.config.load_encoder:
self._enc_conv_num = _count_conv3d(self.encoder)
self._enc_conv_idx = 0
self._enc_feat_map = [None] * self._enc_conv_num
def encode(self, x: torch.Tensor) -> torch.Tensor:
if self.use_feature_cache:
@@ -114,12 +114,11 @@ class ComposedPipelineBase(ABC):
"""
raise NotImplementedError
@abstractmethod
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""
Initialize the pipeline.
"""
raise NotImplementedError
return
def load_modules(self, fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
"""
@@ -6,8 +6,6 @@ This module contains an implementation of the Hunyuan video diffusion pipeline
using the modular pipeline architecture.
"""
from diffusers.image_processor import VaeImageProcessor
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
@@ -57,7 +55,8 @@ class HunyuanVideoPipeline(ComposedPipelineBase):
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler")))
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
@@ -67,20 +66,5 @@ class HunyuanVideoPipeline(ComposedPipelineBase):
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""
Initialize the pipeline.
"""
vae_scale_factor = 2**(len(self.get_module("vae").block_out_channels) -
1)
fastvideo_args.vae_scale_factor = vae_scale_factor
self.image_processor = VaeImageProcessor(
vae_scale_factor=vae_scale_factor)
self.add_module("image_processor", self.image_processor)
num_channels_latents = self.get_module("transformer").in_channels
fastvideo_args.num_channels_latents = num_channels_latents
EntryClass = HunyuanVideoPipeline
@@ -36,6 +36,8 @@ class ForwardBatch:
# Text inputs
prompt: Optional[Union[str, List[str]]] = None
negative_prompt: Optional[Union[str, List[str]]] = None
prompt_path: Optional[str] = None
output_path: str = "outputs/"
# Primary encoder embeddings
prompt_embeds: List[torch.Tensor] = field(default_factory=list)
@@ -49,6 +51,7 @@ class ForwardBatch:
# Batch info
batch_size: Optional[int] = None
num_videos_per_prompt: int = 1
seed: Optional[int] = None
seeds: Optional[List[int]] = None
# Tracking if embeddings are already processed
@@ -60,7 +63,6 @@ class ForwardBatch:
image_latent: Optional[torch.Tensor] = None
# Latent dimensions
num_channels_latents: Optional[int] = None
height_latents: Optional[int] = None
width_latents: Optional[int] = None
num_frames: int = 1 # Default for image models
@@ -68,6 +70,7 @@ class ForwardBatch:
# Original dimensions (before VAE scaling)
height: Optional[int] = None
width: Optional[int] = None
fps: Optional[int] = None
# Timesteps
timesteps: Optional[torch.Tensor] = None
@@ -95,7 +98,9 @@ class ForwardBatch:
# Extra parameters that might be needed by specific pipeline implementations
extra: Dict[str, Any] = field(default_factory=dict)
device: torch.device = field(default_factory=lambda: torch.device("cuda"))
# Misc
save_video: bool = True
return_frames: bool = False
def __post_init__(self):
"""Initialize dependent fields after dataclass initialization."""
@@ -53,12 +53,12 @@ class CLIPImageEncodingStage(PipelineStage):
The batch with encoded prompt embeddings.
"""
if fastvideo_args.use_cpu_offload:
self.image_encoder = self.image_encoder.to(batch.device)
self.image_encoder = self.image_encoder.to(fastvideo_args.device)
image = load_image(batch.image_path)
image_inputs = self.image_processor(
images=image, return_tensors="pt").to(batch.device)
images=image, return_tensors="pt").to(fastvideo_args.device)
with set_forward_context(current_timestep=0, attn_metadata=None):
image_embeds = self.image_encoder(**image_inputs)
@@ -52,7 +52,7 @@ class CLIPTextEncodingStage(PipelineStage):
The batch with encoded prompt embeddings.
"""
if fastvideo_args.use_cpu_offload:
self.text_encoder = self.text_encoder.to(batch.device)
self.text_encoder = self.text_encoder.to(fastvideo_args.device)
text_inputs = self.tokenizer(
batch.prompt,
@@ -63,7 +63,7 @@ class CLIPTextEncodingStage(PipelineStage):
)
with set_forward_context(current_timestep=0, attn_metadata=None):
outputs = self.text_encoder(input_ids=text_inputs["input_ids"].to(
batch.device), )
fastvideo_args.device), )
prompt_embeds = outputs["pooler_output"]
batch.prompt_embeds.append(prompt_embeds)
@@ -79,7 +79,7 @@ class CLIPTextEncodingStage(PipelineStage):
with set_forward_context(current_timestep=0, attn_metadata=None):
negative_outputs = self.text_encoder(
input_ids=negative_text_inputs["input_ids"].to(
batch.device), )
fastvideo_args.device), )
negative_prompt_embeds = negative_outputs["pooler_output"]
assert batch.negative_prompt_embeds is not None
+3 -2
View File
@@ -7,6 +7,7 @@ import torch
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.utils import PRECISION_TO_TYPE
@@ -22,8 +23,8 @@ class DecodingStage(PipelineStage):
output format (e.g., pixel values).
"""
def __init__(self, vae) -> None:
self.vae = vae
def __init__(self, vae: ParallelTiledVAE) -> None:
self.vae: ParallelTiledVAE = vae
def forward(
self,
+2 -2
View File
@@ -66,7 +66,7 @@ class DenoisingStage(PipelineStage):
"""
# If use cpu offload, need to load the model back into gpu again
if fastvideo_args.use_cpu_offload:
self.transformer = self.transformer.to(batch.device)
self.transformer = self.transformer.to(fastvideo_args.device)
# Prepare extra step kwargs for scheduler
extra_step_kwargs = self.prepare_extra_func_kwargs(
self.scheduler.step,
@@ -165,7 +165,7 @@ class DenoisingStage(PipelineStage):
[fastvideo_args.embedded_cfg_scale] *
latent_model_input.shape[0],
dtype=torch.float32,
device=batch.device,
device=fastvideo_args.device,
).to(target_dtype) * 1000.0 if fastvideo_args.embedded_cfg_scale
is not None else None)
+11 -9
View File
@@ -9,6 +9,7 @@ import torch
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
from fastvideo.v1.models.vision_utils import (get_default_height_width,
load_image, normalize,
numpy_to_pt, pil_to_numpy, resize)
@@ -27,8 +28,8 @@ class EncodingStage(PipelineStage):
input format (e.g., latents).
"""
def __init__(self, vae) -> None:
self.vae = vae
def __init__(self, vae: ParallelTiledVAE) -> None:
self.vae: ParallelTiledVAE = vae
def forward(
self,
@@ -49,6 +50,8 @@ class EncodingStage(PipelineStage):
# TODO(will): remove this once we add input/output validation for stages
if image_path is None:
raise ValueError("Image Path must be provided")
assert batch.height is not None
assert batch.width is not None
latent_height = batch.height // self.vae.spatial_compression_ratio
latent_width = batch.width // self.vae.spatial_compression_ratio
@@ -57,16 +60,15 @@ class EncodingStage(PipelineStage):
image,
vae_scale_factor=self.vae.spatial_compression_ratio,
height=batch.height,
width=batch.width).to(batch.device, dtype=torch.float32)
width=batch.width).to(fastvideo_args.device, dtype=torch.float32)
image = image.unsqueeze(2)
video_condition = torch.cat([
image,
image.new_zeros(image.shape[0], image.shape[1],
fastvideo_args.num_frames - 1, batch.height,
batch.width)
batch.num_frames - 1, batch.height, batch.width)
],
dim=2)
video_condition = video_condition.to(device=batch.device,
video_condition = video_condition.to(device=fastvideo_args.device,
dtype=torch.float32)
# Setup VAE precision
@@ -106,9 +108,9 @@ class EncodingStage(PipelineStage):
else:
latent_condition = latent_condition * self.vae.scaling_factor
mask_lat_size = torch.ones(1, 1, fastvideo_args.num_frames,
latent_height, latent_width)
mask_lat_size[:, :, list(range(1, fastvideo_args.num_frames))] = 0
mask_lat_size = torch.ones(1, 1, batch.num_frames, latent_height,
latent_width)
mask_lat_size[:, :, list(range(1, batch.num_frames))] = 0
first_frame_mask = mask_lat_size[:, :, 0:1]
first_frame_mask = torch.repeat_interleave(
first_frame_mask,
@@ -24,9 +24,10 @@ class InputValidationStage(PipelineStage):
def _generate_seeds(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs):
"""Generate seeds for the inference"""
seed = fastvideo_args.seed
num_videos_per_prompt = fastvideo_args.num_videos
seed = batch.seed
num_videos_per_prompt = batch.num_videos_per_prompt
assert seed is not None
seeds = [seed + i for i in range(num_videos_per_prompt)]
batch.seeds = seeds
# Peiyuan: using GPU seed will cause A100 and H100 to generate different results...
@@ -85,12 +86,4 @@ class InputValidationStage(PipelineStage):
f"Guidance scale must be positive, but got {batch.guidance_scale}"
)
# Set device if not already set
if batch.device is None:
batch.device = self.device
# Set data type if not already set
if batch.data_type is None:
batch.data_type = fastvideo_args.precision
return batch
@@ -6,7 +6,6 @@ from diffusers.utils.torch_utils import randn_tensor
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
@@ -21,10 +20,10 @@ class LatentPreparationStage(PipelineStage):
denoised during the diffusion process.
"""
def __init__(self, scheduler, vae=None) -> None:
def __init__(self, scheduler, transformer) -> None:
super().__init__()
self.scheduler = scheduler
self.vae = vae
self.transformer = transformer
def forward(
self,
@@ -42,9 +41,10 @@ class LatentPreparationStage(PipelineStage):
The batch with prepared latent variables.
"""
latent_num_frames = None
# Adjust video length based on VAE version if needed
if hasattr(self, 'adjust_video_length'):
batch = self.adjust_video_length(self.vae, batch, fastvideo_args)
latent_num_frames = self.adjust_video_length(batch, fastvideo_args)
# Determine batch size
if isinstance(batch.prompt, list):
batch_size = len(batch.prompt)
@@ -58,10 +58,10 @@ class LatentPreparationStage(PipelineStage):
# Get required parameters
dtype = batch.prompt_embeds[0].dtype
device = batch.device
device = fastvideo_args.device
generator = batch.generator
latents = batch.latents
num_frames = batch.num_frames
num_frames = latent_num_frames if latent_num_frames is not None else batch.num_frames
height = batch.height
width = batch.width
@@ -69,16 +69,15 @@ class LatentPreparationStage(PipelineStage):
if height is None or width is None:
raise ValueError("Height and width must be provided")
assert fastvideo_args.num_channels_latents is not None
assert fastvideo_args.vae_scale_factor is not None
# Calculate latent shape
shape = (
batch_size,
fastvideo_args.num_channels_latents,
self.transformer.num_channels_latents,
num_frames,
height // fastvideo_args.vae_scale_factor,
width // fastvideo_args.vae_scale_factor,
height //
fastvideo_args.vae_config.arch_config.spatial_compression_ratio,
width //
fastvideo_args.vae_config.arch_config.spatial_compression_ratio,
)
# Validate generator if it's a list
@@ -106,8 +105,8 @@ class LatentPreparationStage(PipelineStage):
return batch
def adjust_video_length(self, vae: ParallelTiledVAE, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
def adjust_video_length(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> int:
"""
Adjust video length based on VAE version.
@@ -119,7 +118,7 @@ class LatentPreparationStage(PipelineStage):
The batch with adjusted video length.
"""
video_length = batch.num_frames
temporal_scale_factor = vae.temporal_compression_ratio if vae is not None else 4
temporal_scale_factor = fastvideo_args.vae_config.arch_config.temporal_compression_ratio
# TODO
batch.num_frames = (video_length - 1) // temporal_scale_factor + 1
return batch
latent_num_frames = (video_length - 1) // temporal_scale_factor + 1
return latent_num_frames
@@ -74,7 +74,7 @@ class LlamaEncodingStage(PipelineStage):
The batch with encoded prompt embeddings.
"""
if fastvideo_args.use_cpu_offload:
self.text_encoder = self.text_encoder.to(batch.device)
self.text_encoder = self.text_encoder.to(fastvideo_args.device)
text = prompt_template_video["template"].format(batch.prompt)
text_inputs = self.tokenizer(
@@ -87,7 +87,7 @@ class LlamaEncodingStage(PipelineStage):
hidden_state_skip_layer = 2
with set_forward_context(current_timestep=0, attn_metadata=None):
outputs = self.text_encoder(
input_ids=text_inputs["input_ids"].to(batch.device),
input_ids=text_inputs["input_ids"].to(fastvideo_args.device),
output_hidden_states=hidden_state_skip_layer is not None,
)
@@ -111,7 +111,7 @@ class LlamaEncodingStage(PipelineStage):
with set_forward_context(current_timestep=0, attn_metadata=None):
negative_outputs = self.text_encoder(
input_ids=negative_text_inputs["input_ids"].to(
batch.device),
fastvideo_args.device),
output_hidden_states=hidden_state_skip_layer is not None,
)
+3 -3
View File
@@ -51,7 +51,7 @@ class T5EncodingStage(PipelineStage):
The batch with encoded prompt embeddings.
"""
if fastvideo_args.use_cpu_offload:
self.text_encoder = self.text_encoder.to(batch.device)
self.text_encoder = self.text_encoder.to(fastvideo_args.device)
text = batch.prompt
text_inputs = self.tokenizer(
@@ -62,7 +62,7 @@ class T5EncodingStage(PipelineStage):
add_special_tokens=True,
return_attention_mask=True,
return_tensors="pt",
).to(batch.device)
).to(fastvideo_args.device)
text_input_ids, mask = text_inputs.input_ids, text_inputs.attention_mask
seq_lens = mask.gt(0).sum(dim=1).long()
with set_forward_context(current_timestep=0, attn_metadata=None):
@@ -89,7 +89,7 @@ class T5EncodingStage(PipelineStage):
add_special_tokens=True,
return_attention_mask=True,
return_tensors="pt",
).to(batch.device)
).to(fastvideo_args.device)
text_input_ids, mask = negative_text_inputs.input_ids, negative_text_inputs.attention_mask
seq_lens = mask.gt(0).sum(dim=1).long()
with set_forward_context(current_timestep=0, attn_metadata=None):
@@ -42,7 +42,7 @@ class TimestepPreparationStage(PipelineStage):
The batch with prepared timesteps.
"""
scheduler = self.scheduler
device = batch.device
device = fastvideo_args.device
num_inference_steps = batch.num_inference_steps
timesteps = batch.timesteps
sigmas = batch.sigmas
+1 -11
View File
@@ -55,7 +55,7 @@ class WanImageToVideoPipeline(ComposedPipelineBase):
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae")))
transformer=self.get_module("transformer")))
self.add_stage(stage_name="image_latent_preparation_stage",
stage=EncodingStage(vae=self.get_module("vae")))
@@ -68,15 +68,5 @@ class WanImageToVideoPipeline(ComposedPipelineBase):
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""
Initialize the pipeline.
"""
vae_scale_factor = self.get_module("vae").spatial_compression_ratio
fastvideo_args.vae_scale_factor = vae_scale_factor
num_channels_latents = self.get_module("transformer").out_channels
fastvideo_args.num_channels_latents = num_channels_latents
EntryClass = WanImageToVideoPipeline
+1 -11
View File
@@ -48,7 +48,7 @@ class WanPipeline(ComposedPipelineBase):
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae")))
transformer=self.get_module("transformer")))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
@@ -58,15 +58,5 @@ class WanPipeline(ComposedPipelineBase):
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""
Initialize the pipeline.
"""
vae_scale_factor = self.get_module("vae").spatial_compression_ratio
fastvideo_args.vae_scale_factor = vae_scale_factor
num_channels_latents = self.get_module("transformer").in_channels
fastvideo_args.num_channels_latents = num_channels_latents
EntryClass = WanPipeline
@@ -1,3 +1,5 @@
# type: ignore
# SPDX-License-Identifier: Apache-2.0
import os
@@ -14,6 +14,7 @@ from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import maybe_download_model
from fastvideo.v1.configs.models.encoders import CLIPTextConfig
logger = init_logger(__name__)
@@ -39,7 +40,8 @@ def test_clip_encoder():
- Produce nearly identical outputs for the same input prompts
"""
args = FastVideoArgs(model_path="openai/clip-vit-large-patch14",
precision="float16")
text_encoder_precision_2="fp16",
text_encoder_config_2=CLIPTextConfig())
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
logger.info("Loading models from %s", args.model_path)
@@ -60,7 +62,7 @@ def test_clip_encoder():
from fastvideo.v1.models.loader.component_loader import TextEncoderLoader
loader = TextEncoderLoader()
args.device_str = "cuda:0"
model2 = loader.load_model(TEXT_ENCODER_PATH, hf_config, device)
model2 = loader.load(TEXT_ENCODER_PATH, "", args)
# Load the HuggingFace implementation directly
# model2 = CLIPTextModel(hf_config)
@@ -13,6 +13,7 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import TextEncoderLoader
from fastvideo.v1.utils import maybe_download_model
from fastvideo.v1.configs.models.encoders import LlamaConfig
logger = init_logger(__name__)
@@ -39,7 +40,8 @@ def test_llama_encoder():
- Produce nearly identical outputs for the same input prompts
"""
args = FastVideoArgs(model_path="meta-llama/Llama-2-7b-hf",
precision="float16")
precision="float16",
text_encoder_config=LlamaConfig())
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
@@ -57,7 +59,7 @@ def test_llama_encoder():
loader = TextEncoderLoader()
args.device_str = "cuda:0"
device = torch.device(args.device_str)
model2 = loader.load_model(TEXT_ENCODER_PATH, hf_config, device)
model2 = loader.load(TEXT_ENCODER_PATH, "", args)
# Convert to float16 and move to device
model2 = model2.to(torch.float16)
@@ -10,6 +10,8 @@ from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import TextEncoderLoader
from fastvideo.v1.utils import maybe_download_model
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.configs.models.encoders import T5Config
logger = init_logger(__name__)
@@ -36,8 +38,9 @@ def test_t5_encoder():
precision).to(device).eval()
tokenizer = AutoTokenizer.from_pretrained(TOKENIZER_PATH)
args = FastVideoArgs(model_path=TEXT_ENCODER_PATH, text_encoder_config=T5Config(), device_str="cuda")
loader = TextEncoderLoader()
model2 = loader.load_model(TEXT_ENCODER_PATH, hf_config, device)
model2 = loader.load(TEXT_ENCODER_PATH, "", args)
# Convert to float16 and move to device
model2 = model2.to(precision)
@@ -52,7 +52,7 @@ def initialize_identical_weights(model, seed=42):
return model
@pytest.mark.usefixtures("distributed_setup")
@pytest.mark.skip(reason="Incompatible with the new config")
def test_hunyuanvideo_distributed():
# Get tensor parallel info
sp_rank = get_sequence_model_parallel_rank()
@@ -10,11 +10,11 @@ from fastvideo.v1.distributed.parallel_state import (
get_sequence_model_parallel_rank,
get_sequence_model_parallel_world_size)
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.dits.hunyuanvideo import (
HunyuanVideoTransformer3DModel as HunyuanVideoDit)
from fastvideo.v1.models.loader.fsdp_load import load_fsdp_model
from fastvideo.v1.models.loader.component_loader import TransformerLoader
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.utils import maybe_download_model
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.configs.models.dits import HunyuanVideoConfig
logger = init_logger(__name__)
@@ -59,13 +59,15 @@ def test_hunyuanvideo_distributed():
config.pop("_class_name")
config.pop("_diffusers_version")
weight_dir_list = glob.glob(os.path.join(TRANSFORMER_PATH, "*.safetensors"))
weight_dir_list = [str(path) for path in weight_dir_list]
model = load_fsdp_model(HunyuanVideoDit,
init_params=config,
weight_dir_list=weight_dir_list,
device=torch.device(f"cuda:{LOCAL_RANK}"),
cpu_offload=False)
precision_str = "bf16"
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
use_cpu_offload=False,
precision=precision_str)
args.device = torch.device(f"cuda:{LOCAL_RANK}")
args.dit_config = HunyuanVideoConfig()
loader = TransformerLoader()
model = loader.load(TRANSFORMER_PATH, "", args)
model.eval()
@@ -11,6 +11,7 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import TransformerLoader
from fastvideo.v1.utils import maybe_download_model
from fastvideo.v1.configs.models.dits import WanVideoConfig
logger = init_logger(__name__)
@@ -33,6 +34,7 @@ def test_wan_transformer():
use_cpu_offload=False,
precision=precision_str)
args.device = device
args.dit_config = WanVideoConfig()
loader = TransformerLoader()
model2 = loader.load(TRANSFORMER_PATH, "", args).to(device, dtype=precision)
@@ -113,6 +115,8 @@ def test_wan_transformer():
# Check if outputs are similar (allowing for small numerical differences)
max_diff = torch.max(torch.abs(output1 - output2))
mean_diff = torch.mean(torch.abs(output1 - output2))
logger.info("Max Diff: %s", max_diff.item())
logger.info("Mean Diff: %s", mean_diff.item())
assert max_diff < 1e-1, f"Maximum difference between outputs: {max_diff.item()}"
# mean diff
assert mean_diff < 1e-2, f"Mean difference between outputs: {mean_diff.item()}"
+12 -16
View File
@@ -8,8 +8,11 @@ import torch
from safetensors.torch import load_file
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.vaes.hunyuanvae import (
AutoencoderKLHunyuanVideo as MyHunyuanVAE)
# from fastvideo.v1.models.vaes.hunyuanvae import (
# AutoencoderKLHunyuanVideo as MyHunyuanVAE)
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.models.loader.component_loader import VAELoader
from fastvideo.v1.configs.models.vaes import HunyuanVAEConfig
from fastvideo.v1.utils import maybe_download_model
logger = init_logger(__name__)
@@ -31,21 +34,14 @@ REFERENCE_LATENT = -105.51324462890625
@pytest.mark.usefixtures("distributed_setup")
def test_hunyuan_vae():
device = torch.device("cuda:0")
# Initialize the two model implementations
config = json.load(open(CONFIG_PATH))
config.pop("_class_name")
config.pop("_diffusers_version")
model = MyHunyuanVAE(**config).to(torch.bfloat16)
precision = torch.bfloat16
precision_str = "bf16"
args = FastVideoArgs(model_path=VAE_PATH, vae_precision=precision_str)
args.device = device
args.vae_config = HunyuanVAEConfig()
loaded = load_file(os.path.join(VAE_PATH,
"diffusion_pytorch_model.safetensors"))
model.load_state_dict(loaded)
# Set model to eval mode
model.eval()
# Move to GPU
model = model.to(device)
loader = VAELoader()
model = loader.load(VAE_PATH, "", args)
model.enable_tiling(tile_sample_min_height=32,
tile_sample_min_width=32,
+9 -11
View File
@@ -9,6 +9,7 @@ from diffusers import AutoencoderKLWan
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import VAELoader
from fastvideo.v1.configs.models.vaes import WanVAEConfig
from fastvideo.v1.utils import maybe_download_model
logger = init_logger(__name__)
@@ -30,6 +31,7 @@ def test_wan_vae():
precision_str = "bf16"
args = FastVideoArgs(model_path=VAE_PATH, vae_precision=precision_str)
args.device = device
args.vae_config = WanVAEConfig()
loader = VAELoader()
model2 = loader.load(VAE_PATH, "", args)
@@ -77,23 +79,19 @@ def test_wan_vae():
# Test decoding
logger.info("Testing decoding...")
latent1_tensor = latent1.mode()
latents_mean = (torch.tensor(model1.config.latents_mean).view(
mean1 = (torch.tensor(model1.config.latents_mean).view(
1, model1.config.z_dim, 1, 1, 1).to(input_tensor.device,
input_tensor.dtype))
latents_std = 1.0 / torch.tensor(model1.config.latents_std).view(
1, model1.config.z_dim, 1, 1, 1).to(input_tensor.device,
std1 = (1.0 / torch.tensor(model1.config.latents_std).view(
1, model1.config.z_dim, 1, 1, 1)).to(input_tensor.device,
input_tensor.dtype)
latent1_tensor = latent1_tensor / latents_std + latents_mean
latent1_tensor = latent1_tensor / std1 + mean1
output1 = model1.decode(latent1_tensor).sample
mean2 = model2.config.arch_config.shift_factor.to(input_tensor.device, input_tensor.dtype)
std2 = model2.config.arch_config.scaling_factor.to(input_tensor.device, input_tensor.dtype)
latent2_tensor = latent2.mode()
latents_mean = (torch.tensor(model2.config.latents_mean).view(
1, model2.config.z_dim, 1, 1, 1).to(input_tensor.device,
input_tensor.dtype))
latents_std = 1.0 / torch.tensor(model2.config.latents_std).view(
1, model2.config.z_dim, 1, 1, 1).to(input_tensor.device,
input_tensor.dtype)
latent2_tensor = latent2_tensor / latents_std + latents_mean
latent2_tensor = latent2_tensor / std2 + mean2
output2 = model2.decode(latent2_tensor)
# Check if outputs have the same shape
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
+5 -9
View File
@@ -13,7 +13,7 @@ import signal
import sys
import tempfile
import traceback
from dataclasses import asdict, fields
from dataclasses import fields, is_dataclass
from functools import partial, wraps
from typing import (Any, Callable, Dict, List, Optional, Tuple, Type, TypeVar,
Union, cast)
@@ -551,14 +551,10 @@ def run_method(obj: Any, method: Union[str, bytes, Callable], args: tuple[Any],
return func(*args, **kwargs)
def diff_keys(a, b):
return [k for k in asdict(a) if asdict(a)[k] != asdict(b)[k]]
def update_in_place(target, source, ignore_fields=()) -> None:
for f in fields(target):
if hasattr(source, f.name) and f.name not in list(ignore_fields):
setattr(target, f.name, getattr(source, f.name))
def shallow_asdict(obj) -> Dict[str, Any]:
if not is_dataclass(obj):
raise TypeError("Expected dataclass instance")
return {f.name: getattr(obj, f.name) for f in fields(obj)}
def kill_itself_when_parent_died() -> None:
-1
View File
@@ -89,7 +89,6 @@ class Worker:
def execute_forward(self, forward_batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
self.fastvideo_args.num_inference_steps = fastvideo_args.num_inference_steps
output_batch = self.pipeline.forward(forward_batch, self.fastvideo_args)
return cast(ForwardBatch, output_batch)