This commit is contained in:
宣源
2024-09-12 15:40:44 +08:00
parent c28d4b379a
commit a49ebb962c
+19 -2
View File
@@ -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)