'update'
This commit is contained in:
@@ -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()
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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()
|
||||
|
||||
|
||||
@@ -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"))
|
||||
|
||||
Reference in New Issue
Block a user