Compare commits
4
Commits
klin/trackwan
...
wei/api
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
38ee9dc3b4 | ||
|
|
77b013fb8a | ||
|
|
c31efe1234 | ||
|
|
c056b89aea |
@@ -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"
|
||||
]
|
||||
|
||||
@@ -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
|
||||
@@ -0,0 +1,7 @@
|
||||
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
|
||||
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
|
||||
__all__ = [
|
||||
"ArchConfig", "ModelConfig",
|
||||
"VAEArchConfig", "VAEConfig"
|
||||
]
|
||||
@@ -0,0 +1,47 @@
|
||||
from dataclasses import dataclass, fields
|
||||
from typing import Dict, Any
|
||||
|
||||
# 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 & overriden 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
|
||||
|
||||
# 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}'")
|
||||
|
||||
def update_model_config(
|
||||
self,
|
||||
source_model_dict: Dict[str, Any]
|
||||
) -> None:
|
||||
assert "arch_config" not in source_model_dict.keys(), "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:
|
||||
print(f"{type(self).__name__} does not contain field '{key}'!")
|
||||
raise AttributeError(f"Invalid field: {key}")
|
||||
@@ -0,0 +1,7 @@
|
||||
from fastvideo.v1.configs.models.vaes.hunyuanvae import HunyuanVAEConfig, HunyuanVAEArchConfig
|
||||
from fastvideo.v1.configs.models.vaes.wanvae import WanVAEConfig, WanVAEArchConfig
|
||||
|
||||
__all__ = [
|
||||
"HunyuanVAEConfig", "HunyuanVAEArchConfig",
|
||||
"WanVAEConfig", "WanVAEArchConfig"
|
||||
]
|
||||
@@ -0,0 +1,36 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Union
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.configs.models 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,37 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Tuple
|
||||
|
||||
from fastvideo.v1.configs.models.vaes.base import VAEConfig, VAEArchConfig
|
||||
|
||||
@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,72 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.configs.models.vaes.base import VAEConfig, VAEArchConfig
|
||||
|
||||
@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,9 @@
|
||||
from fastvideo.v1.configs.pipelines.hunyuan import HunyuanConfig, FastHunyuanConfig
|
||||
from fastvideo.v1.configs.pipelines.wan import WanT2V480PConfig, WanI2V480PConfig
|
||||
from fastvideo.v1.configs.pipelines.base import BaseConfig, SlidingTileAttnConfig
|
||||
from fastvideo.v1.configs.pipelines.registry import get_pipeline_config_cls_for_name
|
||||
|
||||
__all__ = [
|
||||
"HunyuanConfig", "FastHunyuanConfig", "BaseConfig", "SlidingTileAttnConfig",
|
||||
"WanT2V480PConfig", "WanI2V480PConfig", "get_pipeline_config_cls_for_name"
|
||||
]
|
||||
@@ -0,0 +1,103 @@
|
||||
from dataclasses import dataclass, asdict, fields
|
||||
from typing import Optional, Dict, Any
|
||||
import json
|
||||
|
||||
from fastvideo.v1.configs.models import ModelConfig, VAEConfig
|
||||
from fastvideo.v1.utils import shallow_asdict
|
||||
|
||||
@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 # Deprecated
|
||||
vae_config: VAEConfig = VAEConfig()
|
||||
|
||||
# DiT configuration
|
||||
num_channels_latents: Optional[int] = None # Deprecated
|
||||
|
||||
# Image encoder configuration
|
||||
image_encoder_precision: str = "fp32"
|
||||
|
||||
# Text encoder configuration
|
||||
text_encoder_precision: str = "fp16"
|
||||
text_len: int = -1 # Deprecated
|
||||
hidden_state_skip_layer: int = 0 # Deprecated
|
||||
|
||||
# STA (Spatial-Temporal Attention) parameters
|
||||
mask_strategy_file_path: Optional[str] = None
|
||||
enable_torch_compile: bool = False
|
||||
|
||||
neg_prompt: Optional[str] = None
|
||||
|
||||
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, "r") 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)
|
||||
|
||||
@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
|
||||
@@ -1,13 +1,16 @@
|
||||
from dataclasses import dataclass
|
||||
from fastvideo.v1.configs.pipelines.base import BaseConfig
|
||||
|
||||
from fastvideo.v1.configs.base import BaseConfig
|
||||
|
||||
from fastvideo.v1.configs.models import VAEConfig
|
||||
from fastvideo.v1.configs.models.vaes import HunyuanVAEConfig
|
||||
|
||||
@dataclass
|
||||
class HunyuanConfig(BaseConfig):
|
||||
"""Base configuration for HunYuan pipeline architecture."""
|
||||
|
||||
# HunyuanConfig-specific parameters with defaults
|
||||
# VAE
|
||||
vae_config: VAEConfig = HunyuanVAEConfig()
|
||||
# Denoising stage
|
||||
embedded_cfg_scale: int = 6
|
||||
flow_shift: int = 7
|
||||
@@ -27,6 +30,10 @@ class HunyuanConfig(BaseConfig):
|
||||
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
|
||||
class FastHunyuanConfig(HunyuanConfig):
|
||||
@@ -3,9 +3,10 @@
|
||||
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 BaseConfig
|
||||
from fastvideo.v1.configs.pipelines.hunyuan import HunyuanConfig, FastHunyuanConfig
|
||||
from fastvideo.v1.configs.pipelines.wan import WanT2V480PConfig, WanI2V480PConfig
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import (maybe_download_model_index,
|
||||
verify_model_config_and_directory)
|
||||
@@ -1,13 +1,19 @@
|
||||
from dataclasses import dataclass
|
||||
from fastvideo.v1.configs.pipelines.base import BaseConfig
|
||||
|
||||
from fastvideo.v1.configs.base import BaseConfig
|
||||
|
||||
from fastvideo.v1.configs.models import VAEConfig
|
||||
from fastvideo.v1.configs.models.vaes import WanVAEConfig
|
||||
|
||||
@dataclass
|
||||
class WanT2V480PConfig(BaseConfig):
|
||||
"""Base configuration for Wan T2V 1.3B pipeline architecture."""
|
||||
|
||||
# WanConfig-specific parameters with defaults
|
||||
# VAE
|
||||
vae_config: VAEConfig = WanVAEConfig()
|
||||
vae_tiling: bool = False
|
||||
vae_sp: bool = False
|
||||
|
||||
# Video parameters
|
||||
height: int = 480
|
||||
width: int = 832
|
||||
@@ -31,6 +37,9 @@ class WanT2V480PConfig(BaseConfig):
|
||||
|
||||
# WanConfig-specific added parameters
|
||||
|
||||
def __post_init__(self):
|
||||
self.vae_config.load_encoder = False
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
@dataclass
|
||||
class WanI2V480PConfig(WanT2V480PConfig):
|
||||
@@ -43,3 +52,7 @@ class WanI2V480PConfig(WanT2V480PConfig):
|
||||
|
||||
# Precision for each component
|
||||
image_encoder_precision: str = "fp32"
|
||||
|
||||
def __post_init__(self):
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
@@ -0,0 +1,41 @@
|
||||
{
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 81,
|
||||
"fps": 16,
|
||||
"num_inference_steps": 50,
|
||||
"guidance_scale": 3.0,
|
||||
"seed": 1024,
|
||||
"guidance_rescale": 0.0,
|
||||
"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
|
||||
},
|
||||
"num_channels_latents": null,
|
||||
"image_encoder_precision": "fp32",
|
||||
"text_encoder_precision": "fp32",
|
||||
"text_len": 512,
|
||||
"hidden_state_skip_layer": 0,
|
||||
"mask_strategy_file_path": null,
|
||||
"enable_torch_compile": false,
|
||||
"neg_prompt": "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"
|
||||
}
|
||||
@@ -0,0 +1,41 @@
|
||||
{
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 81,
|
||||
"fps": 16,
|
||||
"num_inference_steps": 40,
|
||||
"guidance_scale": 5.0,
|
||||
"seed": 1024,
|
||||
"guidance_rescale": 0.0,
|
||||
"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
|
||||
},
|
||||
"num_channels_latents": null,
|
||||
"image_encoder_precision": "fp32",
|
||||
"text_encoder_precision": "fp32",
|
||||
"text_len": 512,
|
||||
"hidden_state_skip_layer": 0,
|
||||
"mask_strategy_file_path": null,
|
||||
"enable_torch_compile": false,
|
||||
"neg_prompt": "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"
|
||||
}
|
||||
@@ -17,11 +17,11 @@ 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 get_pipeline_config_cls_for_name, BaseConfig
|
||||
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 +52,7 @@ class VideoGenerator:
|
||||
model_path: str,
|
||||
device: Optional[str] = None,
|
||||
torch_dtype: Optional[torch.dtype] = None,
|
||||
pipeline_config: Optional[Union[str | BaseConfig]] = None,
|
||||
**kwargs) -> "VideoGenerator":
|
||||
"""
|
||||
Create a video generator from a pretrained model.
|
||||
@@ -64,22 +65,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, BaseConfig):
|
||||
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,6 +124,7 @@ class VideoGenerator:
|
||||
def generate_video(
|
||||
self,
|
||||
prompt: str,
|
||||
image_path: Optional[str] = None,
|
||||
negative_prompt: Optional[str] = None,
|
||||
output_path: Optional[str] = None,
|
||||
save_video: bool = True,
|
||||
@@ -155,6 +165,8 @@ class VideoGenerator:
|
||||
fastvideo_args = self.fastvideo_args
|
||||
|
||||
# Override parameters if provided
|
||||
if image_path is not None:
|
||||
fastvideo_args.image_path = image_path
|
||||
if negative_prompt is not None:
|
||||
fastvideo_args.neg_prompt = negative_prompt
|
||||
if num_inference_steps is not None:
|
||||
@@ -224,6 +236,7 @@ class VideoGenerator:
|
||||
device = torch.device(fastvideo_args.device_str)
|
||||
batch = ForwardBatch(
|
||||
prompt=prompt,
|
||||
image_path=fastvideo_args.image_path,
|
||||
negative_prompt=fastvideo_args.neg_prompt,
|
||||
num_videos_per_prompt=fastvideo_args.num_videos,
|
||||
height=fastvideo_args.height,
|
||||
|
||||
@@ -10,6 +10,8 @@ from typing import List, Optional
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import FlexibleArgumentParser
|
||||
|
||||
from fastvideo.v1.configs.models import VAEConfig
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -50,9 +52,10 @@ class FastVideoArgs:
|
||||
|
||||
# VAE configuration
|
||||
vae_precision: str = "fp16"
|
||||
vae_tiling: bool = True
|
||||
vae_sp: bool = False
|
||||
vae_scale_factor: 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()
|
||||
|
||||
# DiT configuration
|
||||
num_channels_latents: Optional[int] = None
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import dataclasses
|
||||
from dataclasses import asdict
|
||||
import glob
|
||||
import os
|
||||
import time
|
||||
@@ -310,9 +311,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(
|
||||
@@ -322,7 +325,7 @@ 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)
|
||||
|
||||
@@ -343,6 +346,10 @@ class TransformerLoader(ComponentLoader):
|
||||
"Only diffusers format is supported.")
|
||||
model_config.pop("_diffusers_version")
|
||||
|
||||
# Config from Diffusers supercedes fastvideo's model config
|
||||
# dit_config = fastvideo_args.dit_config
|
||||
# model_config.update(dit_config)
|
||||
|
||||
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
|
||||
|
||||
# Find all safetensors files
|
||||
|
||||
@@ -11,6 +11,7 @@ from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from fastvideo.v1.distributed import (get_sequence_model_parallel_rank,
|
||||
get_sequence_model_parallel_world_size)
|
||||
from fastvideo.v1.configs.models import VAEConfig
|
||||
|
||||
|
||||
class ParallelTiledVAE(ABC):
|
||||
@@ -20,29 +21,36 @@ 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.arch_config = config.arch_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 self.arch_config.temporal_compression_ratio
|
||||
|
||||
@property
|
||||
def spatial_compression_ratio(self) -> int:
|
||||
return self.arch_config.spatial_compression_ratio
|
||||
|
||||
@property
|
||||
def scaling_factor(self) -> Union[float, torch.tensor]:
|
||||
return self.arch_config.scaling_factor
|
||||
|
||||
@abstractmethod
|
||||
def _encode(self, *args, **kwargs) -> torch.Tensor:
|
||||
@@ -408,6 +416,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 +451,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"""
|
||||
|
||||
@@ -15,7 +15,7 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from typing import Optional, Tuple, Union
|
||||
from typing import Optional, Tuple, Union, cast
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
@@ -24,8 +24,8 @@ import torch.nn.functional as F
|
||||
import torch.utils.checkpoint
|
||||
|
||||
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
|
||||
from fastvideo.v1.configs.models.vaes import HunyuanVAEConfig, HunyuanVAEArchConfig
|
||||
|
||||
|
||||
def prepare_causal_attention_mask(
|
||||
@@ -773,95 +773,53 @@ 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__()
|
||||
ParallelTiledVAE.__init__(self, config)
|
||||
arch_config: HunyuanVAEArchConfig = cast(HunyuanVAEArchConfig, config.arch_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
|
||||
|
||||
if load_encoder:
|
||||
self.block_out_channels = arch_config.block_out_channels
|
||||
|
||||
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=arch_config.in_channels,
|
||||
out_channels=arch_config.latent_channels,
|
||||
down_block_types=arch_config.down_block_types,
|
||||
block_out_channels=arch_config.block_out_channels,
|
||||
layers_per_block=arch_config.layers_per_block,
|
||||
norm_num_groups=arch_config.norm_num_groups,
|
||||
act_fn=arch_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=arch_config.mid_block_add_attention,
|
||||
temporal_compression_ratio=arch_config.temporal_compression_ratio,
|
||||
spatial_compression_ratio=arch_config.spatial_compression_ratio,
|
||||
)
|
||||
self.quant_conv = nn.Conv3d(2 * latent_channels,
|
||||
2 * latent_channels,
|
||||
self.quant_conv = nn.Conv3d(2 * arch_config.latent_channels,
|
||||
2 * arch_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=arch_config.latent_channels,
|
||||
out_channels=arch_config.out_channels,
|
||||
up_block_types=arch_config.up_block_types,
|
||||
block_out_channels=arch_config.block_out_channels,
|
||||
layers_per_block=arch_config.layers_per_block,
|
||||
norm_num_groups=arch_config.norm_num_groups,
|
||||
act_fn=arch_config.act_fn,
|
||||
time_compression_ratio=arch_config.temporal_compression_ratio,
|
||||
spatial_compression_ratio=arch_config.spatial_compression_ratio,
|
||||
mid_block_add_attention=arch_config.mid_block_add_attention,
|
||||
)
|
||||
self.post_quant_conv = nn.Conv3d(latent_channels,
|
||||
latent_channels,
|
||||
self.post_quant_conv = nn.Conv3d(arch_config.latent_channels,
|
||||
arch_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)
|
||||
|
||||
@@ -16,16 +16,15 @@
|
||||
|
||||
import contextvars
|
||||
from contextlib import contextmanager
|
||||
from typing import Optional, Tuple, Union
|
||||
from typing import Optional, Tuple, Union, cast
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
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)
|
||||
from fastvideo.v1.models.vaes.common import ParallelTiledVAE, DiagonalGaussianDistribution
|
||||
from fastvideo.v1.configs.models.vaes import WanVAEConfig, WanVAEArchConfig
|
||||
|
||||
CACHE_T = 2
|
||||
|
||||
@@ -781,96 +780,34 @@ 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:
|
||||
config: WanVAEConfig,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
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.arch_config = cast(WanVAEArchConfig, self.arch_config)
|
||||
|
||||
self.z_dim = self.arch_config.z_dim
|
||||
self.temperal_downsample = list(self.arch_config.temperal_downsample)
|
||||
self.temperal_upsample = list(self.arch_config.temperal_downsample)[::-1]
|
||||
self.latents_mean = list(self.arch_config.latents_mean)
|
||||
self.latents_std = list(self.arch_config.latents_std)
|
||||
self.shift_factor = self.arch_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(self.arch_config.base_dim, self.z_dim * 2, self.arch_config.dim_mult,
|
||||
self.arch_config.num_res_blocks, self.arch_config.attn_scales,
|
||||
self.temperal_downsample, self.arch_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(self.arch_config.base_dim, self.z_dim, self.arch_config.dim_mult,
|
||||
self.arch_config.num_res_blocks, self.arch_config.attn_scales,
|
||||
self.temperal_upsample, self.arch_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 +818,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:
|
||||
|
||||
@@ -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
|
||||
@@ -71,14 +69,6 @@ class HunyuanVideoPipeline(ComposedPipelineBase):
|
||||
"""
|
||||
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
|
||||
|
||||
|
||||
@@ -10,6 +10,7 @@ from fastvideo.v1.logger import init_logger
|
||||
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
|
||||
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -15,6 +15,7 @@ from fastvideo.v1.models.vision_utils import (get_default_height_width,
|
||||
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
|
||||
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,9 @@ class LatentPreparationStage(PipelineStage):
|
||||
denoised during the diffusion process.
|
||||
"""
|
||||
|
||||
def __init__(self, scheduler, vae=None) -> None:
|
||||
def __init__(self, scheduler) -> None:
|
||||
super().__init__()
|
||||
self.scheduler = scheduler
|
||||
self.vae = vae
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -44,7 +42,7 @@ class LatentPreparationStage(PipelineStage):
|
||||
|
||||
# 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)
|
||||
batch = self.adjust_video_length(batch, fastvideo_args)
|
||||
# Determine batch size
|
||||
if isinstance(batch.prompt, list):
|
||||
batch_size = len(batch.prompt)
|
||||
@@ -70,15 +68,14 @@ class LatentPreparationStage(PipelineStage):
|
||||
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,
|
||||
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,7 +103,7 @@ class LatentPreparationStage(PipelineStage):
|
||||
|
||||
return batch
|
||||
|
||||
def adjust_video_length(self, vae: ParallelTiledVAE, batch: ForwardBatch,
|
||||
def adjust_video_length(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
"""
|
||||
Adjust video length based on VAE version.
|
||||
@@ -119,7 +116,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
|
||||
|
||||
@@ -54,8 +54,7 @@ class WanImageToVideoPipeline(ComposedPipelineBase):
|
||||
|
||||
self.add_stage(stage_name="latent_preparation_stage",
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae")))
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="image_latent_preparation_stage",
|
||||
stage=EncodingStage(vae=self.get_module("vae")))
|
||||
@@ -72,9 +71,6 @@ class WanImageToVideoPipeline(ComposedPipelineBase):
|
||||
"""
|
||||
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
|
||||
|
||||
|
||||
@@ -47,8 +47,7 @@ class WanPipeline(ComposedPipelineBase):
|
||||
|
||||
self.add_stage(stage_name="latent_preparation_stage",
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae")))
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
@@ -62,9 +61,6 @@ class WanPipeline(ComposedPipelineBase):
|
||||
"""
|
||||
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
|
||||
|
||||
|
||||
@@ -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,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
-11
@@ -13,7 +13,7 @@ import signal
|
||||
import sys
|
||||
import tempfile
|
||||
import traceback
|
||||
from dataclasses import asdict, fields
|
||||
from dataclasses import asdict, fields, is_dataclass
|
||||
from functools import partial, wraps
|
||||
from typing import (Any, Callable, Dict, List, Optional, Tuple, Type, TypeVar,
|
||||
Union, cast)
|
||||
@@ -550,16 +550,10 @@ def run_method(obj: Any, method: Union[str, bytes, Callable], args: tuple[Any],
|
||||
func = partial(method, obj) # type: ignore
|
||||
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):
|
||||
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:
|
||||
# if sys.platform == "linux":
|
||||
|
||||
Reference in New Issue
Block a user