MCPcopy Create free account
hub / github.com/UVA-Computer-Vision-Lab/FrameINO / forward

Method forward

architecture/embeddings.py:718–805  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls 1

toMethod · 0.45

Tested by

no test coverage detected