This commit is contained in:
wailovet
2024-10-10 00:17:14 +08:00
parent 5111a465ab
commit e93b92b47f
7 changed files with 358 additions and 55 deletions
+204
View File
@@ -0,0 +1,204 @@
# https://github.com/TheDenk/cogvideox-controlnet/blob/main/cogvideo_controlnet.py
from typing import Any, Dict, Optional, Tuple, Union
import torch
from torch import nn
from einops import rearrange
import torch.nn.functional as F
from .custom_cogvideox_transformer_3d import Transformer2DModelOutput, CogVideoXBlock
from diffusers.utils import is_torch_version
from diffusers.loaders import PeftAdapterMixin
from diffusers.models.embeddings import CogVideoXPatchEmbed, TimestepEmbedding, Timesteps
from diffusers.models.modeling_utils import ModelMixin
from diffusers.configuration_utils import ConfigMixin, register_to_config
class CogVideoXControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin):
@register_to_config
def __init__(
self,
num_attention_heads: int = 30,
attention_head_dim: int = 64,
vae_channels: int = 16,
in_channels: int = 3,
downscale_coef: int = 8,
flip_sin_to_cos: bool = True,
freq_shift: int = 0,
time_embed_dim: int = 512,
num_layers: int = 8,
dropout: float = 0.0,
attention_bias: bool = True,
sample_width: int = 90,
sample_height: int = 60,
sample_frames: int = 49,
patch_size: int = 2,
temporal_compression_ratio: int = 4,
max_text_seq_length: int = 226,
activation_fn: str = "gelu-approximate",
timestep_activation_fn: str = "silu",
norm_elementwise_affine: bool = True,
norm_eps: float = 1e-5,
spatial_interpolation_scale: float = 1.875,
temporal_interpolation_scale: float = 1.0,
use_rotary_positional_embeddings: bool = False,
use_learned_positional_embeddings: bool = False,
):
super().__init__()
inner_dim = num_attention_heads * attention_head_dim
if not use_rotary_positional_embeddings and use_learned_positional_embeddings:
raise ValueError(
"There are no CogVideoX checkpoints available with disable rotary embeddings and learned positional "
"embeddings. If you're using a custom model and/or believe this should be supported, please open an "
"issue at https://github.com/huggingface/diffusers/issues."
)
start_channels = in_channels * (downscale_coef ** 2)
input_channels = [start_channels, start_channels // 2, start_channels // 4]
self.unshuffle = nn.PixelUnshuffle(downscale_coef)
self.controlnet_encode_first = nn.Sequential(
nn.Conv2d(input_channels[0], input_channels[1], kernel_size=1, stride=1, padding=0),
nn.GroupNorm(2, input_channels[1]),
nn.ReLU(),
)
self.controlnet_encode_second = nn.Sequential(
nn.Conv2d(input_channels[1], input_channels[2], kernel_size=1, stride=1, padding=0),
nn.GroupNorm(2, input_channels[2]),
nn.ReLU(),
)
# 1. Patch embedding
self.patch_embed = CogVideoXPatchEmbed(
patch_size=patch_size,
in_channels=vae_channels + input_channels[2],
embed_dim=inner_dim,
bias=True,
sample_width=sample_width,
sample_height=sample_height,
sample_frames=sample_frames,
temporal_compression_ratio=temporal_compression_ratio,
spatial_interpolation_scale=spatial_interpolation_scale,
temporal_interpolation_scale=temporal_interpolation_scale,
use_positional_embeddings=not use_rotary_positional_embeddings,
use_learned_positional_embeddings=use_learned_positional_embeddings,
)
self.embedding_dropout = nn.Dropout(dropout)
# 2. Time embeddings
self.time_proj = Timesteps(inner_dim, flip_sin_to_cos, freq_shift)
self.time_embedding = TimestepEmbedding(inner_dim, time_embed_dim, timestep_activation_fn)
# 3. Define spatio-temporal transformers blocks
self.transformer_blocks = nn.ModuleList(
[
CogVideoXBlock(
dim=inner_dim,
num_attention_heads=num_attention_heads,
attention_head_dim=attention_head_dim,
time_embed_dim=time_embed_dim,
dropout=dropout,
activation_fn=activation_fn,
attention_bias=attention_bias,
norm_elementwise_affine=norm_elementwise_affine,
norm_eps=norm_eps,
)
for _ in range(num_layers)
]
)
self.gradient_checkpointing = False
def compress_time(self, x, num_frames):
x = rearrange(x, '(b f) c h w -> b f c h w', f=num_frames)
batch_size, frames, channels, height, width = x.shape
x = rearrange(x, 'b f c h w -> (b h w) c f')
if x.shape[-1] % 2 == 1:
x_first, x_rest = x[..., 0], x[..., 1:]
if x_rest.shape[-1] > 0:
x_rest = F.avg_pool1d(x_rest, kernel_size=2, stride=2)
x = torch.cat([x_first[..., None], x_rest], dim=-1)
else:
x = F.avg_pool1d(x, kernel_size=2, stride=2)
x = rearrange(x, '(b h w) c f -> (b f) c h w', b=batch_size, h=height, w=width)
return x
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
controlnet_states: torch.Tensor,
timestep: Union[int, float, torch.LongTensor],
image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
timestep_cond: Optional[torch.Tensor] = None,
return_dict: bool = True,
):
batch_size, num_frames, channels, height, width = controlnet_states.shape
# 0. Controlnet encoder
controlnet_states = rearrange(controlnet_states, 'b f c h w -> (b f) c h w')
controlnet_states = self.unshuffle(controlnet_states)
controlnet_states = self.controlnet_encode_first(controlnet_states)
controlnet_states = self.compress_time(controlnet_states, num_frames=num_frames)
num_frames = controlnet_states.shape[0] // batch_size
controlnet_states = self.controlnet_encode_second(controlnet_states)
controlnet_states = self.compress_time(controlnet_states, num_frames=num_frames)
controlnet_states = rearrange(controlnet_states, '(b f) c h w -> b f c h w', b=batch_size)
hidden_states = torch.cat([hidden_states, controlnet_states], dim=2)
# controlnet_states = self.controlnext_encoder(controlnet_states, timestep=timestep)
# 1. Time embedding
timesteps = timestep
t_emb = self.time_proj(timesteps)
# timesteps does not contain any weights and will always return f32 tensors
# but time_embedding might actually be running in fp16. so we need to cast here.
# there might be better ways to encapsulate this.
t_emb = t_emb.to(dtype=hidden_states.dtype)
emb = self.time_embedding(t_emb, timestep_cond)
hidden_states = self.patch_embed(encoder_hidden_states, hidden_states)
hidden_states = self.embedding_dropout(hidden_states)
text_seq_length = encoder_hidden_states.shape[1]
encoder_hidden_states = hidden_states[:, :text_seq_length]
hidden_states = hidden_states[:, text_seq_length:]
controlnet_hidden_states = ()
# 3. Transformer blocks
for i, block in enumerate(self.transformer_blocks):
if self.training and self.gradient_checkpointing:
def create_custom_forward(module):
def custom_forward(*inputs):
return module(*inputs)
return custom_forward
ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {}
hidden_states, encoder_hidden_states = torch.utils.checkpoint.checkpoint(
create_custom_forward(block),
hidden_states,
encoder_hidden_states,
emb,
image_rotary_emb,
**ckpt_kwargs,
)
else:
hidden_states, encoder_hidden_states = block(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
temb=emb,
image_rotary_emb=image_rotary_emb,
)
controlnet_hidden_states += (hidden_states,)
if not return_dict:
return (controlnet_hidden_states,)
return Transformer2DModelOutput(sample=controlnet_hidden_states)
@@ -829,8 +829,6 @@ class CogVideoX_Fun_Pipeline_Control(VideoSysPipeline):
))
counter = torch.zeros_like(latent_model_input)
noise_pred = torch.zeros_like(latent_model_input)
if do_classifier_free_guidance:
noise_uncond = torch.zeros_like(latent_model_input)
image_rotary_emb = (
self._prepare_rotary_positional_embeddings(height, width, context_frames, device)
@@ -851,17 +849,6 @@ class CogVideoX_Fun_Pipeline_Control(VideoSysPipeline):
return_dict=False,
control_latents=partial_control_latents,
)[0]
# uncond
if do_classifier_free_guidance:
noise_uncond[:, c, :, :, :] += self.transformer(
hidden_states=partial_latent_model_input,
encoder_hidden_states=prompt_embeds,
timestep=timestep,
image_rotary_emb=image_rotary_emb,
return_dict=False,
control_latents=partial_control_latents,
)[0]
counter[:, c, :, :, :] += 1
noise_pred = noise_pred.float()
@@ -984,9 +984,6 @@ class CogVideoX_Fun_Pipeline_Inpaint(VideoSysPipeline):
latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
# Calculate the current step percentage
current_step_percentage = i / num_inference_steps
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
timestep = t.expand(latent_model_input.shape[0])
@@ -995,8 +992,6 @@ class CogVideoX_Fun_Pipeline_Inpaint(VideoSysPipeline):
))
counter = torch.zeros_like(latent_model_input)
noise_pred = torch.zeros_like(latent_model_input)
if do_classifier_free_guidance:
noise_uncond = torch.zeros_like(latent_model_input)
image_rotary_emb = (
self._prepare_rotary_positional_embeddings(height, width, context_frames, device)
@@ -1020,15 +1015,6 @@ class CogVideoX_Fun_Pipeline_Inpaint(VideoSysPipeline):
)[0]
counter[:, c, :, :, :] += 1
if do_classifier_free_guidance:
noise_uncond[:, c, :, :, :] += self.transformer(
hidden_states=partial_latent_model_input,
encoder_hidden_states=prompt_embeds,
timestep=timestep,
image_rotary_emb=image_rotary_emb,
return_dict=False,
inpaint_latents=partial_inpaint_latents,
)[0]
noise_pred = noise_pred.float()
+14
View File
@@ -19,6 +19,8 @@ import torch
from torch import nn
import torch.nn.functional as F
import numpy as np
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.utils import is_torch_version, logging
from diffusers.utils.torch_utils import maybe_allow_in_graph
@@ -566,6 +568,8 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin):
timestep: Union[int, float, torch.LongTensor],
timestep_cond: Optional[torch.Tensor] = None,
image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
controlnet_states: torch.Tensor = None,
controlnet_weights: Optional[Union[float, int, list, np.ndarray, torch.FloatTensor]] = 1.0,
return_dict: bool = True,
):
batch_size, num_frames, channels, height, width = hidden_states.shape
@@ -615,6 +619,16 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin):
image_rotary_emb=image_rotary_emb,
)
if (controlnet_states is not None) and (i < len(controlnet_states)):
controlnet_states_block = controlnet_states[i]
controlnet_block_weight = 1.0
if isinstance(controlnet_weights, (list, np.ndarray)) or torch.is_tensor(controlnet_weights):
controlnet_block_weight = controlnet_weights[i]
elif isinstance(controlnet_weights, (float, int)):
controlnet_block_weight = controlnet_weights
hidden_states = hidden_states + controlnet_states_block * controlnet_block_weight
if not self.config.use_rotary_positional_embeddings:
# CogVideoX-2B
hidden_states = self.norm_final(hidden_states)
+3 -1
View File
@@ -18,7 +18,9 @@ from .cogvideox_fun.fun_pab_transformer_3d import CogVideoXTransformer3DModel as
from .cogvideox_fun.autoencoder_magvit import AutoencoderKLCogVideoX as AutoencoderKLCogVideoXFun
from .cogvideox_fun.utils import get_image_to_video_latent, ASPECT_RATIO_512, get_closest_ratio, to_pil
from .cogvideox_fun.pipeline_cogvideox_inpaint import CogVideoX_Fun_Pipeline_Inpaint
from diffusers.models import AutoencoderKLCogVideoX, CogVideoXTransformer3DModel
# from diffusers.models import AutoencoderKLCogVideoX, CogVideoXTransformer3DModel
from diffusers.models import AutoencoderKLCogVideoX
from .custom_cogvideox_transformer_3d import CogVideoXTransformer3DModel
from diffusers.schedulers import CogVideoXDDIMScheduler
+111 -22
View File
@@ -387,6 +387,8 @@ class CogVideoXPipeline(VideoSysPipeline):
context_stride: Optional[int] = None,
context_overlap: Optional[int] = None,
freenoise: Optional[bool] = True,
controlnet: Optional[dict] = None,
):
"""
Function invoked when calling the pipeline for generation.
@@ -536,7 +538,7 @@ class CogVideoXPipeline(VideoSysPipeline):
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
comfy_pbar = ProgressBar(num_inference_steps)
# 8.5. Temporal tiling prep
# 8. context schedule and temporal tiling
if context_schedule is not None and context_schedule == "temporal_tiling":
t_tile_length = context_frames
t_tile_overlap = context_overlap
@@ -562,7 +564,34 @@ class CogVideoXPipeline(VideoSysPipeline):
if self.transformer.config.use_rotary_positional_embeddings
else None
)
# 9. Controlnet
if controlnet is not None:
self.controlnet = controlnet["control_model"].to(device)
if self.transformer.dtype == torch.float8_e4m3fn:
for name, param in self.controlnet.named_parameters():
if "patch_embed" not in name and param.data.dtype != torch.float8_e4m3fn:
param.data = param.data.to(torch.float8_e4m3fn)
else:
self.controlnet.to(self.transformer.dtype)
if getattr(self.transformer, 'fp8_matmul_enabled', False):
from .fp8_optimization import convert_fp8_linear
if not hasattr(self.controlnet, 'fp8_matmul_enabled') or not self.controlnet.fp8_matmul_enabled:
convert_fp8_linear(self.controlnet, torch.float16)
setattr(self.controlnet, "fp8_matmul_enabled", True)
control_frames = controlnet["control_frames"].to(device).to(self.controlnet.dtype).contiguous()
control_frames = torch.cat([control_frames] * 2) if do_classifier_free_guidance else control_frames
control_weights = controlnet["control_weights"]
print("Controlnet enabled with weights: ", control_weights)
control_start = controlnet["control_start"]
control_end = controlnet["control_end"]
else:
controlnet_states = None
control_weights= None
# 10. Denoising loop
with self.progress_bar(total=num_inference_steps) as progress_bar:
old_pred_original_sample = None # for DPM-solver++
for i, t in enumerate(timesteps):
@@ -666,8 +695,6 @@ class CogVideoXPipeline(VideoSysPipeline):
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
counter = torch.zeros_like(latent_model_input)
noise_pred = torch.zeros_like(latent_model_input)
if do_classifier_free_guidance:
noise_uncond = torch.zeros_like(latent_model_input)
if image_cond_latents is not None:
latent_image_input = torch.cat([image_cond_latents] * 2) if do_classifier_free_guidance else image_cond_latents
@@ -676,39 +703,79 @@ class CogVideoXPipeline(VideoSysPipeline):
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
timestep = t.expand(latent_model_input.shape[0])
context_queue = list(context(
i, num_inference_steps, latents.shape[1], context_frames, context_stride, context_overlap,
))
current_step_percentage = i / num_inference_steps
# use same rotary embeddings for all context windows
image_rotary_emb = (
self._prepare_rotary_positional_embeddings(height, width, context_frames, device)
if self.transformer.config.use_rotary_positional_embeddings
else None
)
for c in context_queue:
partial_latent_model_input = latent_model_input[:, c, :, :, :]
# predict noise model_output
noise_pred[:, c, :, :, :] += self.transformer(
hidden_states=partial_latent_model_input,
encoder_hidden_states=prompt_embeds,
timestep=timestep,
image_rotary_emb=image_rotary_emb,
return_dict=False,
)[0]
context_queue = list(context(
i, num_inference_steps, latents.shape[1], context_frames, context_stride, context_overlap,
))
# uncond
if do_classifier_free_guidance:
noise_uncond[:, c, :, :, :] += self.transformer(
if controlnet is not None:
# controlnet frames are not temporally compressed, so try to match the context frames that are
control_context_queue = list(context(
i,
num_inference_steps,
control_frames.shape[1],
context_frames * self.vae_scale_factor_temporal,
context_stride * self.vae_scale_factor_temporal,
context_overlap * self.vae_scale_factor_temporal,
))
for c, control_c in zip(context_queue, control_context_queue):
partial_latent_model_input = latent_model_input[:, c, :, :, :]
partial_control_frames = control_frames[:, control_c, :, :, :]
controlnet_states = None
if (control_start <= current_step_percentage <= control_end):
# extract controlnet hidden state
controlnet_states = self.controlnet(
hidden_states=partial_latent_model_input,
encoder_hidden_states=prompt_embeds,
image_rotary_emb=image_rotary_emb,
controlnet_states=partial_control_frames,
timestep=timestep,
return_dict=False,
)[0]
if isinstance(controlnet_states, (tuple, list)):
controlnet_states = [x.to(dtype=self.controlnet.dtype) for x in controlnet_states]
else:
controlnet_states = controlnet_states.to(dtype=self.controlnet.dtype)
# predict noise model_output
noise_pred[:, c, :, :, :] += self.transformer(
hidden_states=partial_latent_model_input,
encoder_hidden_states=prompt_embeds,
timestep=timestep,
image_rotary_emb=image_rotary_emb,
return_dict=False,
controlnet_states=controlnet_states,
controlnet_weights=control_weights,
)[0]
counter[:, c, :, :, :] += 1
noise_pred = noise_pred.float()
counter[:, c, :, :, :] += 1
noise_pred = noise_pred.float()
else:
for c in context_queue:
partial_latent_model_input = latent_model_input[:, c, :, :, :]
# predict noise model_output
noise_pred[:, c, :, :, :] += self.transformer(
hidden_states=partial_latent_model_input,
encoder_hidden_states=prompt_embeds,
timestep=timestep,
image_rotary_emb=image_rotary_emb,
return_dict=False
)[0]
counter[:, c, :, :, :] += 1
noise_pred = noise_pred.float()
noise_pred /= counter
if do_classifier_free_guidance:
@@ -744,6 +811,26 @@ class CogVideoXPipeline(VideoSysPipeline):
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
timestep = t.expand(latent_model_input.shape[0])
current_step_percentage = i / num_inference_steps
if controlnet is not None:
controlnet_states = None
if (control_start <= current_step_percentage <= control_end):
# extract controlnet hidden state
controlnet_states = self.controlnet(
hidden_states=latent_model_input,
encoder_hidden_states=prompt_embeds,
image_rotary_emb=image_rotary_emb,
controlnet_states=control_frames,
timestep=timestep,
return_dict=False,
)[0]
if isinstance(controlnet_states, (tuple, list)):
controlnet_states = [x.to(dtype=self.vae.dtype) for x in controlnet_states]
else:
controlnet_states = controlnet_states.to(dtype=self.vae.dtype)
# predict noise model_output
noise_pred = self.transformer(
hidden_states=latent_model_input,
@@ -751,6 +838,8 @@ class CogVideoXPipeline(VideoSysPipeline):
timestep=timestep,
image_rotary_emb=image_rotary_emb,
return_dict=False,
controlnet_states=controlnet_states,
controlnet_weights=control_weights,
)[0]
noise_pred = noise_pred.float()
+26 -5
View File
@@ -27,7 +27,11 @@ from .modules.embeddings import apply_rotary_emb
#from .modules.embeddings import CogVideoXPatchEmbed
from .modules.normalization import AdaLayerNorm, CogVideoXLayerNormZero
try:
from sageattention import sageattn
SAGEATTN_IS_AVAVILABLE = True
except:
SAGEATTN_IS_AVAVILABLE = False
class CogVideoXAttnProcessor2_0:
r"""
@@ -98,9 +102,12 @@ class CogVideoXAttnProcessor2_0:
key[:, :, text_seq_length : emb_len + text_seq_length], image_rotary_emb
)
hidden_states = F.scaled_dot_product_attention(
if SAGEATTN_IS_AVAVILABLE:
hidden_states = sageattn(query, key, value, is_causal=False)
else:
hidden_states = F.scaled_dot_product_attention(
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
)
)
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn_heads * head_dim)
@@ -170,9 +177,12 @@ class FusedCogVideoXAttnProcessor2_0:
if not attn.is_cross_attention:
key[:, :, text_seq_length:] = apply_rotary_emb(key[:, :, text_seq_length:], image_rotary_emb)
hidden_states = F.scaled_dot_product_attention(
if SAGEATTN_IS_AVAVILABLE:
hidden_states = sageattn(query, key, value, is_causal=False)
else:
hidden_states = F.scaled_dot_product_attention(
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
)
)
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
@@ -512,6 +522,8 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin):
timestep_cond: Optional[torch.Tensor] = None,
image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
return_dict: bool = True,
controlnet_states: torch.Tensor = None,
controlnet_weights: Optional[Union[float, int, list, torch.FloatTensor]] = 1.0,
):
# if self.parallel_manager.cp_size > 1:
# (
@@ -587,6 +599,15 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin):
image_rotary_emb=image_rotary_emb,
timestep=timesteps if enable_pab() else None,
)
if (controlnet_states is not None) and (i < len(controlnet_states)):
controlnet_states_block = controlnet_states[i]
controlnet_block_weight = 1.0
if isinstance(controlnet_weights, (list)) or torch.is_tensor(controlnet_weights):
controlnet_block_weight = controlnet_weights[i]
elif isinstance(controlnet_weights, (float, int)):
controlnet_block_weight = controlnet_weights
hidden_states = hidden_states + controlnet_states_block * controlnet_block_weight
#if self.parallel_manager.sp_size > 1:
# hidden_states = gather_sequence(hidden_states, self.parallel_manager.sp_group, dim=1, pad=get_pad("pad"))