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, w
(self, text_embeds: torch.Tensor, image_embeds: torch.Tensor)
| 716 | |
| 717 | |
| 718 | def forward(self, text_embeds: torch.Tensor, image_embeds: torch.Tensor): |
| 719 | r""" |
| 720 | Args: |
| 721 | text_embeds (`torch.Tensor`): |
| 722 | Input text embeddings. Expected shape: (batch_size, seq_length, embedding_dim). |
| 723 | image_embeds (`torch.Tensor`): |
| 724 | Input image embeddings. Expected shape: (batch_size, num_frames, channels, height, width). |
| 725 | """ |
| 726 | |
| 727 | text_embeds = self.text_proj(text_embeds) |
| 728 | |
| 729 | text_batch_size, text_seq_length, text_channels = text_embeds.shape |
| 730 | batch_size, num_frames, channels, height, width = image_embeds.shape |
| 731 | |
| 732 | |
| 733 | if self.patch_size_t is None: |
| 734 | image_embeds = image_embeds.reshape(-1, channels, height, width) |
| 735 | image_embeds = self.proj(image_embeds) # HACK: Project channels to 1920 dim, which is the same dim of Text |
| 736 | image_embeds = image_embeds.view(batch_size, num_frames, *image_embeds.shape[1:]) |
| 737 | image_embeds = image_embeds.flatten(3).transpose(2, 3) # [batch, num_frames, height x width, channels] |
| 738 | image_embeds = image_embeds.flatten(1, 2) # [batch, num_frames x height x width, channels] |
| 739 | else: |
| 740 | p = self.patch_size |
| 741 | p_t = self.patch_size_t |
| 742 | |
| 743 | image_embeds = image_embeds.permute(0, 1, 3, 4, 2) |
| 744 | image_embeds = image_embeds.reshape( |
| 745 | batch_size, num_frames // p_t, p_t, height // p, p, width // p, p, channels |
| 746 | ) |
| 747 | image_embeds = image_embeds.permute(0, 1, 3, 5, 7, 2, 4, 6).flatten(4, 7).flatten(1, 3) |
| 748 | image_embeds = self.proj(image_embeds) |
| 749 | |
| 750 | embeds = torch.cat( |
| 751 | [text_embeds, image_embeds], dim=1 |
| 752 | ).contiguous() # [batch, seq_length + num_frames x height x width, channels] |
| 753 | |
| 754 | |
| 755 | |
| 756 | if self.use_positional_embeddings or self.use_learned_positional_embeddings: |
| 757 | # if self.use_learned_positional_embeddings and (self.sample_width != width or self.sample_height != height): |
| 758 | # raise ValueError( |
| 759 | # "It is currently not possible to generate videos at a different resolution that the defaults. This should only be the case with 'THUDM/CogVideoX-5b-I2V'." |
| 760 | # "If you think this is incorrect, please open an issue at https://github.com/huggingface/diffusers/issues." |
| 761 | # ) |
| 762 | |
| 763 | pre_time_compression_frames = (num_frames - 1) * self.temporal_compression_ratio + 1 |
| 764 | post_time_compression_frames = (self.sample_frames - 1) // self.temporal_compression_ratio + 1 |
| 765 | post_patch_height = self.sample_height // self.patch_size |
| 766 | post_patch_width = self.sample_width // self.patch_size |
| 767 | seq_length = height * width * num_frames // (self.patch_size**2) |
| 768 | |
| 769 | |
| 770 | |
| 771 | # Directly reuse the pos_embedding available |
| 772 | if self.use_FrameIn: |
| 773 | first_frame_token_num = (self.pos_embedding.shape[1] - self.max_text_seq_length) // (num_frames - 1) # Minus the number of frames |
| 774 | # NOTE: the following place has bug where it overlap with the text token |
| 775 | pos_embeds = torch.cat([self.pos_embedding, self.pos_embedding[:, text_seq_length:text_seq_length+first_frame_token_num].clone()], dim=1) # Append the pos_embeds in the token-wise dimension |