Compare commits

...
23 changed files with 2445 additions and 96 deletions
+1
View File
@@ -20,6 +20,7 @@ exclude: |
fastvideo/train\.py|
fastvideo/utils/.*|
examples/.*|
fastvideo/v1/models/schedulers/scheduling_flow_match_euler_discrete.py|
.github/workflows/fastvideo-publish.yml|
.github/workflows/sta-publish.yml|
.github/workflows/build-image-template.yml
+8
View File
@@ -42,6 +42,14 @@ class ModelConfig:
# This should be used only when loading from transformers/diffusers
def update_model_arch(self, source_model_dict: Dict[str, Any]) -> None:
# Remove all keys that start with "_"
keys_to_remove = [
key for key in list(source_model_dict.keys())
if str(key).startswith("_")
]
for key in keys_to_remove:
source_model_dict.pop(key)
arch_config = self.arch_config
valid_fields = {f.name for f in fields(arch_config)}
+2 -1
View File
@@ -1,4 +1,5 @@
from fastvideo.v1.configs.models.dits.flux import FluxImageConfig
from fastvideo.v1.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
from fastvideo.v1.configs.models.dits.wanvideo import WanVideoConfig
__all__ = ["HunyuanVideoConfig", "WanVideoConfig"]
__all__ = ["HunyuanVideoConfig", "WanVideoConfig", "FluxImageConfig"]
+130
View File
@@ -0,0 +1,130 @@
from dataclasses import dataclass, field
from typing import Optional, Tuple
import torch
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 FluxImageArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
_param_names_mapping: dict = field(
default_factory=lambda: {
# 1. context_embedder to txt_in mapping:
r"^context_embedder\.(.*)$":
r"txt_in.\1",
# 2. x_embedder to img_in mapping:
r"^x_embedder\.(.*)$":
r"img_in.\1",
# 3. 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"txt2_in.fc_in.\1",
r"^time_text_embed\.text_embedder\.linear_2\.(.*)$":
r"txt2_in.fc_out.\1",
# 4. 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",
# 5. 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",
# 6. 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 = 1
in_channels: int = 64
out_channels: Optional[int] = None
num_layers: int = 19
num_single_layers: int = 38
attention_head_dim: int = 128
num_attention_heads: int = 24
joint_attention_dim: int = 4096
pooled_projection_dim: int = 768
guidance_embeds: bool = False
axes_dims_rope: Tuple[int, ...] = (16, 56, 56)
rope_theta: int = 10000
dtype: Optional[torch.dtype] = torch.bfloat16
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 // 4
@dataclass
class FluxImageConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=FluxImageArchConfig)
prefix: str = "Flux"
@@ -1,7 +1,9 @@
from fastvideo.v1.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
from fastvideo.v1.configs.models.vaes.image_vae import ImageVAEConfig
from fastvideo.v1.configs.models.vaes.wanvae import WanVAEConfig
__all__ = [
"HunyuanVAEConfig",
"WanVAEConfig",
"ImageVAEConfig",
]
@@ -0,0 +1,41 @@
from dataclasses import dataclass, field
from typing import Optional, Tuple
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class ImageVAEArchConfig(VAEArchConfig):
in_channels: int = 3
out_channels: int = 3
down_block_types: Tuple[str] = ("DownEncoderBlock2D", )
up_block_types: Tuple[str] = ("UpDecoderBlock2D", )
block_out_channels: Tuple[int] = (64, )
layers_per_block: int = 1
act_fn: str = "silu"
latent_channels: int = 4
norm_num_groups: int = 32
sample_size: int = 32
scaling_factor: float = 0.18215
shift_factor: Optional[float] = None
latents_mean: Optional[Tuple[float]] = None
latents_std: Optional[Tuple[float]] = None
force_upcast: float = True
use_quant_conv: bool = True
use_post_quant_conv: bool = True
mid_block_add_attention: bool = True
def __post_init__(self):
self.spatial_compression_ratio: int = 2**(len(self.block_out_channels) -
1)
self.temporal_compression_ratio = 1
@dataclass
class ImageVAEConfig(VAEConfig):
arch_config: VAEArchConfig = field(default_factory=ImageVAEArchConfig)
# overrides VAEConfig
use_tiling: bool = False
use_temporal_tiling: bool = False
use_parallel_tiling: bool = False
+68
View File
@@ -0,0 +1,68 @@
from dataclasses import dataclass, field
from typing import Callable, Tuple
import torch
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.v1.configs.models.dits import FluxImageConfig
from fastvideo.v1.configs.models.encoders import (BaseEncoderOutput,
CLIPTextConfig, T5Config)
from fastvideo.v1.configs.models.vaes import ImageVAEConfig
from fastvideo.v1.configs.pipelines.base import PipelineConfig
def t5_preprocess_text(prompt: str) -> str:
return prompt
def t5_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor:
hidden_state: torch.tensor = outputs.last_hidden_state
assert torch.isnan(hidden_state).sum() == 0
prompt_embeds_tensor: torch.tensor = torch.stack([
torch.cat([u, u.new_zeros(512 - u.size(0), u.size(1))])
for u in hidden_state
],
dim=0)
return prompt_embeds_tensor
def clip_preprocess_text(prompt: str) -> str:
return prompt
def clip_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor:
pooler_output: torch.tensor = outputs.pooler_output
return pooler_output
@dataclass
class FluxConfig(PipelineConfig):
"""Base configuration for Flux pipeline architecture."""
# FluxConfig-specific parameters with defaults
# DiT
dit_config: DiTConfig = field(default_factory=FluxImageConfig)
# VAE
vae_config: VAEConfig = field(default_factory=ImageVAEConfig)
# Denoising stage
embedded_cfg_scale: float = 3.5
# Text encoding stage
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
default_factory=lambda: (CLIPTextConfig(), T5Config()))
preprocess_text_funcs: Tuple[Callable[[str], str], ...] = field(
default_factory=lambda: (clip_preprocess_text, t5_preprocess_text))
postprocess_text_funcs: Tuple[
Callable[[BaseEncoderOutput], torch.tensor],
...] = field(default_factory=lambda:
(clip_postprocess_text, t5_postprocess_text))
# Precision for each component
precision: str = "bf16"
vae_precision: str = "fp16"
text_encoder_precisions: Tuple[str, ...] = field(
default_factory=lambda: ("bf16", "bf16"))
def __post_init__(self):
self.vae_config.load_encoder = False
self.vae_config.load_decoder = True
@@ -4,6 +4,7 @@ import os
from typing import Callable, Dict, Optional, Type
from fastvideo.v1.configs.pipelines.base import PipelineConfig
from fastvideo.v1.configs.pipelines.flux import FluxConfig
from fastvideo.v1.configs.pipelines.hunyuan import (FastHunyuanConfig,
HunyuanConfig)
from fastvideo.v1.configs.pipelines.wan import (WanI2V480PConfig,
@@ -24,6 +25,8 @@ WEIGHT_CONFIG_REGISTRY: Dict[str, Type[PipelineConfig]] = {
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V480PConfig,
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V720PConfig,
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V720PConfig,
"black-forest-labs/FLUX.1-dev": FluxConfig,
"black-forest-labs/FLUX.1-schnell": FluxConfig,
# Add other specific weight variants
}
@@ -32,6 +35,7 @@ PIPELINE_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(),
"flux": lambda id: "flux" in id.lower(),
# Add other pipeline architecture detectors
}
@@ -42,6 +46,7 @@ PIPELINE_FALLBACK_CONFIG: Dict[str, Type[PipelineConfig]] = {
"wanpipeline":
WanT2V480PConfig, # Base Wan config as fallback for any Wan variant
"wanimagetovideo": WanI2V480PConfig,
"flux": FluxConfig,
# Other fallbacks by architecture
}
@@ -0,0 +1,30 @@
from fastvideo import VideoGenerator
# from fastvideo.v1.configs.sample import SamplingParam
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
generator = VideoGenerator.from_pretrained(
"FastVideo/FastHunyuan-diffusers",
# if num_gpus > 1, FastVideo will automatically handle distributed setup
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!
prompt2 = "A beautiful woman in a blue dress walking down a street"
video2 = generator.generate_video(prompt2)
if __name__ == "__main__":
main()
+560
View File
@@ -0,0 +1,560 @@
from typing import List, Optional, Tuple, Union
import torch
import torch.nn as nn
from diffusers.models.normalization import RMSNorm
from fastvideo.v1.attention import DistributedAttention
from fastvideo.v1.configs.models.dits import FluxImageConfig
from fastvideo.v1.layers.layernorm import (LayerNormScaleShift, ScaleResidual,
ScaleResidualLayerNormScaleShift)
from fastvideo.v1.layers.linear import ReplicatedLinear
from fastvideo.v1.layers.mlp import MLP
from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
get_rotary_pos_embed)
from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
TimestepEmbedder)
from fastvideo.v1.models.dits.base import BaseDiT
from fastvideo.v1.platforms import _Backend
class MMDoubleStreamBlock(nn.Module):
"""
A multimodal DiT block with separate modulation for text and image/video,
using distributed attention and linear layers.
"""
def __init__(
self,
hidden_size: int,
num_attention_heads: int,
mlp_ratio: float = 4.0,
dtype: Optional[torch.dtype] = None,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
prefix: str = "",
):
super().__init__()
self.deterministic = False
self.num_attention_heads = num_attention_heads
head_dim = hidden_size // num_attention_heads
mlp_hidden_dim = int(hidden_size * mlp_ratio)
# Image modulation components
self.img_mod = ModulateProjection(
hidden_size,
factor=6,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.img_mod",
)
# Fused operations for image stream
self.img_attn_norm = LayerNormScaleShift(hidden_size,
norm_type="layer",
elementwise_affine=False,
dtype=dtype)
self.img_attn_residual_mlp_norm = ScaleResidualLayerNormScaleShift(
hidden_size,
norm_type="layer",
elementwise_affine=False,
dtype=dtype)
self.img_mlp_residual = ScaleResidual()
# Image attention components
self.img_attn_qkv = ReplicatedLinear(hidden_size,
hidden_size * 3,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.img_attn_qkv")
self.img_attn_q_norm = RMSNorm(head_dim, eps=1e-6)
self.img_attn_k_norm = RMSNorm(head_dim, eps=1e-6)
self.img_attn_proj = ReplicatedLinear(hidden_size,
hidden_size,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.img_attn_proj")
self.img_mlp = MLP(hidden_size,
mlp_hidden_dim,
bias=True,
dtype=dtype,
prefix=f"{prefix}.img_mlp")
# Text modulation components
self.txt_mod = ModulateProjection(
hidden_size,
factor=6,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.txt_mod",
)
# Fused operations for text stream
self.txt_attn_norm = LayerNormScaleShift(hidden_size,
norm_type="layer",
elementwise_affine=False,
dtype=dtype)
self.txt_attn_residual_mlp_norm = ScaleResidualLayerNormScaleShift(
hidden_size,
norm_type="layer",
elementwise_affine=False,
dtype=dtype)
self.txt_mlp_residual = ScaleResidual()
# Text attention components
self.txt_attn_qkv = ReplicatedLinear(hidden_size,
hidden_size * 3,
bias=True,
params_dtype=dtype)
# QK norm layers for text
self.txt_attn_q_norm = RMSNorm(head_dim, eps=1e-6)
self.txt_attn_k_norm = RMSNorm(head_dim, eps=1e-6)
self.txt_attn_proj = ReplicatedLinear(hidden_size,
hidden_size,
bias=True,
params_dtype=dtype)
self.txt_mlp = MLP(hidden_size, mlp_hidden_dim, bias=True, dtype=dtype)
# Distributed attention
self.attn = DistributedAttention(
num_heads=num_attention_heads,
head_size=head_dim,
causal=False,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn")
def forward(
self,
img: torch.Tensor,
txt: torch.Tensor,
vec: torch.Tensor,
freqs_cis_img: Tuple[torch.Tensor, torch.Tensor],
) -> Tuple[torch.Tensor, torch.Tensor]:
# Process modulation vectors
img_mod_outputs = self.img_mod(vec)
(
img_attn_shift,
img_attn_scale,
img_attn_gate,
img_mlp_shift,
img_mlp_scale,
img_mlp_gate,
) = torch.chunk(img_mod_outputs, 6, dim=-1)
txt_mod_outputs = self.txt_mod(vec)
(
txt_attn_shift,
txt_attn_scale,
txt_attn_gate,
txt_mlp_shift,
txt_mlp_scale,
txt_mlp_gate,
) = torch.chunk(txt_mod_outputs, 6, dim=-1)
# Prepare image for attention using fused operation
img_attn_input = self.img_attn_norm(img, img_attn_shift, img_attn_scale)
# Get QKV for image
img_qkv, _ = self.img_attn_qkv(img_attn_input)
batch_size, image_seq_len = img_qkv.shape[0], img_qkv.shape[1]
# Split QKV
img_qkv = img_qkv.view(batch_size, image_seq_len, 3,
self.num_attention_heads, -1)
img_q, img_k, img_v = img_qkv[:, :, 0], img_qkv[:, :, 1], img_qkv[:, :,
2]
# Apply QK-Norm if needed
img_q = self.img_attn_q_norm(img_q).to(img_v)
img_k = self.img_attn_k_norm(img_k).to(img_v)
# Apply rotary embeddings for image
cos, sin = freqs_cis_img
img_q, img_k = _apply_rotary_emb(
img_q, cos, sin,
is_neox_style=False), _apply_rotary_emb(img_k,
cos,
sin,
is_neox_style=False)
# Prepare text for attention using fused operation
txt_attn_input = self.txt_attn_norm(txt, txt_attn_shift, txt_attn_scale)
# Get QKV for text
txt_qkv, _ = self.txt_attn_qkv(txt_attn_input)
batch_size, text_seq_len = txt_qkv.shape[0], txt_qkv.shape[1]
# Split QKV
txt_qkv = txt_qkv.view(batch_size, text_seq_len, 3,
self.num_attention_heads, -1)
txt_q, txt_k, txt_v = txt_qkv[:, :, 0], txt_qkv[:, :, 1], txt_qkv[:, :,
2]
# Apply QK-Norm if needed
txt_q = self.txt_attn_q_norm(txt_q).to(txt_q.dtype)
txt_k = self.txt_attn_k_norm(txt_k).to(txt_k.dtype)
# Run distributed attention
img_attn, txt_attn = self.attn(img_q, img_k, img_v, txt_q, txt_k, txt_v)
img_attn_out, _ = self.img_attn_proj(
img_attn.view(batch_size, image_seq_len, -1))
# Use fused operation for residual connection, normalization, and modulation
img_mlp_input, img_residual = self.img_attn_residual_mlp_norm(
img, img_attn_out, img_attn_gate, img_mlp_shift, img_mlp_scale)
# Process image MLP
img_mlp_out = self.img_mlp(img_mlp_input)
img = self.img_mlp_residual(img_residual, img_mlp_out, img_mlp_gate)
# Process text attention output
txt_attn_out, _ = self.txt_attn_proj(
txt_attn.reshape(batch_size, text_seq_len, -1))
# Use fused operation for residual connection, normalization, and modulation
txt_mlp_input, txt_residual = self.txt_attn_residual_mlp_norm(
txt, txt_attn_out, txt_attn_gate, txt_mlp_shift, txt_mlp_scale)
# Process text MLP
txt_mlp_out = self.txt_mlp(txt_mlp_input)
txt = self.txt_mlp_residual(txt_residual, txt_mlp_out, txt_mlp_gate)
return img, txt
class MMSingleStreamBlock(nn.Module):
"""
A DiT block with parallel linear layers using distributed attention
and tensor parallelism.
"""
def __init__(
self,
hidden_size: int,
num_attention_heads: int,
mlp_ratio: float = 4.0,
dtype: Optional[torch.dtype] = None,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
prefix: str = "",
):
super().__init__()
self.deterministic = False
self.hidden_size = hidden_size
self.num_attention_heads = num_attention_heads
head_dim = hidden_size // num_attention_heads
mlp_hidden_dim = int(hidden_size * mlp_ratio)
self.mlp_hidden_dim = mlp_hidden_dim
# Combined QKV and MLP input projection
self.linear1 = ReplicatedLinear(hidden_size,
hidden_size * 3 + mlp_hidden_dim,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.linear1")
# Combined projection and MLP output
self.linear2 = ReplicatedLinear(hidden_size + mlp_hidden_dim,
hidden_size,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.linear2")
# QK norm layers
self.q_norm = RMSNorm(head_dim, eps=1e-6)
self.k_norm = RMSNorm(head_dim, eps=1e-6)
# Fused operations with better naming
self.input_norm_scale_shift = LayerNormScaleShift(
hidden_size,
norm_type="layer",
eps=1e-6,
elementwise_affine=False,
dtype=dtype)
self.output_residual = ScaleResidual()
# Activation function
self.mlp_act = nn.GELU(approximate="tanh")
# Modulation
self.modulation = ModulateProjection(hidden_size,
factor=3,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.modulation")
# Distributed attention
self.attn = DistributedAttention(
num_heads=num_attention_heads,
head_size=head_dim,
causal=False,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn")
def forward(
self,
x: torch.Tensor,
vec: torch.Tensor,
txt_len: int,
freqs_cis_img: Tuple[torch.Tensor, torch.Tensor],
) -> torch.Tensor:
# Process modulation
mod_shift, mod_scale, mod_gate = self.modulation(vec).chunk(3, dim=-1)
# Apply pre-norm and modulation using fused operation
x_mod = self.input_norm_scale_shift(x, mod_shift, mod_scale)
# Get combined projections
linear1_out, _ = self.linear1(x_mod)
# Split into QKV and MLP parts
qkv, mlp = torch.split(linear1_out,
[3 * self.hidden_size, self.mlp_hidden_dim],
dim=-1)
# Process QKV
batch_size, seq_len = qkv.shape[0], qkv.shape[1]
qkv = qkv.view(batch_size, seq_len, 3, self.num_attention_heads, -1)
q, k, v = qkv[:, :, 0], qkv[:, :, 1], qkv[:, :, 2]
# Apply QK-Norm
q = self.q_norm(q).to(v.dtype)
k = self.k_norm(k).to(v.dtype)
# Split into image and text parts
img_q, txt_q = q[:, :-txt_len], q[:, -txt_len:]
img_k, txt_k = k[:, :-txt_len], k[:, -txt_len:]
img_v, txt_v = v[:, :-txt_len], v[:, -txt_len:]
# Apply rotary embeddings to image parts
cos, sin = freqs_cis_img
img_q, img_k = _apply_rotary_emb(
img_q, cos, sin,
is_neox_style=False), _apply_rotary_emb(img_k,
cos,
sin,
is_neox_style=False)
# Run distributed attention
img_attn_output, txt_attn_output = self.attn(img_q, img_k, img_v, txt_q,
txt_k, txt_v)
attn_output = torch.cat((img_attn_output, txt_attn_output),
dim=1).view(batch_size, seq_len, -1)
# Process MLP activation
mlp_output = self.mlp_act(mlp)
# Combine attention and MLP outputs
combined = torch.cat((attn_output, mlp_output), dim=-1)
# Final projection
output, _ = self.linear2(combined)
# Apply residual connection with gating using fused operation
return self.output_residual(x, output, mod_gate)
class FinalLayer(nn.Module):
"""
The final layer of DiT that projects features to pixel space.
"""
def __init__(self,
hidden_size,
patch_size,
out_channels,
dtype=None,
prefix: str = "") -> None:
super().__init__()
# Normalization
self.norm_final = nn.LayerNorm(hidden_size,
eps=1e-6,
elementwise_affine=False,
dtype=dtype)
output_dim = patch_size**3 * out_channels
self.linear = ReplicatedLinear(hidden_size,
output_dim,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.linear")
# Modulation
self.adaLN_modulation = ModulateProjection(
hidden_size,
factor=2,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.adaLN_modulation")
def forward(self, img, vec):
scale, shift = self.adaLN_modulation(vec).chunk(2, dim=-1)
img = self.norm_final(img) * (1.0 +
scale.unsqueeze(1)) + shift.unsqueeze(1)
img, _ = self.linear(img)
return img
class FluxTransformer2DModel(BaseDiT):
_fsdp_shard_conditions = FluxImageConfig()._fsdp_shard_conditions
_supported_attention_backends = FluxImageConfig(
)._supported_attention_backends
_param_names_mapping = FluxImageConfig()._param_names_mapping
def __init__(self, config: FluxImageConfig) -> None:
super().__init__(config=config)
self.hidden_size = config.hidden_size
self.num_attention_heads = config.num_attention_heads
self.num_channels_latents = config.num_channels_latents
self.text_states_dim = config.joint_attention_dim
self.text_states_dim_2 = config.pooled_projection_dim
self.rope_dim_list = list(config.axes_dims_rope)
self.rope_theta = config.rope_theta
self.out_channels = config.out_channels
self.patch_size = config.patch_size
self.img_in = ReplicatedLinear(config.in_channels,
self.hidden_size,
params_dtype=config.dtype,
prefix=f"{config.prefix}.img_in")
self.txt_in = ReplicatedLinear(self.text_states_dim,
self.hidden_size,
params_dtype=config.dtype,
prefix=f"{config.prefix}.txt_in")
self.time_in = TimestepEmbedder(self.hidden_size,
act_layer="silu",
dtype=config.dtype,
prefix=f"{config.prefix}.time_in")
self.txt2_in = MLP(self.text_states_dim_2,
self.hidden_size,
self.hidden_size,
act_type="silu",
dtype=config.dtype,
prefix=f"{config.prefix}.txt2_in")
self.guidance_in = (TimestepEmbedder(
self.hidden_size,
act_layer="silu",
dtype=config.dtype,
prefix=f"{config.prefix}.guidance_in")
if config.guidance_embeds else None)
# Double blocks
self.double_blocks = nn.ModuleList([
MMDoubleStreamBlock(
self.hidden_size,
self.num_attention_heads,
dtype=config.dtype,
supported_attention_backends=self._supported_attention_backends,
prefix=f"{config.prefix}.double_blocks.{i}")
for i in range(config.num_layers)
])
# Single blocks
self.single_blocks = nn.ModuleList([
MMSingleStreamBlock(
self.hidden_size,
self.num_attention_heads,
dtype=config.dtype,
supported_attention_backends=self._supported_attention_backends,
prefix=f"{config.prefix}.single_blocks.{i+config.num_layers}")
for i in range(config.num_single_layers)
])
self.final_layer = FinalLayer(config.hidden_size,
self.patch_size,
self.out_channels,
dtype=config.dtype,
prefix=f"{config.prefix}.final_layer")
self.__post_init__()
def forward(self,
hidden_states: torch.Tensor,
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
timestep: torch.LongTensor,
encoder_hidden_states_image: Optional[Union[
torch.Tensor, List[torch.Tensor]]] = None,
guidance=None,
**kwargs):
"""
Forward pass of the FluxTransformer2DModel.
Args:
hidden_states: Input image latents [B, N, C]
encoder_hidden_states: Text embeddings [B, L, D]
timestep: Diffusion timestep
guidance: Guidance scale for CFG
Returns:
Tuple of (output)
"""
h = kwargs.pop("height_latents") or None
w = kwargs.pop("width_latents") or None
assert h is not None and w is not None
img = x = hidden_states
# Match diffusers implementation by multiplying timestep by 1000
t = timestep.to(img.dtype)
# Split text embeddings - first token is global, rest are per-token
txt = encoder_hidden_states[1]
text_states_2 = encoder_hidden_states[0]
# Get spatial dimensions
# _, _, oh, ow = img.shape
th, tw = (h // self.patch_size // 2, w // self.patch_size // 2)
# Get rotary embeddings
freqs_cos_img, freqs_sin_img = get_rotary_pos_embed(
(1, th, tw),
self.hidden_size,
self.num_attention_heads,
self.rope_dim_list,
self.rope_theta,
shard_dim=1,
)
freqs_cos_img = freqs_cos_img.to(img.device)
freqs_sin_img = freqs_sin_img.to(img.device)
freqs_cis_img = (freqs_cos_img, freqs_sin_img)
# Prepare modulation vectors
vec = self.time_in(t)
# Add text modulation
vec = vec + self.txt2_in(text_states_2)
# Add guidance modulation
if self.guidance_in is not None and guidance is not None:
vec = vec + self.guidance_in(guidance)
# embed text and image
img, _ = self.img_in(img)
txt, _ = self.txt_in(txt)
img_seq_len = img.shape[1]
txt_seq_len = txt.shape[1]
# Process through double stream blocks
for index, block in enumerate(self.double_blocks):
double_block_args = [img, txt, vec, freqs_cis_img]
img, txt = block(*double_block_args)
# Merge txt and img to pass through single stream blocks
x = torch.cat((img, txt), 1)
# Process through single stream blocks
for index, block in enumerate(self.single_blocks):
single_block_args = [x, vec, txt_seq_len, freqs_cis_img]
x = block(*single_block_args)
# Extract image features
img = x[:, :img_seq_len, ...]
# Final layer
img = self.final_layer(img, vec)
return img
@@ -371,7 +371,6 @@ class TransformerLoader(ComponentLoader):
raise ValueError(
"Model config does not contain a _class_name attribute. "
"Only diffusers format is supported.")
config.pop("_diffusers_version")
# Config from Diffusers supersedes fastvideo's model config
dit_config = fastvideo_args.dit_config
+4 -1
View File
@@ -23,6 +23,7 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
"HunyuanVideoTransformer3DModel":
("dits", "hunyuanvideo", "HunyuanVideoTransformer3DModel"),
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
"FluxTransformer2DModel": ("dits", "flux", "FluxTransformer2DModel"),
}
_IMAGE_TO_VIDEO_DIT_MODELS = {
@@ -34,6 +35,7 @@ _TEXT_ENCODER_MODELS = {
"CLIPTextModel": ("encoders", "clip", "CLIPTextModel"),
"LlamaModel": ("encoders", "llama", "LlamaModel"),
"UMT5EncoderModel": ("encoders", "t5", "UMT5EncoderModel"),
"T5EncoderModel": ("encoders", "t5", "T5EncoderModel"),
}
_IMAGE_ENCODER_MODELS: dict[str, tuple] = {
@@ -45,12 +47,13 @@ _VAE_MODELS = {
"AutoencoderKLHunyuanVideo":
("vaes", "hunyuanvae", "AutoencoderKLHunyuanVideo"),
"AutoencoderKLWan": ("vaes", "wanvae", "AutoencoderKLWan"),
"AutoencoderKL": ("vaes", "image_vae", "AutoencoderKL"),
}
_SCHEDULERS = {
"FlowMatchEulerDiscreteScheduler":
("schedulers", "scheduling_flow_match_euler_discrete",
"FlowMatchDiscreteScheduler"),
"FlowMatchEulerDiscreteScheduler"),
"UniPCMultistepScheduler":
("schedulers", "scheduling_unipc_multistep", "UniPCMultistepScheduler"),
}
@@ -1,3 +1,4 @@
# type: ignore
# SPDX-License-Identifier: Apache-2.0
# Copyright 2024 Stability AI, Katherine Crowson and The HuggingFace Team. All rights reserved.
@@ -19,13 +20,16 @@
#
# ==============================================================================
import math
from dataclasses import dataclass
from typing import Any, Optional, Tuple, Union
from typing import List, Optional, Tuple, Union
import numpy as np
import scipy
import torch
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.schedulers.scheduling_utils import SchedulerMixin
from diffusers.utils import BaseOutput, logging
from diffusers.utils import BaseOutput, is_scipy_available, logging
from fastvideo.v1.models.schedulers.base import BaseScheduler
@@ -33,7 +37,7 @@ logger = logging.get_logger(__name__) # pylint: disable=invalid-name
@dataclass
class FlowMatchDiscreteSchedulerOutput(BaseOutput):
class FlowMatchEulerDiscreteSchedulerOutput(BaseOutput):
"""
Output class for the scheduler's `step` function output.
@@ -46,7 +50,8 @@ class FlowMatchDiscreteSchedulerOutput(BaseOutput):
prev_sample: torch.FloatTensor
class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
BaseScheduler):
"""
Euler scheduler.
@@ -56,16 +61,37 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
Args:
num_train_timesteps (`int`, defaults to 1000):
The number of diffusion steps to train the model.
timestep_spacing (`str`, defaults to `"linspace"`):
The way the timesteps should be scaled. Refer to Table 2 of the [Common Diffusion Noise Schedules and
Sample Steps are Flawed](https://huggingface.co/papers/2305.08891) for more information.
shift (`float`, defaults to 1.0):
The shift value for the timestep schedule.
reverse (`bool`, defaults to `True`):
Whether to reverse the timestep schedule.
use_dynamic_shifting (`bool`, defaults to False):
Whether to apply timestep shifting on-the-fly based on the image resolution.
base_shift (`float`, defaults to 0.5):
Value to stabilize image generation. Increasing `base_shift` reduces variation and image is more consistent
with desired output.
max_shift (`float`, defaults to 1.15):
Value change allowed to latent vectors. Increasing `max_shift` encourages more variation and image may be
more exaggerated or stylized.
base_image_seq_len (`int`, defaults to 256):
The base image sequence length.
max_image_seq_len (`int`, defaults to 4096):
The maximum image sequence length.
invert_sigmas (`bool`, defaults to False):
Whether to invert the sigmas.
shift_terminal (`float`, defaults to None):
The end value of the shifted timestep schedule.
use_karras_sigmas (`bool`, defaults to False):
Whether to use Karras sigmas for step sizes in the noise schedule during sampling.
use_exponential_sigmas (`bool`, defaults to False):
Whether to use exponential sigmas for step sizes in the noise schedule during sampling.
use_beta_sigmas (`bool`, defaults to False):
Whether to use beta sigmas for step sizes in the noise schedule during sampling.
time_shift_type (`str`, defaults to "exponential"):
The type of dynamic resolution-dependent timestep shifting to apply. Either "exponential" or "linear".
stochastic_sampling (`bool`, defaults to False):
Whether to use stochastic sampling.
"""
_compatibles: list[Any] = []
_compatibles = []
order = 1
@register_to_config
@@ -73,31 +99,62 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
self,
num_train_timesteps: int = 1000,
shift: float = 1.0,
reverse: bool = True,
solver: str = "euler",
n_tokens: Optional[int] = None,
**kwargs,
use_dynamic_shifting: bool = False,
base_shift: Optional[float] = 0.5,
max_shift: Optional[float] = 1.15,
base_image_seq_len: Optional[int] = 256,
max_image_seq_len: Optional[int] = 4096,
invert_sigmas: bool = False,
shift_terminal: Optional[float] = None,
use_karras_sigmas: Optional[bool] = False,
use_exponential_sigmas: Optional[bool] = False,
use_beta_sigmas: Optional[bool] = False,
time_shift_type: str = "exponential",
stochastic_sampling: bool = False,
):
sigmas = torch.linspace(1, 0, num_train_timesteps + 1)
if not reverse:
sigmas = sigmas.flip(0)
self.sigmas = sigmas
# the value fed to model
self.timesteps = (sigmas[:-1] *
num_train_timesteps).to(dtype=torch.float32)
self._step_index: int | None = None
self._begin_index = 0
self.supported_solver = ["euler"]
if solver not in self.supported_solver:
if self.config.use_beta_sigmas and not is_scipy_available():
raise ImportError(
"Make sure to install scipy if you want to use beta sigmas.")
if sum([
self.config.use_beta_sigmas, self.config.use_exponential_sigmas,
self.config.use_karras_sigmas
]) > 1:
raise ValueError(
f"Solver {solver} not supported. Supported solvers: {self.supported_solver}"
"Only one of `config.use_beta_sigmas`, `config.use_exponential_sigmas`, `config.use_karras_sigmas` can be used."
)
if time_shift_type not in {"exponential", "linear"}:
raise ValueError(
"`time_shift_type` must either be 'exponential' or 'linear'.")
BaseScheduler.__init__(self)
timesteps = np.linspace(1,
num_train_timesteps,
num_train_timesteps,
dtype=np.float32)[::-1].copy()
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32)
sigmas = timesteps / num_train_timesteps
if not use_dynamic_shifting:
# when use_dynamic_shifting is True, we apply the timestep shifting on the fly based on the image resolution
sigmas = shift * sigmas / (1 + (shift - 1) * sigmas)
self.timesteps = sigmas * num_train_timesteps
self._step_index = None
self._begin_index = None
self._shift = shift
self.sigmas = sigmas.to(
"cpu") # to avoid too much CPU/GPU communication
self.sigma_min = self.sigmas[-1].item()
self.sigma_max = self.sigmas[0].item()
@property
def shift(self):
"""
The value used for shifting.
"""
return self._shift
@property
def step_index(self):
@@ -124,42 +181,207 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
"""
self._begin_index = begin_index
def set_shift(self, shift: float):
self._shift = shift
def scale_noise(
self,
sample: torch.FloatTensor,
timestep: Union[float, torch.FloatTensor],
noise: Optional[torch.FloatTensor] = None,
) -> torch.FloatTensor:
"""
Forward process in flow-matching
Args:
sample (`torch.FloatTensor`):
The input sample.
timestep (`int`, *optional*):
The current timestep in the diffusion chain.
Returns:
`torch.FloatTensor`:
A scaled input sample.
"""
# Make sure sigmas and timesteps have the same device and dtype as original_samples
sigmas = self.sigmas.to(device=sample.device, dtype=sample.dtype)
if sample.device.type == "mps" and torch.is_floating_point(timestep):
# mps does not support float64
schedule_timesteps = self.timesteps.to(sample.device,
dtype=torch.float32)
timestep = timestep.to(sample.device, dtype=torch.float32)
else:
schedule_timesteps = self.timesteps.to(sample.device)
timestep = timestep.to(sample.device)
# self.begin_index is None when scheduler is used for training, or pipeline does not implement set_begin_index
if self.begin_index is None:
step_indices = [
self.index_for_timestep(t, schedule_timesteps) for t in timestep
]
elif self.step_index is not None:
# add_noise is called after first denoising step (for inpainting)
step_indices = [self.step_index] * timestep.shape[0]
else:
# add noise is called before first denoising step to create initial latent(img2img)
step_indices = [self.begin_index] * timestep.shape[0]
sigma = sigmas[step_indices].flatten()
while len(sigma.shape) < len(sample.shape):
sigma = sigma.unsqueeze(-1)
sample = sigma * noise + (1.0 - sigma) * sample
return sample
def _sigma_to_t(self, sigma):
return sigma * self.config.num_train_timesteps
def time_shift(self, mu: float, sigma: float, t: torch.Tensor):
if self.config.time_shift_type == "exponential":
return self._time_shift_exponential(mu, sigma, t)
elif self.config.time_shift_type == "linear":
return self._time_shift_linear(mu, sigma, t)
def stretch_shift_to_terminal(self, t: torch.Tensor) -> torch.Tensor:
r"""
Stretches and shifts the timestep schedule to ensure it terminates at the configured `shift_terminal` config
value.
Reference:
https://github.com/Lightricks/LTX-Video/blob/a01a171f8fe3d99dce2728d60a73fecf4d4238ae/ltx_video/schedulers/rf.py#L51
Args:
t (`torch.Tensor`):
A tensor of timesteps to be stretched and shifted.
Returns:
`torch.Tensor`:
A tensor of adjusted timesteps such that the final value equals `self.config.shift_terminal`.
"""
one_minus_z = 1 - t
scale_factor = one_minus_z[-1] / (1 - self.config.shift_terminal)
stretched_t = 1 - (one_minus_z / scale_factor)
return stretched_t
def set_timesteps(
self,
num_inference_steps: int,
num_inference_steps: Optional[int] = None,
device: Union[str, torch.device] = None,
n_tokens: int = 0,
sigmas: Optional[List[float]] = None,
mu: Optional[float] = None,
timesteps: Optional[List[float]] = None,
**kwargs,
):
"""
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
Args:
num_inference_steps (`int`):
num_inference_steps (`int`, *optional*):
The number of diffusion steps used when generating samples with a pre-trained model.
device (`str` or `torch.device`, *optional*):
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
n_tokens (`int`, *optional*):
Number of tokens in the input sequence.
sigmas (`List[float]`, *optional*):
Custom values for sigmas to be used for each diffusion step. If `None`, the sigmas are computed
automatically.
mu (`float`, *optional*):
Determines the amount of shifting applied to sigmas when performing resolution-dependent timestep
shifting.
timesteps (`List[float]`, *optional*):
Custom values for timesteps to be used for each diffusion step. If `None`, the timesteps are computed
automatically.
"""
if self.config.use_dynamic_shifting and mu is None:
raise ValueError(
"`mu` must be passed when `use_dynamic_shifting` is set to be `True`"
)
if sigmas is not None and timesteps is not None and len(sigmas) != len(
timesteps):
raise ValueError(
"`sigmas` and `timesteps` should have the same length")
if num_inference_steps is not None:
if (sigmas is not None and len(sigmas) != num_inference_steps) or (
timesteps is not None
and len(timesteps) != num_inference_steps):
raise ValueError(
"`sigmas` and `timesteps` should have the same length as num_inference_steps, if `num_inference_steps` is provided"
)
else:
num_inference_steps = len(sigmas) if sigmas is not None else len(
timesteps)
self.num_inference_steps = num_inference_steps
sigmas = torch.linspace(1, 0, num_inference_steps + 1)
sigmas = self.sd3_time_shift(sigmas)
# 1. Prepare default sigmas
is_timesteps_provided = timesteps is not None
if not self.config.reverse:
sigmas = 1 - sigmas
if is_timesteps_provided:
timesteps = np.array(timesteps).astype(np.float32)
if sigmas is None:
if timesteps is None:
timesteps = np.linspace(self._sigma_to_t(self.sigma_max),
self._sigma_to_t(self.sigma_min),
num_inference_steps)
sigmas = timesteps / self.config.num_train_timesteps
else:
sigmas = np.array(sigmas).astype(np.float32)
num_inference_steps = len(sigmas)
# 2. Perform timestep shifting. Either no shifting is applied, or resolution-dependent shifting of
# "exponential" or "linear" type is applied
if self.config.use_dynamic_shifting:
sigmas = self.time_shift(mu, 1.0, sigmas)
else:
sigmas = self.shift * sigmas / (1 + (self.shift - 1) * sigmas)
# 3. If required, stretch the sigmas schedule to terminate at the configured `shift_terminal` value
if self.config.shift_terminal:
sigmas = self.stretch_shift_to_terminal(sigmas)
# 4. If required, convert sigmas to one of karras, exponential, or beta sigma schedules
if self.config.use_karras_sigmas:
sigmas = self._convert_to_karras(
in_sigmas=sigmas, num_inference_steps=num_inference_steps)
elif self.config.use_exponential_sigmas:
sigmas = self._convert_to_exponential(
in_sigmas=sigmas, num_inference_steps=num_inference_steps)
elif self.config.use_beta_sigmas:
sigmas = self._convert_to_beta(
in_sigmas=sigmas, num_inference_steps=num_inference_steps)
# 5. Convert sigmas and timesteps to tensors and move to specified device
sigmas = torch.from_numpy(sigmas).to(dtype=torch.float32, device=device)
if not is_timesteps_provided:
timesteps = sigmas * self.config.num_train_timesteps
else:
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32,
device=device)
# 6. Append the terminal sigma value.
# If a model requires inverted sigma schedule for denoising but timesteps without inversion, the
# `invert_sigmas` flag can be set to `True`. This case is only required in Mochi
if self.config.invert_sigmas:
sigmas = 1.0 - sigmas
timesteps = sigmas * self.config.num_train_timesteps
sigmas = torch.cat([sigmas, torch.ones(1, device=sigmas.device)])
else:
sigmas = torch.cat([sigmas, torch.zeros(1, device=sigmas.device)])
self.timesteps = timesteps
self.sigmas = sigmas
self.timesteps = (sigmas[:-1] * self.config.num_train_timesteps).to(
dtype=torch.float32, device=device)
# Reset step index
self._step_index = None
self._begin_index = None
def index_for_timestep(self, timestep, schedule_timesteps=None) -> int:
def scale_model_input(self,
sample: torch.Tensor,
timestep: Optional[int] = None) -> torch.Tensor:
return sample
def index_for_timestep(self, timestep, schedule_timesteps=None):
if schedule_timesteps is None:
schedule_timesteps = self.timesteps
@@ -171,14 +393,9 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
# case we start in the middle of the denoising schedule (e.g. for image-to-image)
pos = 1 if len(indices) > 1 else 0
idx: int = indices[pos].item()
return indices[pos].item()
return idx
def set_shift(self, shift: float) -> None:
self.config.shift = shift
def _init_step_index(self, timestep) -> None:
def _init_step_index(self, timestep):
if self.begin_index is None:
if isinstance(timestep, torch.Tensor):
timestep = timestep.to(self.timesteps.device)
@@ -186,22 +403,19 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
else:
self._step_index = self._begin_index
def scale_model_input(self,
sample: torch.Tensor,
timestep: Optional[int] = None) -> torch.Tensor:
return sample
def sd3_time_shift(self, t: torch.Tensor):
return (self.config.shift * t) / (1 + (self.config.shift - 1) * t)
def step(
self,
model_output: torch.FloatTensor,
timestep: Union[float, torch.FloatTensor],
sample: torch.FloatTensor,
s_churn: float = 0.0,
s_tmin: float = 0.0,
s_tmax: float = float("inf"),
s_noise: float = 1.0,
generator: Optional[torch.Generator] = None,
per_token_timesteps: Optional[torch.Tensor] = None,
return_dict: bool = True,
**kwargs,
) -> Union[FlowMatchDiscreteSchedulerOutput, Tuple]:
) -> Union[FlowMatchEulerDiscreteSchedulerOutput, Tuple]:
"""
Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion
process from the learned model outputs (most often the predicted noise).
@@ -213,24 +427,30 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
The current discrete timestep in the diffusion chain.
sample (`torch.FloatTensor`):
A current instance of a sample created by the diffusion process.
s_churn (`float`):
s_tmin (`float`):
s_tmax (`float`):
s_noise (`float`, defaults to 1.0):
Scaling factor for noise added to the sample.
generator (`torch.Generator`, *optional*):
A random number generator.
n_tokens (`int`, *optional*):
Number of tokens in the input sequence.
per_token_timesteps (`torch.Tensor`, *optional*):
The timesteps for each token in the sample.
return_dict (`bool`):
Whether or not to return a [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or
tuple.
Whether or not to return a
[`~schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteSchedulerOutput`] or tuple.
Returns:
[`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or `tuple`:
If return_dict is `True`, [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] is
returned, otherwise a tuple is returned where the first element is the sample tensor.
[`~schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteSchedulerOutput`] or `tuple`:
If return_dict is `True`,
[`~schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteSchedulerOutput`] is returned,
otherwise a tuple is returned where the first element is the sample tensor.
"""
if isinstance(timestep, (int, torch.IntTensor, torch.LongTensor)):
raise ValueError((
"Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
" `FlowMatchEulerDiscreteScheduler.step()` is not supported. Make sure to pass"
" one of the `scheduler.timesteps` as a timestep."), )
if self.step_index is None:
@@ -239,24 +459,132 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
# Upcast to avoid precision issues when computing prev_sample
sample = sample.to(torch.float32)
assert self.step_index is not None
dt = self.sigmas[self.step_index + 1] - self.sigmas[self.step_index]
if per_token_timesteps is not None:
per_token_sigmas = per_token_timesteps / self.config.num_train_timesteps
if self.config.solver == "euler":
prev_sample = sample + model_output.to(torch.float32) * dt
sigmas = self.sigmas[:, None, None]
lower_mask = sigmas < per_token_sigmas[None] - 1e-6
lower_sigmas = lower_mask * sigmas
lower_sigmas, _ = lower_sigmas.max(dim=0)
current_sigma = per_token_sigmas[..., None]
next_sigma = lower_sigmas[..., None]
dt = current_sigma - next_sigma
else:
raise ValueError(
f"Solver {self.config.solver} not supported. Supported solvers: {self.supported_solver}"
)
sigma_idx = self.step_index
sigma = self.sigmas[sigma_idx]
sigma_next = self.sigmas[sigma_idx + 1]
current_sigma = sigma
next_sigma = sigma_next
dt = sigma_next - sigma
if self.config.stochastic_sampling:
x0 = sample - current_sigma * model_output
noise = torch.randn_like(sample)
prev_sample = (1.0 - next_sigma) * x0 + next_sigma * noise
else:
prev_sample = sample + dt * model_output
# upon completion increase step index by one
assert self._step_index is not None
self._step_index += 1
if per_token_timesteps is None:
# Cast sample back to model compatible dtype
prev_sample = prev_sample.to(model_output.dtype)
if not return_dict:
return (prev_sample, )
return FlowMatchDiscreteSchedulerOutput(prev_sample=prev_sample)
return FlowMatchEulerDiscreteSchedulerOutput(prev_sample=prev_sample)
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_karras
def _convert_to_karras(self, in_sigmas: torch.Tensor,
num_inference_steps) -> torch.Tensor:
"""Constructs the noise schedule of Karras et al. (2022)."""
# Hack to make sure that other schedulers which copy this function don't break
# TODO: Add this logic to the other schedulers
if hasattr(self.config, "sigma_min"):
sigma_min = self.config.sigma_min
else:
sigma_min = None
if hasattr(self.config, "sigma_max"):
sigma_max = self.config.sigma_max
else:
sigma_max = None
sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item()
sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item()
rho = 7.0 # 7.0 is the value used in the paper
ramp = np.linspace(0, 1, num_inference_steps)
min_inv_rho = sigma_min**(1 / rho)
max_inv_rho = sigma_max**(1 / rho)
sigmas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho))**rho
return sigmas
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_exponential
def _convert_to_exponential(self, in_sigmas: torch.Tensor,
num_inference_steps: int) -> torch.Tensor:
"""Constructs an exponential noise schedule."""
# Hack to make sure that other schedulers which copy this function don't break
# TODO: Add this logic to the other schedulers
if hasattr(self.config, "sigma_min"):
sigma_min = self.config.sigma_min
else:
sigma_min = None
if hasattr(self.config, "sigma_max"):
sigma_max = self.config.sigma_max
else:
sigma_max = None
sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item()
sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item()
sigmas = np.exp(
np.linspace(math.log(sigma_max), math.log(sigma_min),
num_inference_steps))
return sigmas
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_beta
def _convert_to_beta(self,
in_sigmas: torch.Tensor,
num_inference_steps: int,
alpha: float = 0.6,
beta: float = 0.6) -> torch.Tensor:
"""From "Beta Sampling is All You Need" [arXiv:2407.12173] (Lee et. al, 2024)"""
# Hack to make sure that other schedulers which copy this function don't break
# TODO: Add this logic to the other schedulers
if hasattr(self.config, "sigma_min"):
sigma_min = self.config.sigma_min
else:
sigma_min = None
if hasattr(self.config, "sigma_max"):
sigma_max = self.config.sigma_max
else:
sigma_max = None
sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item()
sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item()
sigmas = np.array([
sigma_min + (ppf * (sigma_max - sigma_min)) for ppf in [
scipy.stats.beta.ppf(timestep, alpha, beta)
for timestep in 1 - np.linspace(0, 1, num_inference_steps)
]
])
return sigmas
def _time_shift_exponential(self, mu, sigma, t):
return math.exp(mu) / (math.exp(mu) + (1 / t - 1)**sigma)
def _time_shift_linear(self, mu, sigma, t):
return mu / (mu + (1 / t - 1)**sigma)
def __len__(self):
return self.config.num_train_timesteps
+556
View File
@@ -0,0 +1,556 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from diffusers
# Copyright 2024 The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# Copyright 2024 The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from typing import Dict, Optional, Union
import torch
import torch.nn as nn
from diffusers.models.attention_processor import (ADDED_KV_ATTENTION_PROCESSORS,
CROSS_ATTENTION_PROCESSORS,
Attention, AttentionProcessor,
AttnAddedKVProcessor,
AttnProcessor,
FusedAttnProcessor2_0)
from diffusers.models.autoencoders.vae import Decoder, Encoder
from fastvideo.v1.configs.models.vaes import ImageVAEConfig
from fastvideo.v1.models.vaes.common import DiagonalGaussianDistribution
class AutoencoderKL(nn.Module):
r"""
A VAE model with KL loss for encoding images into latents and decoding latent representations into images.
This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented
for all models (such as downloading or saving).
Parameters:
in_channels (int, *optional*, defaults to 3): Number of channels in the input image.
out_channels (int, *optional*, defaults to 3): Number of channels in the output.
down_block_types (`Tuple[str]`, *optional*, defaults to `("DownEncoderBlock2D",)`):
Tuple of downsample block types.
up_block_types (`Tuple[str]`, *optional*, defaults to `("UpDecoderBlock2D",)`):
Tuple of upsample block types.
block_out_channels (`Tuple[int]`, *optional*, defaults to `(64,)`):
Tuple of block output channels.
act_fn (`str`, *optional*, defaults to `"silu"`): The activation function to use.
latent_channels (`int`, *optional*, defaults to 4): Number of channels in the latent space.
sample_size (`int`, *optional*, defaults to `32`): Sample input size.
scaling_factor (`float`, *optional*, defaults to 0.18215):
The component-wise standard deviation of the trained latent space computed using the first batch of the
training set. This is used to scale the latent space to have unit variance when training the diffusion
model. The latents are scaled with the formula `z = z * scaling_factor` before being passed to the
diffusion model. When decoding, the latents are scaled back to the original scale with the formula: `z = 1
/ scaling_factor * z`. For more details, refer to sections 4.3.2 and D.1 of the [High-Resolution Image
Synthesis with Latent Diffusion Models](https://arxiv.org/abs/2112.10752) paper.
force_upcast (`bool`, *optional*, default to `True`):
If enabled it will force the VAE to run in float32 for high image resolution pipelines, such as SD-XL. VAE
can be fine-tuned / trained to a lower range without losing too much precision in which case
`force_upcast` can be set to `False` - see: https://huggingface.co/madebyollin/sdxl-vae-fp16-fix
mid_block_add_attention (`bool`, *optional*, default to `True`):
If enabled, the mid_block of the Encoder and Decoder will have attention blocks. If set to false, the
mid_block will only have resnet blocks
"""
_supports_gradient_checkpointing = True
_no_split_modules = ["BasicTransformerBlock", "ResnetBlock2D"]
def __init__(
self,
config: ImageVAEConfig,
):
nn.Module.__init__(self)
self.shift_factor = config.shift_factor
self.scaling_factor = config.scaling_factor
if config.load_encoder:
# pass init params to Encoder
self.encoder = Encoder(
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,
act_fn=config.act_fn,
norm_num_groups=config.norm_num_groups,
double_z=True,
mid_block_add_attention=config.mid_block_add_attention,
)
self.quant_conv = nn.Conv2d(2 * config.latent_channels, 2 *
config.latent_channels,
1) if config.use_quant_conv else None
if config.load_decoder:
# pass init params to Decoder
self.decoder = Decoder(
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,
mid_block_add_attention=config.mid_block_add_attention,
)
self.post_quant_conv = nn.Conv2d(
config.latent_channels, config.latent_channels,
1) if config.use_post_quant_conv else None
self.use_slicing = False
self.use_tiling = False
# only relevant if vae tiling is enabled
self.tile_sample_min_size = config.sample_size
sample_size = (config.sample_size[0] if isinstance(
config.sample_size, (list, tuple)) else config.sample_size)
self.tile_latent_min_size = int(
sample_size / (2**(len(config.block_out_channels) - 1)))
self.tile_overlap_factor = 0.25
def enable_tiling(self, use_tiling: bool = True):
r"""
Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to
compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow
processing larger images.
"""
self.use_tiling = use_tiling
def disable_tiling(self):
r"""
Disable tiled VAE decoding. If `enable_tiling` was previously enabled, this method will go back to computing
decoding in one step.
"""
self.enable_tiling(False)
def enable_slicing(self):
r"""
Enable sliced VAE decoding. When this option is enabled, the VAE will split the input tensor in slices to
compute decoding in several steps. This is useful to save some memory and allow larger batch sizes.
"""
self.use_slicing = True
def disable_slicing(self):
r"""
Disable sliced VAE decoding. If `enable_slicing` was previously enabled, this method will go back to computing
decoding in one step.
"""
self.use_slicing = False
@property
# Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.attn_processors
def attn_processors(self) -> Dict[str, AttentionProcessor]:
r"""
Returns:
`dict` of attention processors: A dictionary containing all attention processors used in the model with
indexed by its weight name.
"""
# set recursively
processors: dict[str, AttentionProcessor] = {}
def fn_recursive_add_processors(name: str, module: torch.nn.Module,
processors: Dict[str,
AttentionProcessor]):
if hasattr(module, "get_processor"):
processors[f"{name}.processor"] = module.get_processor()
for sub_name, child in module.named_children():
fn_recursive_add_processors(f"{name}.{sub_name}", child,
processors)
return processors
for name, module in self.named_children():
fn_recursive_add_processors(name, module, processors)
return processors
# Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_attn_processor
def set_attn_processor(self, processor: Union[AttentionProcessor,
Dict[str,
AttentionProcessor]]):
r"""
Sets the attention processor to use to compute attention.
Parameters:
processor (`dict` of `AttentionProcessor` or only `AttentionProcessor`):
The instantiated processor class or a dictionary of processor classes that will be set as the processor
for **all** `Attention` layers.
If `processor` is a dict, the key needs to define the path to the corresponding cross attention
processor. This is strongly recommended when setting trainable attention processors.
"""
count = len(self.attn_processors.keys())
if isinstance(processor, dict) and len(processor) != count:
raise ValueError(
f"A dict of processors was passed, but the number of processors {len(processor)} does not match the"
f" number of attention layers: {count}. Please make sure to pass {count} processor classes."
)
def fn_recursive_attn_processor(name: str, module: torch.nn.Module,
processor):
if hasattr(module, "set_processor"):
if not isinstance(processor, dict):
module.set_processor(processor)
else:
module.set_processor(processor.pop(f"{name}.processor"))
for sub_name, child in module.named_children():
fn_recursive_attn_processor(f"{name}.{sub_name}", child,
processor)
for name, module in self.named_children():
fn_recursive_attn_processor(name, module, processor)
# Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_default_attn_processor
def set_default_attn_processor(self):
"""
Disables custom attention processors and sets the default attention implementation.
"""
if all(proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS
for proc in self.attn_processors.values()):
processor = AttnAddedKVProcessor()
elif all(proc.__class__ in CROSS_ATTENTION_PROCESSORS
for proc in self.attn_processors.values()):
processor = AttnProcessor()
else:
raise ValueError(
f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}"
)
self.set_attn_processor(processor)
def _encode(self, x: torch.Tensor) -> torch.Tensor:
batch_size, num_channels, height, width = x.shape
if self.use_tiling and (width > self.tile_sample_min_size
or height > self.tile_sample_min_size):
return self._tiled_encode(x)
enc = self.encoder(x)
if self.quant_conv is not None:
enc = self.quant_conv(enc)
return enc
def encode(self, x: torch.Tensor) -> torch.Tensor:
"""
Encode a batch of images into latents.
Args:
x (`torch.Tensor`): Input batch of images.
return_dict (`bool`, *optional*, defaults to `True`):
Whether to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple.
Returns:
The latent representations of the encoded images. If `return_dict` is True, a
[`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is returned.
"""
if self.use_slicing and x.shape[0] > 1:
encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)]
enc = torch.cat(encoded_slices)
else:
enc = self._encode(x)
enc = DiagonalGaussianDistribution(enc)
return enc
def _decode(self, z: torch.Tensor) -> torch.Tensor:
if self.use_tiling and (z.shape[-1] > self.tile_latent_min_size
or z.shape[-2] > self.tile_latent_min_size):
return self.tiled_decode(z)
if self.post_quant_conv is not None:
z = self.post_quant_conv(z)
dec = self.decoder(z)
return dec
def decode(self, z: torch.Tensor) -> torch.Tensor:
"""
Decode a batch of images.
Args:
z (`torch.Tensor`): Input batch of latent vectors.
return_dict (`bool`, *optional*, defaults to `True`):
Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple.
Returns:
[`~models.vae.DecoderOutput`] or `tuple`:
If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is
returned.
"""
if self.use_slicing and z.shape[0] > 1:
decoded_slices = [
self._decode(z_slice).sample for z_slice in z.split(1)
]
decoded = torch.cat(decoded_slices)
else:
decoded = self._decode(z)
return decoded
def blend_v(self, a: torch.Tensor, b: torch.Tensor,
blend_extent: int) -> torch.Tensor:
blend_extent = min(a.shape[2], b.shape[2], blend_extent)
for y in range(blend_extent):
b[:, :, y, :] = a[:, :, -blend_extent + y, :] * (
1 - y / blend_extent) + b[:, :, y, :] * (y / blend_extent)
return b
def blend_h(self, a: torch.Tensor, b: torch.Tensor,
blend_extent: int) -> torch.Tensor:
blend_extent = min(a.shape[3], b.shape[3], blend_extent)
for x in range(blend_extent):
b[:, :, :, x] = a[:, :, :, -blend_extent + x] * (
1 - x / blend_extent) + b[:, :, :, x] * (x / blend_extent)
return b
def _tiled_encode(self, x: torch.Tensor) -> torch.Tensor:
r"""Encode a batch of images using a tiled encoder.
When this option is enabled, the VAE will split the input tensor into tiles to compute encoding in several
steps. This is useful to keep memory use constant regardless of image size. The end result of tiled encoding is
different from non-tiled encoding because each tile uses a different encoder. To avoid tiling artifacts, the
tiles overlap and are blended together to form a smooth output. You may still see tile-sized changes in the
output, but they should be much less noticeable.
Args:
x (`torch.Tensor`): Input batch of images.
Returns:
`torch.Tensor`:
The latent representation of the encoded videos.
"""
overlap_size = int(self.tile_sample_min_size *
(1 - self.tile_overlap_factor))
blend_extent = int(self.tile_latent_min_size * self.tile_overlap_factor)
row_limit = self.tile_latent_min_size - blend_extent
# Split the image into 512x512 tiles and encode them separately.
rows = []
for i in range(0, x.shape[2], overlap_size):
row = []
for j in range(0, x.shape[3], overlap_size):
tile = x[:, :, i:i + self.tile_sample_min_size,
j:j + self.tile_sample_min_size]
tile = self.encoder(tile)
if self.quant_conv:
tile = self.quant_conv(tile)
row.append(tile)
rows.append(row)
result_rows = []
for i, row in enumerate(rows):
result_row = []
for j, tile in enumerate(row):
# blend the above tile and the left tile
# to the current tile and add the current tile to the result row
if i > 0:
tile = self.blend_v(rows[i - 1][j], tile, blend_extent)
if j > 0:
tile = self.blend_h(row[j - 1], tile, blend_extent)
result_row.append(tile[:, :, :row_limit, :row_limit])
result_rows.append(torch.cat(result_row, dim=3))
enc = torch.cat(result_rows, dim=2)
return enc
def tiled_encode(self, x: torch.Tensor) -> torch.Tensor:
r"""Encode a batch of images using a tiled encoder.
When this option is enabled, the VAE will split the input tensor into tiles to compute encoding in several
steps. This is useful to keep memory use constant regardless of image size. The end result of tiled encoding is
different from non-tiled encoding because each tile uses a different encoder. To avoid tiling artifacts, the
tiles overlap and are blended together to form a smooth output. You may still see tile-sized changes in the
output, but they should be much less noticeable.
Args:
x (`torch.Tensor`): Input batch of images.
return_dict (`bool`, *optional*, defaults to `True`):
Whether or not to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple.
Returns:
[`~models.autoencoder_kl.AutoencoderKLOutput`] or `tuple`:
If return_dict is True, a [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain
`tuple` is returned.
"""
overlap_size = int(self.tile_sample_min_size *
(1 - self.tile_overlap_factor))
blend_extent = int(self.tile_latent_min_size * self.tile_overlap_factor)
row_limit = self.tile_latent_min_size - blend_extent
# Split the image into 512x512 tiles and encode them separately.
rows = []
for i in range(0, x.shape[2], overlap_size):
row = []
for j in range(0, x.shape[3], overlap_size):
tile = x[:, :, i:i + self.tile_sample_min_size,
j:j + self.tile_sample_min_size]
tile = self.encoder(tile)
if self.quant_conv:
tile = self.quant_conv(tile)
row.append(tile)
rows.append(row)
result_rows = []
for i, row in enumerate(rows):
result_row = []
for j, tile in enumerate(row):
# blend the above tile and the left tile
# to the current tile and add the current tile to the result row
if i > 0:
tile = self.blend_v(rows[i - 1][j], tile, blend_extent)
if j > 0:
tile = self.blend_h(row[j - 1], tile, blend_extent)
result_row.append(tile[:, :, :row_limit, :row_limit])
result_rows.append(torch.cat(result_row, dim=3))
moments = torch.cat(result_rows, dim=2)
enc = DiagonalGaussianDistribution(moments)
return enc
def tiled_decode(self, z: torch.Tensor) -> torch.Tensor:
r"""
Decode a batch of images using a tiled decoder.
Args:
z (`torch.Tensor`): Input batch of latent vectors.
return_dict (`bool`, *optional*, defaults to `True`):
Whether or not to return a [`~models.vae.DecoderOutput`] instead of a plain tuple.
Returns:
[`~models.vae.DecoderOutput`] or `tuple`:
If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is
returned.
"""
overlap_size = int(self.tile_latent_min_size *
(1 - self.tile_overlap_factor))
blend_extent = int(self.tile_sample_min_size * self.tile_overlap_factor)
row_limit = self.tile_sample_min_size - blend_extent
# Split z into overlapping 64x64 tiles and decode them separately.
# The tiles have an overlap to avoid seams between tiles.
rows = []
for i in range(0, z.shape[2], overlap_size):
row = []
for j in range(0, z.shape[3], overlap_size):
tile = z[:, :, i:i + self.tile_latent_min_size,
j:j + self.tile_latent_min_size]
if self.post_quant_conv:
tile = self.post_quant_conv(tile)
decoded = self.decoder(tile)
row.append(decoded)
rows.append(row)
result_rows = []
for i, row in enumerate(rows):
result_row = []
for j, tile in enumerate(row):
# blend the above tile and the left tile
# to the current tile and add the current tile to the result row
if i > 0:
tile = self.blend_v(rows[i - 1][j], tile, blend_extent)
if j > 0:
tile = self.blend_h(row[j - 1], tile, blend_extent)
result_row.append(tile[:, :, :row_limit, :row_limit])
result_rows.append(torch.cat(result_row, dim=3))
dec = torch.cat(result_rows, dim=2)
return dec
def forward(
self,
sample: torch.Tensor,
sample_posterior: bool = False,
generator: Optional[torch.Generator] = None,
) -> torch.Tensor:
r"""
Args:
sample (`torch.Tensor`): Input sample.
sample_posterior (`bool`, *optional*, defaults to `False`):
Whether to sample from the posterior.
"""
x = sample
posterior = self.encode(x).latent_dist
if sample_posterior:
z = posterior.sample(generator=generator)
else:
z = posterior.mode()
dec = self.decode(z).sample
return dec
# Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.fuse_qkv_projections
def fuse_qkv_projections(self):
"""
Enables fused QKV projections. For self-attention modules, all projection matrices (i.e., query, key, value)
are fused. For cross-attention modules, key and value projection matrices are fused.
<Tip warning={true}>
This API is 🧪 experimental.
</Tip>
"""
self.original_attn_processors = None
for _, attn_processor in self.attn_processors.items():
if "Added" in str(attn_processor.__class__.__name__):
raise ValueError(
"`fuse_qkv_projections()` is not supported for models having added KV projections."
)
self.original_attn_processors = self.attn_processors
for module in self.modules():
if isinstance(module, Attention):
module.fuse_projections(fuse=True)
self.set_attn_processor(FusedAttnProcessor2_0())
# Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.unfuse_qkv_projections
def unfuse_qkv_projections(self):
"""Disables the fused QKV projection if enabled.
<Tip warning={true}>
This API is 🧪 experimental.
</Tip>
"""
if self.original_attn_processors is not None:
self.set_attn_processor(self.original_attn_processors)
@@ -128,8 +128,12 @@ class ComposedPipelineBase(ABC):
modules_config = deepcopy(self.config)
# remove keys that are not pipeline modules
modules_config.pop("_class_name")
modules_config.pop("_diffusers_version")
keys_to_remove = [
key for key in list(modules_config.keys())
if str(key).startswith("_")
]
for key in keys_to_remove:
modules_config.pop(key)
# some sanity checks
assert len(
@@ -0,0 +1,359 @@
import importlib
import numpy as np
import torch
from einops import rearrange
from fastvideo.v1.attention import get_attn_backend
from fastvideo.v1.distributed import (get_sequence_model_parallel_rank,
get_sequence_model_parallel_world_size)
from fastvideo.v1.distributed.communication_op import (
sequence_model_parallel_all_gather)
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.forward_context import set_forward_context
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.pipelines.stages.denoising import DenoisingStage
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.utils import PRECISION_TO_TYPE
logger = init_logger(__name__)
class TimestepsPreparationPreStage(PipelineStage):
def __init__(self, scheduler) -> None:
super().__init__()
self.scheduler = scheduler
def calculate_shift(self,
image_seq_len,
base_seq_len: int = 256,
max_seq_len: int = 4096,
base_shift: float = 0.5,
max_shift: float = 1.15):
m = (max_shift - base_shift) / (max_seq_len - base_seq_len)
b = base_shift - m * base_seq_len
mu = image_seq_len * m + b
return mu
def forward(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
assert batch.height is not None
assert batch.width is not None
batch.sigmas = np.linspace(
1.0, 1 / batch.num_inference_steps,
batch.num_inference_steps) if batch.sigmas is None else batch.sigmas
spatial_compression_ratio = fastvideo_args.vae_config.arch_config.spatial_compression_ratio
batch.extra_set_timesteps_kwargs["mu"] = (
batch.extra_set_timesteps_kwargs.get("mu", None)
or self.calculate_shift(
(batch.height // spatial_compression_ratio // 2) *
(batch.width // spatial_compression_ratio // 2),
self.scheduler.config.get("base_image_seq_len", 256),
self.scheduler.config.get("max_image_seq_len", 4096),
self.scheduler.config.get("base_shift", 0.5),
self.scheduler.config.get("max_shift", 1.15),
))
return batch
class DenoisingPreprocessingStage(PipelineStage):
def __init__(self) -> None:
super().__init__()
def forward(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
# [B, in_channels // 4, 1, H, W] -> [B, H // 2, W // 2, in_channels]
assert batch.latents is not None
b, c, _, h, w = batch.latents.shape
latents = batch.latents.view(b, c, h // 2, 2, w // 2, 2)
latents = latents.permute(0, 2, 4, 1, 3, 5)
latents = latents.reshape(b, (h // 2) * (w // 2), c * 4)
batch.latents = latents
return batch
class DenoisingPostprocessingStage(PipelineStage):
def __init__(self) -> None:
super().__init__()
def forward(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
assert batch.latents is not None
assert batch.height_latents is not None
assert batch.width_latents is not None
latents = batch.latents
# Skip decoding if output type is latent
if fastvideo_args.output_type == "latent":
latents = latents
else:
# [B, (H // 2) * (W // 2), in_channels] -> [B, in_channels // 4, 1, H, W]
b, _, c = latents.shape
h, w = batch.height_latents, batch.width_latents
latents = latents.view(b, h // 2, w // 2, c // 4, 2, 2)
latents = latents.permute(0, 3, 1, 4, 2, 5)
latents = latents.reshape(b, c // 4, h, w)
# latents = latents.squeeze(2)
batch.latents = latents
return batch
st_attn_available = False
spec = importlib.util.find_spec("st_attn")
if spec is not None:
st_attn_available = True
from fastvideo.v1.attention.backends.sliding_tile_attn import (
SlidingTileAttentionBackend)
class FluxDenoisingStage(DenoisingStage):
"""
Stage for running the denoising loop in diffusion pipelines.
This stage handles the iterative denoising process that transforms
the initial noise into the final output.
"""
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Run the denoising loop.
Args:
batch: The current batch information.
fastvideo_args: The inference arguments.
Returns:
The batch with denoised latents.
"""
# Prepare extra step kwargs for scheduler
extra_step_kwargs = self.prepare_extra_func_kwargs(
self.scheduler.step,
{
"generator": batch.generator,
"eta": batch.eta
},
)
# Setup precision and autocast settings
target_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
# Handle sequence parallelism if enabled
world_size, rank = get_sequence_model_parallel_world_size(
), get_sequence_model_parallel_rank()
sp_group = world_size > 1
if sp_group:
latents = rearrange(batch.latents,
"b (n s) c -> b n s c",
n=world_size).contiguous()
latents = latents[:, rank, :, :]
batch.latents = latents
if batch.image_latent is not None:
image_latent = rearrange(batch.image_latent,
"b (n s) c -> b n s c",
n=world_size).contiguous()
image_latent = image_latent[:, rank, :, :]
batch.image_latent = image_latent
# Get timesteps and calculate warmup steps
timesteps = batch.timesteps
# TODO(will): remove this once we add input/output validation for stages
if timesteps is None:
raise ValueError("Timesteps must be provided")
num_inference_steps = batch.num_inference_steps
num_warmup_steps = len(
timesteps) - num_inference_steps * self.scheduler.order
# Create 3D list for mask strategy
def dict_to_3d_list(mask_strategy, t_max=50, l_max=60, h_max=24):
result = [[[None for _ in range(h_max)] for _ in range(l_max)]
for _ in range(t_max)]
if mask_strategy is None:
return result
for key, value in mask_strategy.items():
t, layer, h = map(int, key.split('_'))
result[t][layer][h] = value
return result
# Prepare image latents and embeddings for I2V generation
image_embeds = batch.image_embeds
if len(image_embeds) > 0:
assert torch.isnan(image_embeds[0]).sum() == 0
image_embeds = [
image_embed.to(target_dtype) for image_embed in image_embeds
]
image_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{
"encoder_hidden_states_image": image_embeds,
"height_latents": batch.height_latents,
"width_latents": batch.width_latents,
},
)
# Get latents and embeddings
latents = batch.latents
prompt_embeds = batch.prompt_embeds
assert torch.isnan(prompt_embeds[0]).sum() == 0
if batch.do_classifier_free_guidance:
neg_prompt_embeds = batch.negative_prompt_embeds
assert neg_prompt_embeds is not None
assert torch.isnan(neg_prompt_embeds[0]).sum() == 0
# Run denoising loop
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
# Skip if interrupted
if hasattr(self, 'interrupt') and self.interrupt:
continue
# Expand latents for I2V
latent_model_input = latents.to(target_dtype)
if batch.image_latent is not None:
latent_model_input = torch.cat(
[latent_model_input, batch.image_latent],
dim=1).to(target_dtype)
assert torch.isnan(latent_model_input).sum() == 0
latent_model_input = self.scheduler.scale_model_input(
latent_model_input, t)
# Prepare inputs for transformer
t_expand = t.repeat(latent_model_input.shape[0])
guidance_expand = (torch.tensor(
[fastvideo_args.embedded_cfg_scale] *
latent_model_input.shape[0],
dtype=torch.float32,
device=fastvideo_args.device,
).to(target_dtype) * 1000.0 if fastvideo_args.embedded_cfg_scale
is not None else None)
# Predict noise residual
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled):
# TODO(will-refactor): all of this should be in the stage's init
attn_head_size = self.transformer.hidden_size // self.transformer.num_attention_heads
self.attn_backend = get_attn_backend(
head_size=attn_head_size,
dtype=torch.float16, # TODO(will): hack
supported_attention_backends=(
_Backend.SLIDING_TILE_ATTN, _Backend.FLASH_ATTN,
_Backend.TORCH_SDPA) # hack
)
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
self.attn_metadata_builder_cls = self.attn_backend.get_builder_cls(
)
if self.attn_metadata_builder_cls is not None:
self.attn_metadata_builder = self.attn_metadata_builder_cls(
)
# TODO(will): clean this up
attn_metadata = self.attn_metadata_builder.build(
current_timestep=i,
forward_batch=batch,
fastvideo_args=fastvideo_args,
)
assert attn_metadata is not None, "attn_metadata cannot be None"
else:
attn_metadata = None
else:
attn_metadata = None
# TODO(will): finalize the interface. vLLM uses this to
# support torch dynamo compilation. They pass in
# attn_metadata, vllm_config, and num_tokens. We can pass in
# fastvideo_args or training_args, and attn_metadata.
with set_forward_context(
current_timestep=i,
attn_metadata=attn_metadata,
# fastvideo_args=fastvideo_args
):
# Run transformer
noise_pred = self.transformer(
latent_model_input,
prompt_embeds,
t_expand,
guidance=guidance_expand,
**image_kwargs,
)
# Apply guidance
if batch.do_classifier_free_guidance:
with set_forward_context(
current_timestep=i,
attn_metadata=attn_metadata,
# fastvideo_args=fastvideo_args
):
# Run transformer
noise_pred_uncond = self.transformer(
latent_model_input,
neg_prompt_embeds,
t_expand,
guidance=guidance_expand,
**image_kwargs,
)
noise_pred_text = noise_pred
noise_pred = noise_pred_uncond + batch.guidance_scale * (
noise_pred_text - noise_pred_uncond)
# Apply guidance rescale if needed
if batch.guidance_rescale > 0.0:
# Based on 3.4. in https://arxiv.org/pdf/2305.08891.pdf
noise_pred = self.rescale_noise_cfg(
noise_pred,
noise_pred_text,
guidance_rescale=batch.guidance_rescale,
)
# Compute the previous noisy sample
latents = self.scheduler.step(noise_pred,
t,
latents,
**extra_step_kwargs,
return_dict=False)[0]
# Update progress bar
if i == len(timesteps) - 1 or (
(i + 1) > num_warmup_steps and
(i + 1) % self.scheduler.order == 0
and progress_bar is not None):
progress_bar.update()
# Gather results if using sequence parallelism
if sp_group:
latents = sequence_model_parallel_all_gather(latents, dim=1)
# Update batch with final latents
batch.latents = latents
if fastvideo_args.use_cpu_offload:
self.transformer.to('cpu')
torch.cuda.empty_cache()
return batch
class ImageOutputStage(PipelineStage):
def __init__(self) -> None:
super().__init__()
def forward(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
output = batch.output
output = output.unsqueeze(2)
batch.output = output
return batch
@@ -0,0 +1,82 @@
# SPDX-License-Identifier: Apache-2.0
"""
Flux image diffusion pipeline implementation.
This module contains an implementation of the Flux image diffusion pipeline
using the modular pipeline architecture.
"""
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.flux.custom_stages import (
DenoisingPostprocessingStage, DenoisingPreprocessingStage,
FluxDenoisingStage, ImageOutputStage, TimestepsPreparationPreStage)
from fastvideo.v1.pipelines.stages import (ConditioningStage, DecodingStage,
InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
logger = init_logger(__name__)
class FluxPipeline(ComposedPipelineBase):
_required_config_modules = [
"text_encoder", "text_encoder_2", "tokenizer", "tokenizer_2", "vae",
"transformer", "scheduler"
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[
self.get_module("text_encoder"),
self.get_module("text_encoder_2"),
],
tokenizers=[
self.get_module("tokenizer"),
self.get_module("tokenizer_2"),
],
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="timesteps_preparation_pre_stage",
stage=TimestepsPreparationPreStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="modulation",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
self.add_stage(stage_name="denoising_preprocessing_stage",
stage=DenoisingPreprocessingStage())
self.add_stage(stage_name="denoising_stage",
stage=FluxDenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="denoising_postprocessing_stage",
stage=DenoisingPostprocessingStage())
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
self.add_stage(stage_name="output_stage", stage=ImageOutputStage())
EntryClass = FluxPipeline
@@ -79,6 +79,7 @@ class ForwardBatch:
# Timesteps
timesteps: Optional[torch.Tensor] = None
timestep: Optional[Union[torch.Tensor, float, int]] = None
extra_set_timesteps_kwargs: Dict[str, Any] = field(default_factory=dict)
step_index: Optional[int] = None
# Scheduler parameters
+29 -10
View File
@@ -64,6 +64,7 @@ class DenoisingStage(PipelineStage):
Returns:
The batch with denoised latents.
"""
assert batch.latents is not None
# Prepare extra step kwargs for scheduler
extra_step_kwargs = self.prepare_extra_func_kwargs(
self.scheduler.step,
@@ -83,16 +84,28 @@ class DenoisingStage(PipelineStage):
), get_sequence_model_parallel_rank()
sp_group = world_size > 1
if sp_group:
latents = rearrange(batch.latents,
"b t (n s) h w -> b t n s h w",
n=world_size).contiguous()
latents = latents[:, :, rank, :, :, :]
if batch.latents.shape[2] == 1:
latents = rearrange(batch.latents,
"b t f (n s) w -> b t f n s w",
n=world_size).contiguous()
latents = latents[:, :, :, rank, :, :]
else:
latents = rearrange(batch.latents,
"b t (n s) h w -> b t n s h w",
n=world_size).contiguous()
latents = latents[:, :, rank, :, :, :]
batch.latents = latents
if batch.image_latent is not None:
image_latent = rearrange(batch.image_latent,
"b t (n s) h w -> b t n s h w",
n=world_size).contiguous()
image_latent = image_latent[:, :, rank, :, :, :]
if batch.image_latent.shape[2] == 1:
image_latent = rearrange(batch.image_latent,
"b t f (n s) w -> b t f n s w",
n=world_size).contiguous()
image_latent = image_latent[:, :, :, rank, :, :]
else:
image_latent = rearrange(batch.image_latent,
"b t (n s) h w -> b t n s h w",
n=world_size).contiguous()
image_latent = image_latent[:, :, rank, :, :, :]
batch.image_latent = image_latent
# Get timesteps and calculate warmup steps
@@ -261,7 +274,11 @@ class DenoisingStage(PipelineStage):
# Gather results if using sequence parallelism
if sp_group:
latents = sequence_model_parallel_all_gather(latents, dim=2)
if latents.shape[2] == 1:
# image latents
latents = sequence_model_parallel_all_gather(latents, dim=3)
else:
latents = sequence_model_parallel_all_gather(latents, dim=2)
# Update batch with final latents
batch.latents = latents
@@ -285,7 +302,9 @@ class DenoisingStage(PipelineStage):
"""
extra_step_kwargs = {}
for k, v in kwargs.items():
accepts = k in set(inspect.signature(func).parameters.keys())
accepts = (k in set(inspect.signature(func).parameters.keys())
or "kwargs" in set(
inspect.signature(func).parameters.keys()))
if accepts:
extra_step_kwargs[k] = v
return extra_step_kwargs
@@ -102,7 +102,8 @@ class LatentPreparationStage(PipelineStage):
# Update batch with prepared latents
batch.latents = latents
batch.height_latents = latents.shape[-2]
batch.width_latents = latents.shape[-1]
return batch
def adjust_video_length(self, batch: ForwardBatch,
@@ -47,9 +47,9 @@ class TimestepPreparationStage(PipelineStage):
timesteps = batch.timesteps
sigmas = batch.sigmas
n_tokens = batch.n_tokens
extra_set_timesteps_kwargs = batch.extra_set_timesteps_kwargs
# Prepare extra kwargs for set_timesteps
extra_set_timesteps_kwargs = {}
if n_tokens is not None and "n_tokens" in inspect.signature(
scheduler.set_timesteps).parameters:
extra_set_timesteps_kwargs["n_tokens"] = n_tokens
@@ -0,0 +1,151 @@
# SPDX-License-Identifier: Apache-2.0
import os
import numpy as np
import pytest
import torch
from diffusers import FluxTransformer2DModel
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.models.loader.component_loader import TransformerLoader
from fastvideo.v1.utils import maybe_download_model
from fastvideo.v1.configs.models.dits import FluxImageConfig
logger = init_logger(__name__)
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29503"
BASE_MODEL_PATH = "black-forest-labs/FLUX.1-dev"
# BASE_MODEL_PATH = "/home/test/.cache/huggingface/hub/models--black-forest-labs--FLUX.1-dev/snapshots/0ef5fff789c832c5c7f4e127f94c8b54bbcced44/"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join(
'data', BASE_MODEL_PATH))
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
@pytest.mark.usefixtures("distributed_setup")
def test_flux_transformer():
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
precision = torch.bfloat16
precision_str = "bf16"
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
use_cpu_offload=False,
precision=precision_str)
args.device = device
args.dit_config = FluxImageConfig()
loader = TransformerLoader()
model2 = loader.load(TRANSFORMER_PATH, "", args).to(device, dtype=precision)
model1 = FluxTransformer2DModel.from_pretrained(
TRANSFORMER_PATH, device=device,
torch_dtype=precision).to(device, dtype=precision).requires_grad_(False)
total_params = sum(p.numel() for p in model1.parameters())
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
weight_sum_model1 = sum(
p.to(torch.float64).sum().item() for p in model1.parameters())
# Also calculate mean for more stable comparison
weight_mean_model1 = weight_sum_model1 / total_params
logger.info("Model 1 weight sum: %s", weight_sum_model1)
logger.info("Model 1 weight mean: %s", weight_mean_model1)
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
total_params_model2 = sum(p.numel() for p in model2.parameters())
weight_sum_model2 = sum(
p.to(torch.float64).sum().item() for p in model2.parameters())
# Also calculate mean for more stable comparison
weight_mean_model2 = weight_sum_model2 / total_params_model2
logger.info("Model 2 weight sum: %s", weight_sum_model2)
logger.info("Model 2 weight mean: %s", weight_mean_model2)
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
logger.info("Weight sum difference: %s", weight_sum_diff)
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
logger.info("Weight mean difference: %s", weight_mean_diff)
# Set both models to eval mode
model1 = model1.eval()
model2 = model2.eval()
# Create identical inputs for both models
batch_size = 1
seq_len = 30
# Video latents [B, N, C]
hidden_states = torch.randn(batch_size,
64 * 64,
64,
device=device, dtype=precision)
# Text embeddings [B, L, D] (including global token)
encoder_hidden_states = torch.randn(batch_size,
seq_len,
4096,
device=device,
dtype=precision)
pooled_text_embeds = torch.randn(batch_size,
768,
device=device,
dtype=precision)
# Timestep
timestep = torch.tensor([0.5], device=device, dtype=precision)
guidance_1 = torch.tensor([3.5], device=device, dtype=torch.float32).expand(batch_size)
guidance_2 = torch.tensor([3.5], device=device, dtype=torch.float32).expand(batch_size).to(precision) * 1000
def _prepare_latent_image_ids(batch_size, height, width, device, dtype):
latent_image_ids = torch.zeros(height, width, 3)
latent_image_ids[..., 1] = latent_image_ids[..., 1] + torch.arange(height)[:, None]
latent_image_ids[..., 2] = latent_image_ids[..., 2] + torch.arange(width)[None, :]
latent_image_id_height, latent_image_id_width, latent_image_id_channels = latent_image_ids.shape
latent_image_ids = latent_image_ids.reshape(
latent_image_id_height * latent_image_id_width, latent_image_id_channels
)
return latent_image_ids.to(device=device, dtype=dtype)
with torch.amp.autocast('cuda', dtype=precision):
print("Running model1 inference...")
output1 = model1(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
pooled_projections=pooled_text_embeds,
timestep=timestep,
txt_ids=torch.zeros(encoder_hidden_states.shape[1], 3).to(device=device, dtype=precision),
img_ids=_prepare_latent_image_ids(batch_size, 64, 64, device, precision),
guidance=guidance_1,
return_dict=False,
)[0]
print("Running model2 inference...")
with set_forward_context(
current_timestep=0,
attn_metadata=None,
):
output2 = model2(hidden_states=hidden_states,
encoder_hidden_states=[pooled_text_embeds,
encoder_hidden_states],
guidance=guidance_2,
timestep=timestep * 1000,
height_latents=128,
width_latents=128)
# Check if outputs have the same shape
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
# 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 < 4, f"Maximum difference between outputs: {max_diff.item()}"
assert mean_diff < 5e-1, f"Mean difference between outputs: {mean_diff.item()}"