From a49ebb962c115498e26f9eb9a1ec98270db42147 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=AE=A3=E6=BA=90?= Date: Thu, 12 Sep 2024 15:40:44 +0800 Subject: [PATCH] add swa --- cogvideox/models/transformer3d.py | 21 +++++++++++++++++++-- 1 file changed, 19 insertions(+), 2 deletions(-) diff --git a/cogvideox/models/transformer3d.py b/cogvideox/models/transformer3d.py index 5aa7323..ef59345 100644 --- a/cogvideox/models/transformer3d.py +++ b/cogvideox/models/transformer3d.py @@ -31,6 +31,8 @@ from diffusers.models.modeling_outputs import Transformer2DModelOutput from diffusers.models.modeling_utils import ModelMixin from diffusers.models.normalization import AdaLayerNorm, CogVideoXLayerNormZero +from cogvideox.models.attention_processor import CogVideoXSWAAttnProcessor2_0 + logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -87,6 +89,7 @@ class CogVideoXBlock(nn.Module): ff_inner_dim: Optional[int] = None, ff_bias: bool = True, attention_out_bias: bool = True, + swa: bool = False, ): super().__init__() @@ -101,7 +104,7 @@ class CogVideoXBlock(nn.Module): eps=1e-6, bias=attention_bias, out_bias=attention_out_bias, - processor=CogVideoXAttnProcessor2_0(), + processor=CogVideoXSWAAttnProcessor2_0() if swa else CogVideoXAttnProcessor2_0(), ) # 2. Feed Forward @@ -122,6 +125,9 @@ class CogVideoXBlock(nn.Module): encoder_hidden_states: torch.Tensor, temb: torch.Tensor, image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, + num_frames: int = None, + height: int = None, + width: int = None, ) -> torch.Tensor: text_seq_length = encoder_hidden_states.size(1) @@ -135,6 +141,9 @@ class CogVideoXBlock(nn.Module): hidden_states=norm_hidden_states, encoder_hidden_states=norm_encoder_hidden_states, image_rotary_emb=image_rotary_emb, + num_frames = num_frames, + height = height, + width = width, ) hidden_states = hidden_states + gate_msa * attn_hidden_states @@ -238,6 +247,7 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin): spatial_interpolation_scale: float = 1.875, temporal_interpolation_scale: float = 1.0, use_rotary_positional_embeddings: bool = False, + swa: bool = False, ): super().__init__() inner_dim = num_attention_heads * attention_head_dim @@ -285,6 +295,7 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin): attention_bias=attention_bias, norm_elementwise_affine=norm_elementwise_affine, norm_eps=norm_eps, + swa=swa, ) for _ in range(num_layers) ] @@ -417,6 +428,7 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin): return_dict: bool = True, ): batch_size, num_frames, channels, height, width = hidden_states.shape + p = self.config.patch_size # 1. Time embedding timesteps = timestep @@ -469,7 +481,13 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin): encoder_hidden_states, emb, image_rotary_emb, + num_frames, + height // p, + width // p, **ckpt_kwargs, + num_frames=num_frames, + height=height // p, + width=width // p, ) else: hidden_states, encoder_hidden_states = block( @@ -493,7 +511,6 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin): hidden_states = self.proj_out(hidden_states) # 6. Unpatchify - p = self.config.patch_size output = hidden_states.reshape(batch_size, num_frames, height // p, width // p, channels, p, p) output = output.permute(0, 1, 4, 2, 5, 3, 6).flatten(5, 6).flatten(3, 4)