From 320d69a7c5a4e6dbc33a4a86cc0918b1912d6747 Mon Sep 17 00:00:00 2001 From: bubbliiiing <3323290568@qq.com> Date: Wed, 18 Sep 2024 17:19:44 +0800 Subject: [PATCH] fix some bug in import --- cogvideox/models/transformer3d.py | 40 ++++++++++++++++++++++++++++++- cogvideox/ui/ui.py | 2 +- requirements.txt | 2 +- 3 files changed, 41 insertions(+), 3 deletions(-) diff --git a/cogvideox/models/transformer3d.py b/cogvideox/models/transformer3d.py index a0e7982..b80af91 100644 --- a/cogvideox/models/transformer3d.py +++ b/cogvideox/models/transformer3d.py @@ -27,7 +27,7 @@ from diffusers.utils import is_torch_version, 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, CogVideoXAttnProcessor2_0, FusedCogVideoXAttnProcessor2_0 -from diffusers.models.embeddings import CogVideoXPatchEmbed, TimestepEmbedding, Timesteps, get_3d_sincos_pos_embed +from diffusers.models.embeddings import TimestepEmbedding, Timesteps, get_3d_sincos_pos_embed from diffusers.models.modeling_outputs import Transformer2DModelOutput from diffusers.models.modeling_utils import ModelMixin from diffusers.models.normalization import AdaLayerNorm, CogVideoXLayerNormZero @@ -35,6 +35,44 @@ from diffusers.models.normalization import AdaLayerNorm, CogVideoXLayerNormZero logger = logging.get_logger(__name__) # pylint: disable=invalid-name +class CogVideoXPatchEmbed(nn.Module): + def __init__( + self, + patch_size: int = 2, + in_channels: int = 16, + embed_dim: int = 1920, + text_embed_dim: int = 4096, + bias: bool = True, + ) -> None: + super().__init__() + self.patch_size = patch_size + + self.proj = nn.Conv2d( + in_channels, embed_dim, kernel_size=(patch_size, patch_size), stride=patch_size, bias=bias + ) + self.text_proj = nn.Linear(text_embed_dim, embed_dim) + + def forward(self, text_embeds: torch.Tensor, image_embeds: torch.Tensor): + r""" + Args: + text_embeds (`torch.Tensor`): + Input text embeddings. Expected shape: (batch_size, seq_length, embedding_dim). + image_embeds (`torch.Tensor`): + Input image embeddings. Expected shape: (batch_size, num_frames, channels, height, width). + """ + text_embeds = self.text_proj(text_embeds) + + batch, num_frames, channels, height, width = image_embeds.shape + image_embeds = image_embeds.reshape(-1, channels, height, width) + image_embeds = self.proj(image_embeds) + image_embeds = image_embeds.view(batch, num_frames, *image_embeds.shape[1:]) + image_embeds = image_embeds.flatten(3).transpose(2, 3) # [batch, num_frames, height x width, channels] + image_embeds = image_embeds.flatten(1, 2) # [batch, num_frames x height x width, channels] + + embeds = torch.cat( + [text_embeds, image_embeds], dim=1 + ).contiguous() # [batch, seq_length + num_frames x height x width, channels] + return embeds @maybe_allow_in_graph class CogVideoXBlock(nn.Module): diff --git a/cogvideox/ui/ui.py b/cogvideox/ui/ui.py index 55f7f00..d5b9bc4 100644 --- a/cogvideox/ui/ui.py +++ b/cogvideox/ui/ui.py @@ -1132,7 +1132,7 @@ def post_eas( class CogVideoX_I2VController_EAS: - def __init__(self, edition, config_path, model_name, savedir_sample): + def __init__(self, model_name, savedir_sample): self.savedir_sample = savedir_sample os.makedirs(self.savedir_sample, exist_ok=True) diff --git a/requirements.txt b/requirements.txt index a81f6bc..f5df893 100644 --- a/requirements.txt +++ b/requirements.txt @@ -24,5 +24,5 @@ func_timeout deepspeed accelerate>=0.25.0 gradio>=3.41.2 -diffusers>=0.28.2 +diffusers>=0.30.1 transformers>=4.37.2