maybe
This commit is contained in:
+186
-183
@@ -21,20 +21,18 @@ import torch.nn.functional as F
|
||||
|
||||
import numpy as np
|
||||
from einops import rearrange
|
||||
from functools import reduce
|
||||
from operator import mul
|
||||
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.utils import logging
|
||||
from diffusers.utils.torch_utils import maybe_allow_in_graph
|
||||
from diffusers.models.attention import Attention, FeedForward
|
||||
from diffusers.models.attention_processor import AttentionProcessor
|
||||
from diffusers.models.embeddings import CogVideoXPatchEmbed, TimestepEmbedding, Timesteps
|
||||
from diffusers.models.embeddings import TimestepEmbedding, Timesteps
|
||||
from diffusers.models.modeling_outputs import Transformer2DModelOutput
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from diffusers.models.normalization import AdaLayerNorm, CogVideoXLayerNormZero
|
||||
from diffusers.loaders import PeftAdapterMixin
|
||||
from .embeddings import CogVideoX1_1PatchEmbed
|
||||
from .embeddings import CogVideoXPatchEmbed
|
||||
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
@@ -64,14 +62,6 @@ def fft(tensor):
|
||||
|
||||
return low_freq_fft, high_freq_fft
|
||||
|
||||
def rotate_half(x):
|
||||
x = rearrange(x, "... (d r) -> ... d r", r=2)
|
||||
x1, x2 = x.unbind(dim=-1)
|
||||
x = torch.stack((-x2, x1), dim=-1)
|
||||
return rearrange(x, "... d r -> ... (d r)")
|
||||
|
||||
|
||||
|
||||
class CogVideoXAttnProcessor2_0:
|
||||
r"""
|
||||
Processor for implementing scaled dot-product attention for the CogVideoX model. It applies a rotary embedding on
|
||||
@@ -81,16 +71,7 @@ class CogVideoXAttnProcessor2_0:
|
||||
def __init__(self):
|
||||
if not hasattr(F, "scaled_dot_product_attention"):
|
||||
raise ImportError("CogVideoXAttnProcessor requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
|
||||
def rotary(self, t, rope_args):
|
||||
def reshape_freq(freqs):
|
||||
freqs = freqs[: rope_args["T"], : rope_args["H"], : rope_args["W"]].contiguous()
|
||||
freqs = rearrange(freqs, "t h w d -> (t h w) d")
|
||||
freqs = freqs.unsqueeze(0).unsqueeze(0)
|
||||
return freqs
|
||||
freqs_cos = reshape_freq(self.freqs_cos).to(t.dtype)
|
||||
freqs_sin = reshape_freq(self.freqs_sin).to(t.dtype)
|
||||
|
||||
return t * freqs_cos + rotate_half(t) * freqs_sin
|
||||
|
||||
@torch.compiler.disable()
|
||||
def __call__(
|
||||
self,
|
||||
@@ -99,7 +80,6 @@ class CogVideoXAttnProcessor2_0:
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
image_rotary_emb: Optional[torch.Tensor] = None,
|
||||
rope_args: Optional[dict] = None
|
||||
) -> torch.Tensor:
|
||||
text_seq_length = encoder_hidden_states.size(1)
|
||||
|
||||
@@ -129,127 +109,118 @@ class CogVideoXAttnProcessor2_0:
|
||||
if attn.norm_k is not None:
|
||||
key = attn.norm_k(key)
|
||||
|
||||
|
||||
|
||||
# Apply RoPE if needed
|
||||
if image_rotary_emb is not None:
|
||||
self.freqs_cos = image_rotary_emb[0]
|
||||
self.freqs_sin = image_rotary_emb[1]
|
||||
print("rope args", rope_args) #{'T': 6, 'H': 30, 'W': 45, 'seq_length': 8775}
|
||||
print("freqs_cos", self.freqs_cos.shape) #torch.Size([13, 30, 45, 64])
|
||||
print("freqs_sin", self.freqs_sin.shape)
|
||||
|
||||
|
||||
from diffusers.models.embeddings import apply_rotary_emb
|
||||
|
||||
#query[:, :, text_seq_length:] = apply_rotary_emb(query[:, :, text_seq_length:], image_rotary_emb)
|
||||
query = torch.cat(
|
||||
(query[:, :, : text_seq_length],
|
||||
self.rotary(query[:, :, text_seq_length:],
|
||||
rope_args)),
|
||||
dim=2)
|
||||
|
||||
if not attn.is_cross_attention:
|
||||
#key[:, :, text_seq_length:] = apply_rotary_emb(key[:, :, text_seq_length:], image_rotary_emb)
|
||||
key = torch.cat(
|
||||
(key[ :, :, : text_seq_length],
|
||||
self.rotary(key[:, :, text_seq_length:],
|
||||
rope_args)),
|
||||
dim=2)
|
||||
|
||||
if SAGEATTN_IS_AVAILABLE:
|
||||
hidden_states = sageattn(query, key, value, is_causal=False)
|
||||
else:
|
||||
hidden_states = F.scaled_dot_product_attention(
|
||||
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
|
||||
)
|
||||
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
|
||||
encoder_hidden_states, hidden_states = hidden_states.split(
|
||||
[text_seq_length, hidden_states.size(1) - text_seq_length], dim=1
|
||||
)
|
||||
return hidden_states, encoder_hidden_states
|
||||
|
||||
|
||||
class FusedCogVideoXAttnProcessor2_0:
|
||||
r"""
|
||||
Processor for implementing scaled dot-product attention for the CogVideoX model. It applies a rotary embedding on
|
||||
query and key vectors, but does not include spatial normalization.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
if not hasattr(F, "scaled_dot_product_attention"):
|
||||
raise ImportError("CogVideoXAttnProcessor requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
|
||||
@torch.compiler.disable()
|
||||
def __call__(
|
||||
self,
|
||||
attn: Attention,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
image_rotary_emb: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
text_seq_length = encoder_hidden_states.size(1)
|
||||
|
||||
hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
|
||||
|
||||
batch_size, sequence_length, _ = (
|
||||
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
|
||||
)
|
||||
|
||||
if attention_mask is not None:
|
||||
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
|
||||
attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1])
|
||||
|
||||
qkv = attn.to_qkv(hidden_states)
|
||||
split_size = qkv.shape[-1] // 3
|
||||
query, key, value = torch.split(qkv, split_size, dim=-1)
|
||||
|
||||
inner_dim = key.shape[-1]
|
||||
head_dim = inner_dim // attn.heads
|
||||
|
||||
query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
|
||||
if attn.norm_q is not None:
|
||||
query = attn.norm_q(query)
|
||||
if attn.norm_k is not None:
|
||||
key = attn.norm_k(key)
|
||||
|
||||
# Apply RoPE if needed
|
||||
if image_rotary_emb is not None:
|
||||
from diffusers.models.embeddings import apply_rotary_emb
|
||||
has_nan = torch.isnan(query).any()
|
||||
if has_nan:
|
||||
raise ValueError(f"query before rope has nan: {has_nan}")
|
||||
|
||||
query[:, :, text_seq_length:] = apply_rotary_emb(query[:, :, text_seq_length:], image_rotary_emb)
|
||||
query[:, :, text_seq_length:] = apply_rotary_emb(query[:, :, text_seq_length:], image_rotary_emb)
|
||||
if not attn.is_cross_attention:
|
||||
key[:, :, text_seq_length:] = apply_rotary_emb(key[:, :, text_seq_length:], image_rotary_emb)
|
||||
|
||||
if SAGEATTN_IS_AVAILABLE:
|
||||
hidden_states = sageattn(query, key, value, is_causal=False)
|
||||
else:
|
||||
hidden_states = F.scaled_dot_product_attention(
|
||||
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
|
||||
)
|
||||
#if SAGEATTN_IS_AVAILABLE:
|
||||
# hidden_states = sageattn(query, key, value, is_causal=False)
|
||||
#else:
|
||||
hidden_states = F.scaled_dot_product_attention(
|
||||
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
|
||||
)
|
||||
has_nan = torch.isnan(hidden_states).any()
|
||||
if has_nan:
|
||||
raise ValueError(f"hs after scaled_dot_product_attention has nan: {has_nan}")
|
||||
has_inf = torch.isinf(hidden_states).any()
|
||||
if has_inf:
|
||||
raise ValueError(f"hs after scaled_dot_product_attention has inf: {has_inf}")
|
||||
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
has_nan = torch.isnan(hidden_states).any()
|
||||
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
|
||||
encoder_hidden_states, hidden_states = hidden_states.split(
|
||||
[text_seq_length, hidden_states.size(1) - text_seq_length], dim=1
|
||||
)
|
||||
return hidden_states, encoder_hidden_states
|
||||
|
||||
|
||||
|
||||
# class FusedCogVideoXAttnProcessor2_0:
|
||||
# r"""
|
||||
# Processor for implementing scaled dot-product attention for the CogVideoX model. It applies a rotary embedding on
|
||||
# query and key vectors, but does not include spatial normalization.
|
||||
# """
|
||||
|
||||
# def __init__(self):
|
||||
# if not hasattr(F, "scaled_dot_product_attention"):
|
||||
# raise ImportError("CogVideoXAttnProcessor requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
|
||||
# @torch.compiler.disable()
|
||||
# def __call__(
|
||||
# self,
|
||||
# attn: Attention,
|
||||
# hidden_states: torch.Tensor,
|
||||
# encoder_hidden_states: torch.Tensor,
|
||||
# attention_mask: Optional[torch.Tensor] = None,
|
||||
# image_rotary_emb: Optional[torch.Tensor] = None,
|
||||
# ) -> torch.Tensor:
|
||||
# print("FusedCogVideoXAttnProcessor2_0")
|
||||
# text_seq_length = encoder_hidden_states.size(1)
|
||||
|
||||
# hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
|
||||
|
||||
# batch_size, sequence_length, _ = (
|
||||
# hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
|
||||
# )
|
||||
|
||||
# if attention_mask is not None:
|
||||
# attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
|
||||
# attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1])
|
||||
|
||||
# qkv = attn.to_qkv(hidden_states)
|
||||
# split_size = qkv.shape[-1] // 3
|
||||
# query, key, value = torch.split(qkv, split_size, dim=-1)
|
||||
|
||||
# inner_dim = key.shape[-1]
|
||||
# head_dim = inner_dim // attn.heads
|
||||
|
||||
# query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
# key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
# value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
|
||||
# if attn.norm_q is not None:
|
||||
# query = attn.norm_q(query)
|
||||
# if attn.norm_k is not None:
|
||||
# key = attn.norm_k(key)
|
||||
|
||||
# # Apply RoPE if needed
|
||||
# if image_rotary_emb is not None:
|
||||
# from diffusers.models.embeddings import apply_rotary_emb
|
||||
|
||||
# query[:, :, text_seq_length:] = apply_rotary_emb(query[:, :, text_seq_length:], image_rotary_emb)
|
||||
# if not attn.is_cross_attention:
|
||||
# key[:, :, text_seq_length:] = apply_rotary_emb(key[:, :, text_seq_length:], image_rotary_emb)
|
||||
|
||||
# if SAGEATTN_IS_AVAILABLE:
|
||||
# hidden_states = sageattn(query, key, value, is_causal=False)
|
||||
# else:
|
||||
# hidden_states = F.scaled_dot_product_attention(
|
||||
# query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
|
||||
# )
|
||||
|
||||
# hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
|
||||
|
||||
# # linear proj
|
||||
# hidden_states = attn.to_out[0](hidden_states)
|
||||
# # dropout
|
||||
# hidden_states = attn.to_out[1](hidden_states)
|
||||
|
||||
# encoder_hidden_states, hidden_states = hidden_states.split(
|
||||
# [text_seq_length, hidden_states.size(1) - text_seq_length], dim=1
|
||||
# )
|
||||
# return hidden_states, encoder_hidden_states
|
||||
|
||||
#region Blocks
|
||||
@maybe_allow_in_graph
|
||||
class CogVideoXBlock(nn.Module):
|
||||
|
||||
@@ -344,14 +315,14 @@ class CogVideoXBlock(nn.Module):
|
||||
fuser=None,
|
||||
fastercache_counter=0,
|
||||
fastercache_start_step=15,
|
||||
fastercache_device="cuda:0",
|
||||
rope_args=None
|
||||
fastercache_device="cuda:0"
|
||||
) -> torch.Tensor:
|
||||
text_seq_length = encoder_hidden_states.size(1)
|
||||
# norm & modulate
|
||||
norm_hidden_states, norm_encoder_hidden_states, gate_msa, enc_gate_msa = self.norm1(
|
||||
hidden_states, encoder_hidden_states, temb
|
||||
)
|
||||
|
||||
# Tora Motion-guidance Fuser
|
||||
if video_flow_feature is not None:
|
||||
H, W = video_flow_feature.shape[-2:]
|
||||
@@ -378,7 +349,7 @@ class CogVideoXBlock(nn.Module):
|
||||
attn_hidden_states, attn_encoder_hidden_states = self.attn1(
|
||||
hidden_states=norm_hidden_states,
|
||||
encoder_hidden_states=norm_encoder_hidden_states,
|
||||
image_rotary_emb=image_rotary_emb, rope_args=rope_args
|
||||
image_rotary_emb=image_rotary_emb
|
||||
)
|
||||
if fastercache_counter == fastercache_start_step:
|
||||
self.cached_hidden_states = [attn_hidden_states.to(fastercache_device), attn_hidden_states.to(fastercache_device)]
|
||||
@@ -386,10 +357,18 @@ class CogVideoXBlock(nn.Module):
|
||||
elif fastercache_counter > fastercache_start_step:
|
||||
self.cached_hidden_states[-1].copy_(attn_hidden_states.to(fastercache_device))
|
||||
self.cached_encoder_hidden_states[-1].copy_(attn_encoder_hidden_states.to(fastercache_device))
|
||||
|
||||
|
||||
|
||||
hidden_states = hidden_states + gate_msa * attn_hidden_states
|
||||
encoder_hidden_states = encoder_hidden_states + enc_gate_msa * attn_encoder_hidden_states
|
||||
|
||||
# has_nan = torch.isnan(hidden_states).any()
|
||||
# if has_nan:
|
||||
# raise ValueError(f"hs before norm2 has nan: {has_nan}")
|
||||
# has_inf = torch.isinf(hidden_states).any()
|
||||
# if has_inf:
|
||||
# raise ValueError(f"hs before norm2 has inf: {has_inf}")
|
||||
|
||||
# norm & modulate
|
||||
norm_hidden_states, norm_encoder_hidden_states, gate_ff, enc_gate_ff = self.norm2(
|
||||
hidden_states, encoder_hidden_states, temb
|
||||
@@ -404,7 +383,7 @@ class CogVideoXBlock(nn.Module):
|
||||
|
||||
return hidden_states, encoder_hidden_states
|
||||
|
||||
|
||||
#region Transformer
|
||||
class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
|
||||
"""
|
||||
A Transformer model for video-like data in [CogVideoX](https://github.com/THUDM/CogVideo).
|
||||
@@ -479,6 +458,7 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
|
||||
sample_height: int = 60,
|
||||
sample_frames: int = 49,
|
||||
patch_size: int = 2,
|
||||
patch_size_t: int = 2,
|
||||
temporal_compression_ratio: int = 4,
|
||||
max_text_seq_length: int = 226,
|
||||
activation_fn: str = "gelu-approximate",
|
||||
@@ -489,6 +469,7 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
|
||||
temporal_interpolation_scale: float = 1.0,
|
||||
use_rotary_positional_embeddings: bool = False,
|
||||
use_learned_positional_embeddings: bool = False,
|
||||
patch_bias: bool = True,
|
||||
):
|
||||
super().__init__()
|
||||
inner_dim = num_attention_heads * attention_head_dim
|
||||
@@ -501,12 +482,13 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
|
||||
)
|
||||
|
||||
# 1. Patch embedding
|
||||
self.patch_embed = CogVideoX1_1PatchEmbed(
|
||||
self.patch_embed = CogVideoXPatchEmbed(
|
||||
patch_size=patch_size,
|
||||
patch_size_t=patch_size_t,
|
||||
in_channels=in_channels,
|
||||
embed_dim=inner_dim,
|
||||
text_embed_dim=text_embed_dim,
|
||||
#bias=True,
|
||||
bias=patch_bias,
|
||||
sample_width=sample_width,
|
||||
sample_height=sample_height,
|
||||
sample_frames=sample_frames,
|
||||
@@ -550,7 +532,14 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
|
||||
norm_eps=norm_eps,
|
||||
chunk_dim=1,
|
||||
)
|
||||
self.proj_out = nn.Linear(inner_dim, patch_size * patch_size * patch_size * out_channels)
|
||||
if patch_size_t is None:
|
||||
# For CogVideox 1.0
|
||||
output_dim = patch_size * patch_size * out_channels
|
||||
else:
|
||||
# For CogVideoX 1.5
|
||||
output_dim = patch_size * patch_size * patch_size_t * out_channels
|
||||
|
||||
self.proj_out = nn.Linear(inner_dim, output_dim)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
@@ -626,44 +615,44 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
|
||||
fn_recursive_attn_processor(name, module, processor)
|
||||
|
||||
# Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.fuse_qkv_projections with FusedAttnProcessor2_0->FusedCogVideoXAttnProcessor2_0
|
||||
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.
|
||||
# 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}>
|
||||
# <Tip warning={true}>
|
||||
|
||||
This API is 🧪 experimental.
|
||||
# This API is 🧪 experimental.
|
||||
|
||||
</Tip>
|
||||
"""
|
||||
self.original_attn_processors = None
|
||||
# </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.")
|
||||
# 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
|
||||
# self.original_attn_processors = self.attn_processors
|
||||
|
||||
for module in self.modules():
|
||||
if isinstance(module, Attention):
|
||||
module.fuse_projections(fuse=True)
|
||||
# for module in self.modules():
|
||||
# if isinstance(module, Attention):
|
||||
# module.fuse_projections(fuse=True)
|
||||
|
||||
self.set_attn_processor(FusedCogVideoXAttnProcessor2_0())
|
||||
# self.set_attn_processor(FusedCogVideoXAttnProcessor2_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.
|
||||
# def unfuse_qkv_projections(self):
|
||||
# """Disables the fused QKV projection if enabled.
|
||||
|
||||
<Tip warning={true}>
|
||||
# <Tip warning={true}>
|
||||
|
||||
This API is 🧪 experimental.
|
||||
# This API is 🧪 experimental.
|
||||
|
||||
</Tip>
|
||||
# </Tip>
|
||||
|
||||
"""
|
||||
if self.original_attn_processors is not None:
|
||||
self.set_attn_processor(self.original_attn_processors)
|
||||
# """
|
||||
# if self.original_attn_processors is not None:
|
||||
# self.set_attn_processor(self.original_attn_processors)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -678,9 +667,7 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
|
||||
return_dict: bool = True,
|
||||
):
|
||||
batch_size, num_frames, channels, height, width = hidden_states.shape
|
||||
p = self.config.patch_size
|
||||
print("p", p)
|
||||
|
||||
|
||||
# 1. Time embedding
|
||||
timesteps = timestep
|
||||
t_emb = self.time_proj(timesteps)
|
||||
@@ -691,25 +678,24 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
|
||||
t_emb = t_emb.to(dtype=hidden_states.dtype)
|
||||
emb = self.time_embedding(t_emb, timestep_cond)
|
||||
|
||||
# RoPE
|
||||
seq_length = num_frames * height * width // reduce(mul, [p, p, p])
|
||||
rope_T = num_frames // p
|
||||
rope_H = height // p
|
||||
rope_W = width // p
|
||||
rope_args = {
|
||||
"T": rope_T,
|
||||
"H": rope_H,
|
||||
"W": rope_W,
|
||||
"seq_length": seq_length,
|
||||
}
|
||||
|
||||
# 2. Patch embedding
|
||||
p = self.config.patch_size
|
||||
p_t = self.config.patch_size_t
|
||||
|
||||
# We know that the hidden states height and width will always be divisible by patch_size.
|
||||
# But, the number of frames may not be divisible by patch_size_t. So, we pad with the beginning frames.
|
||||
if p_t is not None:
|
||||
remaining_frames = p_t - num_frames % p_t
|
||||
first_frame = hidden_states[:, :1].repeat(1, 1 + remaining_frames, 1, 1, 1)
|
||||
hidden_states = torch.cat([first_frame, hidden_states[:, 1:]], dim=1)
|
||||
|
||||
hidden_states = self.patch_embed(encoder_hidden_states, hidden_states)
|
||||
hidden_states = self.embedding_dropout(hidden_states)
|
||||
|
||||
|
||||
text_seq_length = encoder_hidden_states.shape[1]
|
||||
encoder_hidden_states = hidden_states[:, :text_seq_length]
|
||||
hidden_states = hidden_states[:, text_seq_length:]
|
||||
|
||||
if self.use_fastercache:
|
||||
self.fastercache_counter+=1
|
||||
if self.fastercache_counter >= self.fastercache_start_step + 3 and self.fastercache_counter % 5 !=0:
|
||||
@@ -754,8 +740,15 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
|
||||
# - It is okay to `channels` use for CogVideoX-2b and CogVideoX-5b (number of input channels is equal to output channels)
|
||||
# - However, for CogVideoX-5b-I2V also takes concatenated input image latents (number of input channels is twice the output channels)
|
||||
|
||||
output = hidden_states.reshape(1, num_frames, height // p, width // p, -1, p, p)
|
||||
output = output.permute(0, 1, 4, 2, 5, 3, 6).flatten(5, 6).flatten(3, 4)
|
||||
if p_t is None:
|
||||
output = hidden_states.reshape(batch_size, num_frames, height // p, width // p, -1, p, p)
|
||||
output = output.permute(0, 1, 4, 2, 5, 3, 6).flatten(5, 6).flatten(3, 4)
|
||||
else:
|
||||
output = hidden_states.reshape(
|
||||
batch_size, (num_frames + p_t - 1) // p_t, height // p, width // p, -1, p_t, p, p
|
||||
)
|
||||
output = output.permute(0, 1, 5, 4, 2, 6, 3, 7).flatten(6, 7).flatten(4, 5).flatten(1, 2)
|
||||
output = output[:, remaining_frames:]
|
||||
|
||||
(bb, tt, cc, hh, ww) = output.shape
|
||||
cond = rearrange(output, "B T C H W -> (B T) C H W", B=bb, C=cc, T=tt, H=hh, W=ww)
|
||||
@@ -777,6 +770,7 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
|
||||
output = torch.cat([output, recovered_uncond])
|
||||
else:
|
||||
for i, block in enumerate(self.transformer_blocks):
|
||||
print("block", i)
|
||||
hidden_states, encoder_hidden_states = block(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
@@ -785,9 +779,11 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
|
||||
video_flow_feature=video_flow_features[i] if video_flow_features is not None else None,
|
||||
fuser = self.fuser_list[i] if self.fuser_list is not None else None,
|
||||
fastercache_counter = self.fastercache_counter,
|
||||
fastercache_device = self.fastercache_device,
|
||||
rope_args=rope_args
|
||||
fastercache_device = self.fastercache_device
|
||||
)
|
||||
has_nan = torch.isnan(hidden_states).any()
|
||||
if has_nan:
|
||||
raise ValueError(f"block output hidden_states has nan: {has_nan}")
|
||||
|
||||
if (controlnet_states is not None) and (i < len(controlnet_states)):
|
||||
controlnet_states_block = controlnet_states[i]
|
||||
@@ -816,9 +812,16 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
|
||||
# Note: we use `-1` instead of `channels`:
|
||||
# - It is okay to `channels` use for CogVideoX-2b and CogVideoX-5b (number of input channels is equal to output channels)
|
||||
# - However, for CogVideoX-5b-I2V also takes concatenated input image latents (number of input channels is twice the output channels)
|
||||
p = self.config.patch_size
|
||||
output = hidden_states.reshape(batch_size, num_frames, height // p, width // p, -1, p, p)
|
||||
output = output.permute(0, 1, 4, 2, 5, 3, 6).flatten(5, 6).flatten(3, 4)
|
||||
|
||||
if p_t is None:
|
||||
output = hidden_states.reshape(batch_size, num_frames, height // p, width // p, -1, p, p)
|
||||
output = output.permute(0, 1, 4, 2, 5, 3, 6).flatten(5, 6).flatten(3, 4)
|
||||
else:
|
||||
output = hidden_states.reshape(
|
||||
batch_size, (num_frames + p_t - 1) // p_t, height // p, width // p, -1, p_t, p, p
|
||||
)
|
||||
output = output.permute(0, 1, 5, 4, 2, 6, 3, 7).flatten(6, 7).flatten(4, 5).flatten(1, 2)
|
||||
output = output[:, remaining_frames:]
|
||||
|
||||
if self.fastercache_counter >= self.fastercache_start_step + 1:
|
||||
(bb, tt, cc, hh, ww) = output.shape
|
||||
|
||||
Reference in New Issue
Block a user