diff --git a/cogvideox_fun/context.py b/cogvideox_fun/context.py
new file mode 100644
index 0000000..6a30fed
--- /dev/null
+++ b/cogvideox_fun/context.py
@@ -0,0 +1,184 @@
+import numpy as np
+from typing import Callable, Optional, List
+
+
+def ordered_halving(val):
+ bin_str = f"{val:064b}"
+ bin_flip = bin_str[::-1]
+ as_int = int(bin_flip, 2)
+
+ return as_int / (1 << 64)
+
+def does_window_roll_over(window: list[int], num_frames: int) -> tuple[bool, int]:
+ prev_val = -1
+ for i, val in enumerate(window):
+ val = val % num_frames
+ if val < prev_val:
+ return True, i
+ prev_val = val
+ return False, -1
+
+def shift_window_to_start(window: list[int], num_frames: int):
+ start_val = window[0]
+ for i in range(len(window)):
+ # 1) subtract each element by start_val to move vals relative to the start of all frames
+ # 2) add num_frames and take modulus to get adjusted vals
+ window[i] = ((window[i] - start_val) + num_frames) % num_frames
+
+def shift_window_to_end(window: list[int], num_frames: int):
+ # 1) shift window to start
+ shift_window_to_start(window, num_frames)
+ end_val = window[-1]
+ end_delta = num_frames - end_val - 1
+ for i in range(len(window)):
+ # 2) add end_delta to each val to slide windows to end
+ window[i] = window[i] + end_delta
+
+def get_missing_indexes(windows: list[list[int]], num_frames: int) -> list[int]:
+ all_indexes = list(range(num_frames))
+ for w in windows:
+ for val in w:
+ try:
+ all_indexes.remove(val)
+ except ValueError:
+ pass
+ return all_indexes
+
+def uniform_looped(
+ step: int = ...,
+ num_steps: Optional[int] = None,
+ num_frames: int = ...,
+ context_size: Optional[int] = None,
+ context_stride: int = 3,
+ context_overlap: int = 4,
+ closed_loop: bool = True,
+):
+ if num_frames <= context_size:
+ yield list(range(num_frames))
+ return
+
+ context_stride = min(context_stride, int(np.ceil(np.log2(num_frames / context_size))) + 1)
+
+ for context_step in 1 << np.arange(context_stride):
+ pad = int(round(num_frames * ordered_halving(step)))
+ for j in range(
+ int(ordered_halving(step) * context_step) + pad,
+ num_frames + pad + (0 if closed_loop else -context_overlap),
+ (context_size * context_step - context_overlap),
+ ):
+ yield [e % num_frames for e in range(j, j + context_size * context_step, context_step)]
+
+#from AnimateDiff-Evolved by Kosinkadink (https://github.com/Kosinkadink/ComfyUI-AnimateDiff-Evolved)
+def uniform_standard(
+ step: int = ...,
+ num_steps: Optional[int] = None,
+ num_frames: int = ...,
+ context_size: Optional[int] = None,
+ context_stride: int = 3,
+ context_overlap: int = 4,
+ closed_loop: bool = True,
+):
+ windows = []
+ if num_frames <= context_size:
+ windows.append(list(range(num_frames)))
+ return windows
+
+ context_stride = min(context_stride, int(np.ceil(np.log2(num_frames / context_size))) + 1)
+
+ for context_step in 1 << np.arange(context_stride):
+ pad = int(round(num_frames * ordered_halving(step)))
+ for j in range(
+ int(ordered_halving(step) * context_step) + pad,
+ num_frames + pad + (0 if closed_loop else -context_overlap),
+ (context_size * context_step - context_overlap),
+ ):
+ windows.append([e % num_frames for e in range(j, j + context_size * context_step, context_step)])
+
+ # now that windows are created, shift any windows that loop, and delete duplicate windows
+ delete_idxs = []
+ win_i = 0
+ while win_i < len(windows):
+ # if window is rolls over itself, need to shift it
+ is_roll, roll_idx = does_window_roll_over(windows[win_i], num_frames)
+ if is_roll:
+ roll_val = windows[win_i][roll_idx] # roll_val might not be 0 for windows of higher strides
+ shift_window_to_end(windows[win_i], num_frames=num_frames)
+ # check if next window (cyclical) is missing roll_val
+ if roll_val not in windows[(win_i+1) % len(windows)]:
+ # need to insert new window here - just insert window starting at roll_val
+ windows.insert(win_i+1, list(range(roll_val, roll_val + context_size)))
+ # delete window if it's not unique
+ for pre_i in range(0, win_i):
+ if windows[win_i] == windows[pre_i]:
+ delete_idxs.append(win_i)
+ break
+ win_i += 1
+
+ # reverse delete_idxs so that they will be deleted in an order that doesn't break idx correlation
+ delete_idxs.reverse()
+ for i in delete_idxs:
+ windows.pop(i)
+ return windows
+
+def static_standard(
+ step: int = ...,
+ num_steps: Optional[int] = None,
+ num_frames: int = ...,
+ context_size: Optional[int] = None,
+ context_stride: int = 3,
+ context_overlap: int = 4,
+ closed_loop: bool = True,
+):
+ windows = []
+ if num_frames <= context_size:
+ windows.append(list(range(num_frames)))
+ return windows
+ # always return the same set of windows
+ delta = context_size - context_overlap
+ for start_idx in range(0, num_frames, delta):
+ # if past the end of frames, move start_idx back to allow same context_length
+ ending = start_idx + context_size
+ if ending >= num_frames:
+ final_delta = ending - num_frames
+ final_start_idx = start_idx - final_delta
+ windows.append(list(range(final_start_idx, final_start_idx + context_size)))
+ break
+ windows.append(list(range(start_idx, start_idx + context_size)))
+ return windows
+
+def get_context_scheduler(name: str) -> Callable:
+ if name == "uniform_looped":
+ return uniform_looped
+ elif name == "uniform_standard":
+ return uniform_standard
+ elif name == "static_standard":
+ return static_standard
+ else:
+ raise ValueError(f"Unknown context_overlap policy {name}")
+
+
+def get_total_steps(
+ scheduler,
+ timesteps: List[int],
+ num_steps: Optional[int] = None,
+ num_frames: int = ...,
+ context_size: Optional[int] = None,
+ context_stride: int = 3,
+ context_overlap: int = 4,
+ closed_loop: bool = True,
+):
+ return sum(
+ len(
+ list(
+ scheduler(
+ i,
+ num_steps,
+ num_frames,
+ context_size,
+ context_stride,
+ context_overlap,
+ )
+ )
+ )
+ for i in range(len(timesteps))
+ )
diff --git a/cogvideox_fun/fun_pab_transformer_3d.py b/cogvideox_fun/fun_pab_transformer_3d.py
index cb524e8..25a3934 100644
--- a/cogvideox_fun/fun_pab_transformer_3d.py
+++ b/cogvideox_fun/fun_pab_transformer_3d.py
@@ -37,6 +37,14 @@ from ..videosys.modules.embeddings import apply_rotary_emb
from ..videosys.core.pab_mgr import enable_pab, if_broadcast_spatial
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
+try:
+ from sageattention import sageattn
+ SAGEATTN_IS_AVAVILABLE = True
+ logger.info("Using sageattn")
+except:
+ logger.info("sageattn not found, using sdpa")
+ SAGEATTN_IS_AVAVILABLE = False
+
class CogVideoXAttnProcessor2_0:
r"""
Processor for implementing scaled dot-product attention for the CogVideoX model. It applies a rotary embedding on
@@ -106,9 +114,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)
@@ -178,9 +189,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)
diff --git a/cogvideox_fun/lora_utils.py b/cogvideox_fun/lora_utils.py
new file mode 100644
index 0000000..ccb3f65
--- /dev/null
+++ b/cogvideox_fun/lora_utils.py
@@ -0,0 +1,525 @@
+# LoRA network module
+# reference:
+# https://github.com/microsoft/LoRA/blob/main/loralib/layers.py
+# https://github.com/cloneofsimo/lora/blob/master/lora_diffusion/lora.py
+# https://github.com/bmaltais/kohya_ss
+
+import hashlib
+import math
+import os
+from collections import defaultdict
+from io import BytesIO
+from typing import List, Optional, Type, Union
+
+import safetensors.torch
+import torch
+import torch.utils.checkpoint
+from diffusers.models.lora import LoRACompatibleConv, LoRACompatibleLinear
+from safetensors.torch import load_file
+from transformers import T5EncoderModel
+
+
+class LoRAModule(torch.nn.Module):
+ """
+ replaces forward method of the original Linear, instead of replacing the original Linear module.
+ """
+
+ def __init__(
+ self,
+ lora_name,
+ org_module: torch.nn.Module,
+ multiplier=1.0,
+ lora_dim=4,
+ alpha=1,
+ dropout=None,
+ rank_dropout=None,
+ module_dropout=None,
+ ):
+ """if alpha == 0 or None, alpha is rank (no scaling)."""
+ super().__init__()
+ self.lora_name = lora_name
+
+ if org_module.__class__.__name__ == "Conv2d":
+ in_dim = org_module.in_channels
+ out_dim = org_module.out_channels
+ else:
+ in_dim = org_module.in_features
+ out_dim = org_module.out_features
+
+ self.lora_dim = lora_dim
+ if org_module.__class__.__name__ == "Conv2d":
+ kernel_size = org_module.kernel_size
+ stride = org_module.stride
+ padding = org_module.padding
+ self.lora_down = torch.nn.Conv2d(in_dim, self.lora_dim, kernel_size, stride, padding, bias=False)
+ self.lora_up = torch.nn.Conv2d(self.lora_dim, out_dim, (1, 1), (1, 1), bias=False)
+ else:
+ self.lora_down = torch.nn.Linear(in_dim, self.lora_dim, bias=False)
+ self.lora_up = torch.nn.Linear(self.lora_dim, out_dim, bias=False)
+
+ if type(alpha) == torch.Tensor:
+ alpha = alpha.detach().float().numpy() # without casting, bf16 causes error
+ alpha = self.lora_dim if alpha is None or alpha == 0 else alpha
+ self.scale = alpha / self.lora_dim
+ self.register_buffer("alpha", torch.tensor(alpha))
+
+ # same as microsoft's
+ torch.nn.init.kaiming_uniform_(self.lora_down.weight, a=math.sqrt(5))
+ torch.nn.init.zeros_(self.lora_up.weight)
+
+ self.multiplier = multiplier
+ self.org_module = org_module # remove in applying
+ self.dropout = dropout
+ self.rank_dropout = rank_dropout
+ self.module_dropout = module_dropout
+
+ def apply_to(self):
+ self.org_forward = self.org_module.forward
+ self.org_module.forward = self.forward
+ del self.org_module
+
+ def forward(self, x, *args, **kwargs):
+ weight_dtype = x.dtype
+ org_forwarded = self.org_forward(x)
+
+ # module dropout
+ if self.module_dropout is not None and self.training:
+ if torch.rand(1) < self.module_dropout:
+ return org_forwarded
+
+ lx = self.lora_down(x.to(self.lora_down.weight.dtype))
+
+ # normal dropout
+ if self.dropout is not None and self.training:
+ lx = torch.nn.functional.dropout(lx, p=self.dropout)
+
+ # rank dropout
+ if self.rank_dropout is not None and self.training:
+ mask = torch.rand((lx.size(0), self.lora_dim), device=lx.device) > self.rank_dropout
+ if len(lx.size()) == 3:
+ mask = mask.unsqueeze(1) # for Text Encoder
+ elif len(lx.size()) == 4:
+ mask = mask.unsqueeze(-1).unsqueeze(-1) # for Conv2d
+ lx = lx * mask
+
+ # scaling for rank dropout: treat as if the rank is changed
+ scale = self.scale * (1.0 / (1.0 - self.rank_dropout)) # redundant for readability
+ else:
+ scale = self.scale
+
+ lx = self.lora_up(lx)
+
+ return org_forwarded.to(weight_dtype) + lx.to(weight_dtype) * self.multiplier * scale
+
+
+def addnet_hash_legacy(b):
+ """Old model hash used by sd-webui-additional-networks for .safetensors format files"""
+ m = hashlib.sha256()
+
+ b.seek(0x100000)
+ m.update(b.read(0x10000))
+ return m.hexdigest()[0:8]
+
+
+def addnet_hash_safetensors(b):
+ """New model hash used by sd-webui-additional-networks for .safetensors format files"""
+ hash_sha256 = hashlib.sha256()
+ blksize = 1024 * 1024
+
+ b.seek(0)
+ header = b.read(8)
+ n = int.from_bytes(header, "little")
+
+ offset = n + 8
+ b.seek(offset)
+ for chunk in iter(lambda: b.read(blksize), b""):
+ hash_sha256.update(chunk)
+
+ return hash_sha256.hexdigest()
+
+
+def precalculate_safetensors_hashes(tensors, metadata):
+ """Precalculate the model hashes needed by sd-webui-additional-networks to
+ save time on indexing the model later."""
+
+ # Because writing user metadata to the file can change the result of
+ # sd_models.model_hash(), only retain the training metadata for purposes of
+ # calculating the hash, as they are meant to be immutable
+ metadata = {k: v for k, v in metadata.items() if k.startswith("ss_")}
+
+ bytes = safetensors.torch.save(tensors, metadata)
+ b = BytesIO(bytes)
+
+ model_hash = addnet_hash_safetensors(b)
+ legacy_hash = addnet_hash_legacy(b)
+ return model_hash, legacy_hash
+
+
+class LoRANetwork(torch.nn.Module):
+ TRANSFORMER_TARGET_REPLACE_MODULE = ["CogVideoXTransformer3DModel"]
+ TEXT_ENCODER_TARGET_REPLACE_MODULE = ["T5LayerSelfAttention", "T5LayerFF", "BertEncoder"]
+ LORA_PREFIX_TRANSFORMER = "lora_unet"
+ LORA_PREFIX_TEXT_ENCODER = "lora_te"
+ def __init__(
+ self,
+ text_encoder: Union[List[T5EncoderModel], T5EncoderModel],
+ unet,
+ multiplier: float = 1.0,
+ lora_dim: int = 4,
+ alpha: float = 1,
+ dropout: Optional[float] = None,
+ module_class: Type[object] = LoRAModule,
+ add_lora_in_attn_temporal: bool = False,
+ varbose: Optional[bool] = False,
+ ) -> None:
+ super().__init__()
+ self.multiplier = multiplier
+
+ self.lora_dim = lora_dim
+ self.alpha = alpha
+ self.dropout = dropout
+
+ print(f"create LoRA network. base dim (rank): {lora_dim}, alpha: {alpha}")
+ print(f"neuron dropout: p={self.dropout}")
+
+ # create module instances
+ def create_modules(
+ is_unet: bool,
+ root_module: torch.nn.Module,
+ target_replace_modules: List[torch.nn.Module],
+ ) -> List[LoRAModule]:
+ prefix = (
+ self.LORA_PREFIX_TRANSFORMER
+ if is_unet
+ else self.LORA_PREFIX_TEXT_ENCODER
+ )
+ loras = []
+ skipped = []
+ for name, module in root_module.named_modules():
+ if module.__class__.__name__ in target_replace_modules:
+ for child_name, child_module in module.named_modules():
+ is_linear = child_module.__class__.__name__ == "Linear" or child_module.__class__.__name__ == "LoRACompatibleLinear"
+ is_conv2d = child_module.__class__.__name__ == "Conv2d" or child_module.__class__.__name__ == "LoRACompatibleConv"
+ is_conv2d_1x1 = is_conv2d and child_module.kernel_size == (1, 1)
+
+ if not add_lora_in_attn_temporal:
+ if "attn_temporal" in child_name:
+ continue
+
+ if is_linear or is_conv2d:
+ lora_name = prefix + "." + name + "." + child_name
+ lora_name = lora_name.replace(".", "_")
+
+ dim = None
+ alpha = None
+
+ if is_linear or is_conv2d_1x1:
+ dim = self.lora_dim
+ alpha = self.alpha
+
+ if dim is None or dim == 0:
+ if is_linear or is_conv2d_1x1:
+ skipped.append(lora_name)
+ continue
+
+ lora = module_class(
+ lora_name,
+ child_module,
+ self.multiplier,
+ dim,
+ alpha,
+ dropout=dropout,
+ )
+ loras.append(lora)
+ return loras, skipped
+
+ text_encoders = text_encoder if type(text_encoder) == list else [text_encoder]
+
+ self.text_encoder_loras = []
+ skipped_te = []
+ for i, text_encoder in enumerate(text_encoders):
+ if text_encoder is not None:
+ text_encoder_loras, skipped = create_modules(False, text_encoder, LoRANetwork.TEXT_ENCODER_TARGET_REPLACE_MODULE)
+ self.text_encoder_loras.extend(text_encoder_loras)
+ skipped_te += skipped
+ print(f"create LoRA for Text Encoder: {len(self.text_encoder_loras)} modules.")
+
+ self.unet_loras, skipped_un = create_modules(True, unet, LoRANetwork.TRANSFORMER_TARGET_REPLACE_MODULE)
+ print(f"create LoRA for U-Net: {len(self.unet_loras)} modules.")
+
+ # assertion
+ names = set()
+ for lora in self.text_encoder_loras + self.unet_loras:
+ assert lora.lora_name not in names, f"duplicated lora name: {lora.lora_name}"
+ names.add(lora.lora_name)
+
+ def apply_to(self, text_encoder, unet, apply_text_encoder=True, apply_unet=True):
+ if apply_text_encoder:
+ print("enable LoRA for text encoder")
+ else:
+ self.text_encoder_loras = []
+
+ if apply_unet:
+ print("enable LoRA for U-Net")
+ else:
+ self.unet_loras = []
+
+ for lora in self.text_encoder_loras + self.unet_loras:
+ lora.apply_to()
+ self.add_module(lora.lora_name, lora)
+
+ def set_multiplier(self, multiplier):
+ self.multiplier = multiplier
+ for lora in self.text_encoder_loras + self.unet_loras:
+ lora.multiplier = self.multiplier
+
+ def load_weights(self, file):
+ if os.path.splitext(file)[1] == ".safetensors":
+ from safetensors.torch import load_file
+
+ weights_sd = load_file(file)
+ else:
+ weights_sd = torch.load(file, map_location="cpu")
+ info = self.load_state_dict(weights_sd, False)
+ return info
+
+ def prepare_optimizer_params(self, text_encoder_lr, unet_lr, default_lr):
+ self.requires_grad_(True)
+ all_params = []
+
+ def enumerate_params(loras):
+ params = []
+ for lora in loras:
+ params.extend(lora.parameters())
+ return params
+
+ if self.text_encoder_loras:
+ param_data = {"params": enumerate_params(self.text_encoder_loras)}
+ if text_encoder_lr is not None:
+ param_data["lr"] = text_encoder_lr
+ all_params.append(param_data)
+
+ if self.unet_loras:
+ param_data = {"params": enumerate_params(self.unet_loras)}
+ if unet_lr is not None:
+ param_data["lr"] = unet_lr
+ all_params.append(param_data)
+
+ return all_params
+
+ def enable_gradient_checkpointing(self):
+ pass
+
+ def get_trainable_params(self):
+ return self.parameters()
+
+ def save_weights(self, file, dtype, metadata):
+ if metadata is not None and len(metadata) == 0:
+ metadata = None
+
+ state_dict = self.state_dict()
+
+ if dtype is not None:
+ for key in list(state_dict.keys()):
+ v = state_dict[key]
+ v = v.detach().clone().to("cpu").to(dtype)
+ state_dict[key] = v
+
+ if os.path.splitext(file)[1] == ".safetensors":
+ from safetensors.torch import save_file
+
+ # Precalculate model hashes to save time on indexing
+ if metadata is None:
+ metadata = {}
+ model_hash, legacy_hash = precalculate_safetensors_hashes(state_dict, metadata)
+ metadata["sshs_model_hash"] = model_hash
+ metadata["sshs_legacy_hash"] = legacy_hash
+
+ save_file(state_dict, file, metadata)
+ else:
+ torch.save(state_dict, file)
+
+def create_network(
+ multiplier: float,
+ network_dim: Optional[int],
+ network_alpha: Optional[float],
+ text_encoder: Union[T5EncoderModel, List[T5EncoderModel]],
+ transformer,
+ neuron_dropout: Optional[float] = None,
+ add_lora_in_attn_temporal: bool = False,
+ **kwargs,
+):
+ if network_dim is None:
+ network_dim = 4 # default
+ if network_alpha is None:
+ network_alpha = 1.0
+
+ network = LoRANetwork(
+ text_encoder,
+ transformer,
+ multiplier=multiplier,
+ lora_dim=network_dim,
+ alpha=network_alpha,
+ dropout=neuron_dropout,
+ add_lora_in_attn_temporal=add_lora_in_attn_temporal,
+ varbose=True,
+ )
+ return network
+
+def merge_lora(transformer, lora_path, multiplier, device='cpu', dtype=torch.float32, state_dict=None):
+ LORA_PREFIX_TRANSFORMER = "lora_unet"
+ LORA_PREFIX_TEXT_ENCODER = "lora_te"
+ if state_dict is None:
+ state_dict = load_file(lora_path, device=device)
+ else:
+ state_dict = state_dict
+ updates = defaultdict(dict)
+ for key, value in state_dict.items():
+ layer, elem = key.split('.', 1)
+ updates[layer][elem] = value
+
+ for layer, elems in updates.items():
+
+ # if "lora_te" in layer:
+ # if transformer_only:
+ # continue
+ # else:
+ # layer_infos = layer.split(LORA_PREFIX_TEXT_ENCODER + "_")[-1].split("_")
+ # curr_layer = pipeline.text_encoder
+ #else:
+ layer_infos = layer.split(LORA_PREFIX_TRANSFORMER + "_")[-1].split("_")
+ curr_layer = transformer
+
+ temp_name = layer_infos.pop(0)
+ while len(layer_infos) > -1:
+ try:
+ curr_layer = curr_layer.__getattr__(temp_name)
+ if len(layer_infos) > 0:
+ temp_name = layer_infos.pop(0)
+ elif len(layer_infos) == 0:
+ break
+ except Exception:
+ if len(layer_infos) == 0:
+ print('Error loading layer')
+ if len(temp_name) > 0:
+ temp_name += "_" + layer_infos.pop(0)
+ else:
+ temp_name = layer_infos.pop(0)
+
+ weight_up = elems['lora_up.weight'].to(dtype)
+ weight_down = elems['lora_down.weight'].to(dtype)
+ if 'alpha' in elems.keys():
+ alpha = elems['alpha'].item() / weight_up.shape[1]
+ else:
+ alpha = 1.0
+
+ curr_layer.weight.data = curr_layer.weight.data.to(device)
+ if len(weight_up.shape) == 4:
+ curr_layer.weight.data += multiplier * alpha * torch.mm(weight_up.squeeze(3).squeeze(2),
+ weight_down.squeeze(3).squeeze(2)).unsqueeze(
+ 2).unsqueeze(3)
+ else:
+ curr_layer.weight.data += multiplier * alpha * torch.mm(weight_up, weight_down)
+
+ return transformer
+
+# TODO: Refactor with merge_lora.
+def unmerge_lora(pipeline, lora_path, multiplier=1, device="cpu", dtype=torch.float32):
+ """Unmerge state_dict in LoRANetwork from the pipeline in diffusers."""
+ LORA_PREFIX_UNET = "lora_unet"
+ LORA_PREFIX_TEXT_ENCODER = "lora_te"
+ state_dict = load_file(lora_path, device=device)
+
+ updates = defaultdict(dict)
+ for key, value in state_dict.items():
+ layer, elem = key.split('.', 1)
+ updates[layer][elem] = value
+
+ for layer, elems in updates.items():
+
+ if "lora_te" in layer:
+ layer_infos = layer.split(LORA_PREFIX_TEXT_ENCODER + "_")[-1].split("_")
+ curr_layer = pipeline.text_encoder
+ else:
+ layer_infos = layer.split(LORA_PREFIX_UNET + "_")[-1].split("_")
+ curr_layer = pipeline.transformer
+
+ temp_name = layer_infos.pop(0)
+ while len(layer_infos) > -1:
+ try:
+ curr_layer = curr_layer.__getattr__(temp_name)
+ if len(layer_infos) > 0:
+ temp_name = layer_infos.pop(0)
+ elif len(layer_infos) == 0:
+ break
+ except Exception:
+ if len(layer_infos) == 0:
+ print('Error loading layer')
+ if len(temp_name) > 0:
+ temp_name += "_" + layer_infos.pop(0)
+ else:
+ temp_name = layer_infos.pop(0)
+
+ weight_up = elems['lora_up.weight'].to(dtype)
+ weight_down = elems['lora_down.weight'].to(dtype)
+ if 'alpha' in elems.keys():
+ alpha = elems['alpha'].item() / weight_up.shape[1]
+ else:
+ alpha = 1.0
+
+ curr_layer.weight.data = curr_layer.weight.data.to(device)
+ if len(weight_up.shape) == 4:
+ curr_layer.weight.data -= multiplier * alpha * torch.mm(weight_up.squeeze(3).squeeze(2),
+ weight_down.squeeze(3).squeeze(2)).unsqueeze(2).unsqueeze(3)
+ else:
+ curr_layer.weight.data -= multiplier * alpha * torch.mm(weight_up, weight_down)
+
+ return pipeline
+
+def load_lora_into_transformer(state_dict, transformer, adapter_name=None):
+ from peft import LoraConfig, inject_adapter_in_model, set_peft_model_state_dict
+ from diffusers.utils.peft_utils import get_peft_kwargs, get_adapter_name
+ from diffusers.utils.import_utils import is_peft_version
+ from diffusers.utils.state_dict_utils import convert_unet_state_dict_to_peft
+ keys = list(state_dict.keys())
+ transformer_keys = [k for k in keys if k.startswith("transformer")]
+ state_dict = {
+ k.replace(f"transformer.", ""): v for k, v in state_dict.items() if k in transformer_keys
+ }
+ if len(state_dict.keys()) > 0:
+ # check with first key if is not in peft format
+ first_key = next(iter(state_dict.keys()))
+ if "lora_A" not in first_key:
+ state_dict = convert_unet_state_dict_to_peft(state_dict)
+ if adapter_name in getattr(transformer, "peft_config", {}):
+ raise ValueError(
+ f"Adapter name {adapter_name} already in use in the transformer - please select a new adapter name."
+ )
+ rank = {}
+ for key, val in state_dict.items():
+ if "lora_B" in key:
+ rank[key] = val.shape[1]
+ lora_config_kwargs = get_peft_kwargs(rank, network_alpha_dict=None, peft_state_dict=state_dict)
+ if "use_dora" in lora_config_kwargs:
+ if lora_config_kwargs["use_dora"] and is_peft_version("<", "0.9.0"):
+ raise ValueError(
+ "You need `peft` 0.9.0 at least to use DoRA-enabled LoRAs. Please upgrade your installation of `peft`."
+ )
+ else:
+ lora_config_kwargs.pop("use_dora")
+ lora_config = LoraConfig(**lora_config_kwargs)
+ # adapter_name
+ if adapter_name is None:
+ adapter_name = get_adapter_name(transformer)
+
+ transformer = inject_adapter_in_model(lora_config, transformer, adapter_name=adapter_name)
+ incompatible_keys = set_peft_model_state_dict(transformer, state_dict, adapter_name)
+ if incompatible_keys is not None:
+ # check only for unexpected keys
+ unexpected_keys = getattr(incompatible_keys, "unexpected_keys", None)
+ if unexpected_keys:
+ print(
+ f"Loading adapter weights from state_dict led to unexpected keys not found in the model: "
+ f" {unexpected_keys}. "
+ )
+ return transformer
\ No newline at end of file
diff --git a/cogvideox_fun/pipeline_cogvideox_control.py b/cogvideox_fun/pipeline_cogvideox_control.py
index 5dcda4a..966e0ee 100644
--- a/cogvideox_fun/pipeline_cogvideox_control.py
+++ b/cogvideox_fun/pipeline_cogvideox_control.py
@@ -214,7 +214,8 @@ class CogVideoX_Fun_Pipeline_Control(VideoSysPipeline):
set_pab_manager(pab_config)
def prepare_latents(
- self, batch_size, num_channels_latents, num_frames, height, width, dtype, device, generator, latents=None
+ self, batch_size, num_channels_latents, num_frames, height, width, dtype, device, generator, timesteps, denoise_strength, num_inference_steps,
+ latents=None, freenoise=True, context_size=None, context_overlap=None
):
shape = (
batch_size,
@@ -228,15 +229,62 @@ class CogVideoX_Fun_Pipeline_Control(VideoSysPipeline):
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
)
+ noise = randn_tensor(shape, generator=generator, device=torch.device("cpu"), dtype=self.vae.dtype)
+ if freenoise:
+ print("Applying FreeNoise")
+ # code and comments from AnimateDiff-Evolved by Kosinkadink (https://github.com/Kosinkadink/ComfyUI-AnimateDiff-Evolved)
+ video_length = num_frames // 4
+ delta = context_size - context_overlap
+ for start_idx in range(0, video_length-context_size, delta):
+ # start_idx corresponds to the beginning of a context window
+ # goal: place shuffled in the delta region right after the end of the context window
+ # if space after context window is not enough to place the noise, adjust and finish
+ place_idx = start_idx + context_size
+ # if place_idx is outside the valid indexes, we are already finished
+ if place_idx >= video_length:
+ break
+ end_idx = place_idx - 1
+ #print("video_length:", video_length, "start_idx:", start_idx, "end_idx:", end_idx, "place_idx:", place_idx, "delta:", delta)
+ # if there is not enough room to copy delta amount of indexes, copy limited amount and finish
+ if end_idx + delta >= video_length:
+ final_delta = video_length - place_idx
+ # generate list of indexes in final delta region
+ list_idx = torch.tensor(list(range(start_idx,start_idx+final_delta)), device=torch.device("cpu"), dtype=torch.long)
+ # shuffle list
+ list_idx = list_idx[torch.randperm(final_delta, generator=generator)]
+ # apply shuffled indexes
+ noise[:, place_idx:place_idx + final_delta, :, :, :] = noise[:, list_idx, :, :, :]
+ break
+ # otherwise, do normal behavior
+ # generate list of indexes in delta region
+ list_idx = torch.tensor(list(range(start_idx,start_idx+delta)), device=torch.device("cpu"), dtype=torch.long)
+ # shuffle list
+ list_idx = list_idx[torch.randperm(delta, generator=generator)]
+ # apply shuffled indexes
+ #print("place_idx:", place_idx, "delta:", delta, "list_idx:", list_idx)
+ noise[:, place_idx:place_idx + delta, :, :, :] = noise[:, list_idx, :, :, :]
if latents is None:
- latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
+ latents = noise.to(device)
else:
latents = latents.to(device)
+ timesteps, num_inference_steps = self.get_timesteps(num_inference_steps, denoise_strength, device)
+ latent_timestep = timesteps[:1]
+
+ noise = randn_tensor(shape, generator=generator, device=device, dtype=self.vae.dtype)
+ frames_needed = noise.shape[1]
+ current_frames = latents.shape[1]
+
+ if frames_needed > current_frames:
+ repeat_factor = frames_needed // current_frames
+ additional_frame = torch.randn((latents.size(0), repeat_factor, latents.size(2), latents.size(3), latents.size(4)), dtype=latents.dtype, device=latents.device)
+ latents = torch.cat((latents, additional_frame), dim=1)
+ elif frames_needed < current_frames:
+ latents = latents[:, :frames_needed, :, :, :]
- # scale the initial noise by the standard deviation required by the scheduler
- latents = latents * self.scheduler.init_noise_sigma
- return latents
+ latents = self.scheduler.add_noise(latents, noise, latent_timestep)
+ latents = latents * self.scheduler.init_noise_sigma # scale the initial noise by the standard deviation required by the scheduler
+ return latents, timesteps, noise
def prepare_control_latents(
self, mask, masked_image, batch_size, height, width, dtype, device, generator, do_classifier_free_guidance
@@ -300,6 +348,16 @@ class CogVideoX_Fun_Pipeline_Control(VideoSysPipeline):
if accepts_generator:
extra_step_kwargs["generator"] = generator
return extra_step_kwargs
+
+ def _gaussian_weights(self, t_tile_length, t_batch_size):
+ from numpy import pi, exp, sqrt
+
+ var = 0.01
+ midpoint = (t_tile_length - 1) / 2 # -1 because index goes from 0 to latent_width - 1
+ t_probs = [exp(-(t-midpoint)*(t-midpoint)/(t_tile_length*t_tile_length)/(2*var)) / sqrt(2*pi*var) for t in range(t_tile_length)]
+ weights = torch.tensor(t_probs)
+ weights = weights.unsqueeze(0).unsqueeze(2).unsqueeze(3).unsqueeze(4).repeat(1, t_batch_size,1, 1, 1)
+ return weights
# Copied from diffusers.pipelines.latte.pipeline_latte.LattePipeline.check_inputs
def check_inputs(
@@ -372,7 +430,10 @@ class CogVideoX_Fun_Pipeline_Control(VideoSysPipeline):
width: int,
num_frames: int,
device: torch.device,
- ) -> Tuple[torch.Tensor, torch.Tensor]:
+ start_frame: Optional[int] = None,
+ end_frame: Optional[int] = None,
+ context_frames: Optional[int] = None,
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
grid_height = height // (self.vae_scale_factor_spatial * self.transformer.config.patch_size)
grid_width = width // (self.vae_scale_factor_spatial * self.transformer.config.patch_size)
base_size_width = 720 // (self.vae_scale_factor_spatial * self.transformer.config.patch_size)
@@ -388,6 +449,19 @@ class CogVideoX_Fun_Pipeline_Control(VideoSysPipeline):
temporal_size=num_frames,
use_real=True,
)
+
+ if start_frame is not None or context_frames is not None:
+ freqs_cos = freqs_cos.view(num_frames, grid_height * grid_width, -1)
+ freqs_sin = freqs_sin.view(num_frames, grid_height * grid_width, -1)
+ if context_frames is not None:
+ freqs_cos = freqs_cos[context_frames]
+ freqs_sin = freqs_sin[context_frames]
+ else:
+ freqs_cos = freqs_cos[start_frame:end_frame]
+ freqs_sin = freqs_sin[start_frame:end_frame]
+
+ freqs_cos = freqs_cos.view(-1, freqs_cos.shape[-1])
+ freqs_sin = freqs_sin.view(-1, freqs_sin.shape[-1])
freqs_cos = freqs_cos.to(device=device)
freqs_sin = freqs_sin.to(device=device)
@@ -430,6 +504,7 @@ class CogVideoX_Fun_Pipeline_Control(VideoSysPipeline):
timesteps: Optional[List[int]] = None,
guidance_scale: float = 6,
use_dynamic_cfg: bool = False,
+ denoise_strength: float = 1.0,
num_videos_per_prompt: int = 1,
eta: float = 0.0,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
@@ -447,6 +522,12 @@ class CogVideoX_Fun_Pipeline_Control(VideoSysPipeline):
control_strength: float = 1.0,
control_start_percent: float = 0.0,
control_end_percent: float = 1.0,
+ scheduler_name: str = "DPM",
+ context_schedule: Optional[str] = None,
+ context_frames: Optional[int] = None,
+ context_stride: Optional[int] = None,
+ context_overlap: Optional[int] = None,
+ freenoise: Optional[bool] = True,
) -> Union[CogVideoX_Fun_PipelineOutput, Tuple]:
"""
Function invoked when calling the pipeline for generation.
@@ -524,10 +605,10 @@ class CogVideoX_Fun_Pipeline_Control(VideoSysPipeline):
`tuple`. When returning a tuple, the first element is a list with the generated images.
"""
- if num_frames > 49:
- raise ValueError(
- "The number of frames must be less than 49 for now due to static positional embeddings. This will be updated in the future to remove this limitation."
- )
+ # if num_frames > 49:
+ # raise ValueError(
+ # "The number of frames must be less than 49 for now due to static positional embeddings. This will be updated in the future to remove this limitation."
+ # )
if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)):
callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs
@@ -576,7 +657,7 @@ class CogVideoX_Fun_Pipeline_Control(VideoSysPipeline):
# 5. Prepare latents.
latent_channels = self.vae.config.latent_channels
- latents = self.prepare_latents(
+ latents, timesteps, noise = self.prepare_latents(
batch_size * num_videos_per_prompt,
latent_channels,
num_frames,
@@ -585,31 +666,20 @@ class CogVideoX_Fun_Pipeline_Control(VideoSysPipeline):
self.vae.dtype,
device,
generator,
+ timesteps,
+ denoise_strength,
+ num_inference_steps,
latents,
+ context_size=context_frames,
+ context_overlap=context_overlap,
+ freenoise=freenoise,
)
if comfyui_progressbar:
pbar.update(1)
- if control_video is not None:
- video_length = control_video.shape[2]
- control_video = self.image_processor.preprocess(rearrange(control_video, "b c f h w -> (b f) c h w"), height=height, width=width)
- control_video = control_video.to(dtype=torch.float32)
- control_video = rearrange(control_video, "(b f) c h w -> b c f h w", f=video_length)
- else:
- control_video = None
- control_video_latents = self.prepare_control_latents(
- None,
- control_video,
- batch_size,
- height,
- width,
- self.vae.dtype,
- device,
- generator,
- do_classifier_free_guidance
- )[1]
+
control_video_latents_input = (
- torch.cat([control_video_latents] * 2) if do_classifier_free_guidance else control_video_latents
+ torch.cat([control_video] * 2) if do_classifier_free_guidance else control_video
)
control_latents = rearrange(control_video_latents_input, "b c f h w -> b f c h w")
@@ -621,16 +691,37 @@ class CogVideoX_Fun_Pipeline_Control(VideoSysPipeline):
# 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
- # 7. Create rotary embeds if required
- image_rotary_emb = (
- self._prepare_rotary_positional_embeddings(height, width, latents.size(1), device)
- if self.transformer.config.use_rotary_positional_embeddings
- else None
- )
+
# 8. Denoising loop
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
+ # 8.5. Temporal tiling prep
+ if context_schedule is not None and context_schedule == "temporal_tiling":
+ t_tile_length = context_frames
+ t_tile_overlap = context_overlap
+ t_tile_weights = self._gaussian_weights(t_tile_length=t_tile_length, t_batch_size=1).to(latents.device).to(self.vae.dtype)
+ use_temporal_tiling = True
+ print("Temporal tiling enabled")
+ elif context_schedule is not None:
+ print(f"Context schedule enabled: {context_frames} frames, {context_stride} stride, {context_overlap} overlap")
+ use_temporal_tiling = False
+ use_context_schedule = True
+ from .context import get_context_scheduler
+ context = get_context_scheduler(context_schedule)
+
+ else:
+ use_temporal_tiling = False
+ use_context_schedule = False
+ print("Temporal tiling and context schedule disabled")
+ # 7. Create rotary embeds if required
+ image_rotary_emb = (
+ self._prepare_rotary_positional_embeddings(height, width, latents.size(1), device)
+ if self.transformer.config.use_rotary_positional_embeddings
+ else None
+ )
+
+
with self.progress_bar(total=num_inference_steps) as progress_bar:
# for DPM-solver++
old_pred_original_sample = None
@@ -638,69 +729,237 @@ class CogVideoX_Fun_Pipeline_Control(VideoSysPipeline):
if self.interrupt:
continue
- 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)
+ if use_temporal_tiling and isinstance(self.scheduler, CogVideoXDDIMScheduler):
+ #temporal tiling code based on https://github.com/mayuelala/FollowYourEmoji/blob/main/models/video_pipeline.py
+ # =====================================================
+ grid_ts = 0
+ cur_t = 0
+ while cur_t < latents.shape[1]:
+ cur_t = max(grid_ts * t_tile_length - t_tile_overlap * grid_ts, 0) + t_tile_length
+ grid_ts += 1
- # Calculate the current step percentage
- current_step_percentage = i / num_inference_steps
+ all_t = latents.shape[1]
+ latents_all_list = []
+ # =====================================================
- # Determine if control_latents should be applied
- apply_control = control_start_percent <= current_step_percentage <= control_end_percent
- current_control_latents = control_latents if apply_control else torch.zeros_like(control_latents)
+ image_rotary_emb = (
+ self._prepare_rotary_positional_embeddings(height, width, context_frames, device)
+ if self.transformer.config.use_rotary_positional_embeddings
+ else None
+ )
- # broadcast to batch dimension in a way that's compatible with ONNX/Core ML
- timestep = t.expand(latent_model_input.shape[0])
+ for t_i in range(grid_ts):
+ if t_i < grid_ts - 1:
+ ofs_t = max(t_i * t_tile_length - t_tile_overlap * t_i, 0)
+ if t_i == grid_ts - 1:
+ ofs_t = all_t - t_tile_length
- # predict noise model_output
- noise_pred = self.transformer(
- hidden_states=latent_model_input,
- encoder_hidden_states=prompt_embeds,
- timestep=timestep,
- image_rotary_emb=image_rotary_emb,
- return_dict=False,
- control_latents=current_control_latents,
- )[0]
- noise_pred = noise_pred.float()
+ input_start_t = ofs_t
+ input_end_t = ofs_t + t_tile_length
- # perform guidance
- if use_dynamic_cfg:
- self._guidance_scale = 1 + guidance_scale * (
- (1 - math.cos(math.pi * ((num_inference_steps - t.item()) / num_inference_steps) ** 5.0)) / 2
- )
- if do_classifier_free_guidance:
- noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
- noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_text - noise_pred_uncond)
+ latents_tile = latents[:, input_start_t:input_end_t,:, :, :]
+ control_latents_tile = control_latents[:, input_start_t:input_end_t, :, :, :]
- # compute the previous noisy sample x_t -> x_t-1
- if not isinstance(self.scheduler, CogVideoXDPMScheduler):
- latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
+ latent_model_input_tile = torch.cat([latents_tile] * 2) if do_classifier_free_guidance else latents_tile
+ latent_model_input_tile = self.scheduler.scale_model_input(latent_model_input_tile, t)
+
+ #t_input = t[None].to(device)
+ t_input = t.expand(latent_model_input_tile.shape[0]) # broadcast to batch dimension in a way that's compatible with ONNX/Core ML
+
+ # predict noise model_output
+ noise_pred = self.transformer(
+ hidden_states=latent_model_input_tile,
+ encoder_hidden_states=prompt_embeds,
+ timestep=t_input,
+ image_rotary_emb=image_rotary_emb,
+ return_dict=False,
+ control_latents=control_latents_tile,
+ )[0]
+ noise_pred = noise_pred.float()
+
+ if do_classifier_free_guidance:
+ noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
+ noise_pred = noise_pred_uncond + self._guidance_scale * (noise_pred_text - noise_pred_uncond)
+
+ # compute the previous noisy sample x_t -> x_t-1
+ latents_tile = self.scheduler.step(noise_pred, t, latents_tile.to(self.vae.dtype), **extra_step_kwargs, return_dict=False)[0]
+ latents_all_list.append(latents_tile)
+
+ # ==========================================
+ latents_all = torch.zeros(latents.shape, device=latents.device, dtype=self.vae.dtype)
+ contributors = torch.zeros(latents.shape, device=latents.device, dtype=self.vae.dtype)
+ # Add each tile contribution to overall latents
+ for t_i in range(grid_ts):
+ if t_i < grid_ts - 1:
+ ofs_t = max(t_i * t_tile_length - t_tile_overlap * t_i, 0)
+ if t_i == grid_ts - 1:
+ ofs_t = all_t - t_tile_length
+
+ input_start_t = ofs_t
+ input_end_t = ofs_t + t_tile_length
+
+ latents_all[:, input_start_t:input_end_t,:, :, :] += latents_all_list[t_i] * t_tile_weights
+ contributors[:, input_start_t:input_end_t,:, :, :] += t_tile_weights
+
+ latents_all /= contributors
+
+ latents = latents_all
+
+ if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
+ progress_bar.update()
+ pbar.update(1)
+ # ==========================================
+ elif use_context_schedule:
+
+ 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
+
+ # Determine if control_latents should be applied
+ apply_control = control_start_percent <= current_step_percentage <= control_end_percent
+ current_control_latents = control_latents if apply_control else torch.zeros_like(control_latents)
+
+ # 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,
+ ))
+ 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)
+ if self.transformer.config.use_rotary_positional_embeddings
+ else None
+ )
+
+ for c in context_queue:
+ partial_latent_model_input = latent_model_input[:, c, :, :, :]
+ partial_control_latents = current_control_latents[:, 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,
+ 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()
+
+ noise_pred /= counter
+ if do_classifier_free_guidance:
+ noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
+ noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_text - noise_pred_uncond)
+
+ # compute the previous noisy sample x_t -> x_t-1
+ if not isinstance(self.scheduler, CogVideoXDPMScheduler):
+ latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
+ else:
+ latents, old_pred_original_sample = self.scheduler.step(
+ noise_pred,
+ old_pred_original_sample,
+ t,
+ timesteps[i - 1] if i > 0 else None,
+ latents,
+ **extra_step_kwargs,
+ return_dict=False,
+ )
+ latents = latents.to(prompt_embeds.dtype)
+
+ # call the callback, if provided
+ if callback_on_step_end is not None:
+ callback_kwargs = {}
+ for k in callback_on_step_end_tensor_inputs:
+ callback_kwargs[k] = locals()[k]
+ callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
+
+ latents = callback_outputs.pop("latents", latents)
+ prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
+ negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds)
+
+ if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
+ progress_bar.update()
+ if comfyui_progressbar:
+ pbar.update(1)
else:
- latents, old_pred_original_sample = self.scheduler.step(
- noise_pred,
- old_pred_original_sample,
- t,
- timesteps[i - 1] if i > 0 else None,
- latents,
- **extra_step_kwargs,
+ 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
+
+ # Determine if control_latents should be applied
+ apply_control = control_start_percent <= current_step_percentage <= control_end_percent
+ current_control_latents = control_latents if apply_control else torch.zeros_like(control_latents)
+
+ # broadcast to batch dimension in a way that's compatible with ONNX/Core ML
+ timestep = t.expand(latent_model_input.shape[0])
+
+ # predict noise model_output
+ noise_pred = self.transformer(
+ hidden_states=latent_model_input,
+ encoder_hidden_states=prompt_embeds,
+ timestep=timestep,
+ image_rotary_emb=image_rotary_emb,
return_dict=False,
- )
- latents = latents.to(prompt_embeds.dtype)
+ control_latents=current_control_latents,
+ )[0]
+ noise_pred = noise_pred.float()
- # call the callback, if provided
- if callback_on_step_end is not None:
- callback_kwargs = {}
- for k in callback_on_step_end_tensor_inputs:
- callback_kwargs[k] = locals()[k]
- callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
+ if do_classifier_free_guidance:
+ noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
+ noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_text - noise_pred_uncond)
- latents = callback_outputs.pop("latents", latents)
- prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
- negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds)
+ # compute the previous noisy sample x_t -> x_t-1
+ if not isinstance(self.scheduler, CogVideoXDPMScheduler):
+ latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
+ else:
+ latents, old_pred_original_sample = self.scheduler.step(
+ noise_pred,
+ old_pred_original_sample,
+ t,
+ timesteps[i - 1] if i > 0 else None,
+ latents,
+ **extra_step_kwargs,
+ return_dict=False,
+ )
+ latents = latents.to(prompt_embeds.dtype)
- if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
- progress_bar.update()
- if comfyui_progressbar:
- pbar.update(1)
+ # call the callback, if provided
+ if callback_on_step_end is not None:
+ callback_kwargs = {}
+ for k in callback_on_step_end_tensor_inputs:
+ callback_kwargs[k] = locals()[k]
+ callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
+
+ latents = callback_outputs.pop("latents", latents)
+ prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
+ negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds)
+
+ if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
+ progress_bar.update()
+ if comfyui_progressbar:
+ pbar.update(1)
# if output_type == "numpy":
# video = self.decode_latents(latents)
diff --git a/cogvideox_fun/pipeline_cogvideox_inpaint.py b/cogvideox_fun/pipeline_cogvideox_inpaint.py
index 5e56432..4c9d505 100644
--- a/cogvideox_fun/pipeline_cogvideox_inpaint.py
+++ b/cogvideox_fun/pipeline_cogvideox_inpaint.py
@@ -277,6 +277,9 @@ class CogVideoX_Fun_Pipeline_Inpaint(VideoSysPipeline):
is_strength_max=True,
return_noise=False,
return_video_latents=False,
+ context_size=None,
+ context_overlap=None,
+ freenoise=False,
):
shape = (
batch_size,
@@ -309,11 +312,47 @@ class CogVideoX_Fun_Pipeline_Inpaint(VideoSysPipeline):
video_latents = rearrange(video_latents, "b c f h w -> b f c h w")
if latents is None:
- noise = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
+ noise = randn_tensor(shape, generator=generator, device=torch.device("cpu"), dtype=dtype)
+ if freenoise:
+ print("Applying FreeNoise")
+ # code and comments from AnimateDiff-Evolved by Kosinkadink (https://github.com/Kosinkadink/ComfyUI-AnimateDiff-Evolved)
+ video_length = video_length // 4
+ delta = context_size - context_overlap
+ for start_idx in range(0, video_length-context_size, delta):
+ # start_idx corresponds to the beginning of a context window
+ # goal: place shuffled in the delta region right after the end of the context window
+ # if space after context window is not enough to place the noise, adjust and finish
+ place_idx = start_idx + context_size
+ # if place_idx is outside the valid indexes, we are already finished
+ if place_idx >= video_length:
+ break
+ end_idx = place_idx - 1
+ #print("video_length:", video_length, "start_idx:", start_idx, "end_idx:", end_idx, "place_idx:", place_idx, "delta:", delta)
+
+ # if there is not enough room to copy delta amount of indexes, copy limited amount and finish
+ if end_idx + delta >= video_length:
+ final_delta = video_length - place_idx
+ # generate list of indexes in final delta region
+ list_idx = torch.tensor(list(range(start_idx,start_idx+final_delta)), device=torch.device("cpu"), dtype=torch.long)
+ # shuffle list
+ list_idx = list_idx[torch.randperm(final_delta, generator=generator)]
+ # apply shuffled indexes
+ noise[:, place_idx:place_idx + final_delta, :, :, :] = noise[:, list_idx, :, :, :]
+ break
+ # otherwise, do normal behavior
+ # generate list of indexes in delta region
+ list_idx = torch.tensor(list(range(start_idx,start_idx+delta)), device=torch.device("cpu"), dtype=torch.long)
+ # shuffle list
+ list_idx = list_idx[torch.randperm(delta, generator=generator)]
+ # apply shuffled indexes
+ #print("place_idx:", place_idx, "delta:", delta, "list_idx:", list_idx)
+ noise[:, place_idx:place_idx + delta, :, :, :] = noise[:, list_idx, :, :, :]
+
# if strength is 1. then initialise the latents to noise, else initial to image + noise
latents = noise if is_strength_max else self.scheduler.add_noise(video_latents, noise, timestep)
# if pure noise then scale the initial latents by the Scheduler's init sigma
latents = latents * self.scheduler.init_noise_sigma if is_strength_max else latents
+ latents = latents.to(device)
else:
noise = latents.to(device)
latents = noise * self.scheduler.init_noise_sigma
@@ -393,7 +432,17 @@ class CogVideoX_Fun_Pipeline_Inpaint(VideoSysPipeline):
if accepts_generator:
extra_step_kwargs["generator"] = generator
return extra_step_kwargs
+
+ def _gaussian_weights(self, t_tile_length, t_batch_size):
+ from numpy import pi, exp, sqrt
+ var = 0.01
+ midpoint = (t_tile_length - 1) / 2 # -1 because index goes from 0 to latent_width - 1
+ t_probs = [exp(-(t-midpoint)*(t-midpoint)/(t_tile_length*t_tile_length)/(2*var)) / sqrt(2*pi*var) for t in range(t_tile_length)]
+ weights = torch.tensor(t_probs)
+ weights = weights.unsqueeze(0).unsqueeze(2).unsqueeze(3).unsqueeze(4).repeat(1, t_batch_size,1, 1, 1)
+ return weights
+
# Copied from diffusers.pipelines.latte.pipeline_latte.LattePipeline.check_inputs
def check_inputs(
self,
@@ -465,7 +514,10 @@ class CogVideoX_Fun_Pipeline_Inpaint(VideoSysPipeline):
width: int,
num_frames: int,
device: torch.device,
- ) -> Tuple[torch.Tensor, torch.Tensor]:
+ start_frame: Optional[int] = None,
+ end_frame: Optional[int] = None,
+ context_frames: Optional[int] = None,
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
grid_height = height // (self.vae_scale_factor_spatial * self.transformer.config.patch_size)
grid_width = width // (self.vae_scale_factor_spatial * self.transformer.config.patch_size)
base_size_width = 720 // (self.vae_scale_factor_spatial * self.transformer.config.patch_size)
@@ -481,6 +533,19 @@ class CogVideoX_Fun_Pipeline_Inpaint(VideoSysPipeline):
temporal_size=num_frames,
use_real=True,
)
+
+ if start_frame is not None or context_frames is not None:
+ freqs_cos = freqs_cos.view(num_frames, grid_height * grid_width, -1)
+ freqs_sin = freqs_sin.view(num_frames, grid_height * grid_width, -1)
+ if context_frames is not None:
+ freqs_cos = freqs_cos[context_frames]
+ freqs_sin = freqs_sin[context_frames]
+ else:
+ freqs_cos = freqs_cos[start_frame:end_frame]
+ freqs_sin = freqs_sin[start_frame:end_frame]
+
+ freqs_cos = freqs_cos.view(-1, freqs_cos.shape[-1])
+ freqs_sin = freqs_sin.view(-1, freqs_sin.shape[-1])
freqs_cos = freqs_cos.to(device=device)
freqs_sin = freqs_sin.to(device=device)
@@ -540,6 +605,11 @@ class CogVideoX_Fun_Pipeline_Inpaint(VideoSysPipeline):
strength: float = 1,
noise_aug_strength: float = 0.0563,
comfyui_progressbar: bool = False,
+ context_schedule: Optional[str] = None,
+ context_frames: Optional[int] = None,
+ context_stride: Optional[int] = None,
+ context_overlap: Optional[int] = None,
+ freenoise: Optional[bool] = True,
) -> Union[CogVideoX_Fun_PipelineOutput, Tuple]:
"""
Function invoked when calling the pipeline for generation.
@@ -617,10 +687,10 @@ class CogVideoX_Fun_Pipeline_Inpaint(VideoSysPipeline):
`tuple`. When returning a tuple, the first element is a list with the generated images.
"""
- if num_frames > 49:
- raise ValueError(
- "The number of frames must be less than 49 for now due to static positional embeddings. This will be updated in the future to remove this limitation."
- )
+ # if num_frames > 49:
+ # raise ValueError(
+ # "The number of frames must be less than 49 for now due to static positional embeddings. This will be updated in the future to remove this limitation."
+ # )
if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)):
callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs
@@ -704,6 +774,9 @@ class CogVideoX_Fun_Pipeline_Inpaint(VideoSysPipeline):
is_strength_max=is_strength_max,
return_noise=True,
return_video_latents=return_image_latents,
+ context_size=context_frames,
+ context_overlap=context_overlap,
+ freenoise=freenoise,
)
if return_image_latents:
latents, noise, image_latents = latents_outputs
@@ -794,11 +867,28 @@ class CogVideoX_Fun_Pipeline_Inpaint(VideoSysPipeline):
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
# 7. Create rotary embeds if required
- image_rotary_emb = (
- self._prepare_rotary_positional_embeddings(height, width, latents.size(1), device)
- if self.transformer.config.use_rotary_positional_embeddings
- else None
- )
+ if context_schedule is not None and context_schedule == "temporal_tiling":
+ t_tile_length = context_frames
+ t_tile_overlap = context_overlap
+ t_tile_weights = self._gaussian_weights(t_tile_length=t_tile_length, t_batch_size=1).to(latents.device).to(self.vae.dtype)
+ use_temporal_tiling = True
+ print("Temporal tiling enabled")
+ elif context_schedule is not None:
+ print(f"Context schedule enabled: {context_frames} frames, {context_stride} stride, {context_overlap} overlap")
+ use_temporal_tiling = False
+ use_context_schedule = True
+ from .context import get_context_scheduler
+ context = get_context_scheduler(context_schedule)
+ else:
+ use_temporal_tiling = False
+ use_context_schedule = False
+ print("Temporal tiling and context schedule disabled")
+ # 7. Create rotary embeds if required
+ image_rotary_emb = (
+ self._prepare_rotary_positional_embeddings(height, width, latents.size(1), device)
+ if self.transformer.config.use_rotary_positional_embeddings
+ else None
+ )
# 8. Denoising loop
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
@@ -809,63 +899,232 @@ class CogVideoX_Fun_Pipeline_Inpaint(VideoSysPipeline):
for i, t in enumerate(timesteps):
if self.interrupt:
continue
+ if use_temporal_tiling and isinstance(self.scheduler, CogVideoXDDIMScheduler):
+ #temporal tiling code based on https://github.com/mayuelala/FollowYourEmoji/blob/main/models/video_pipeline.py
+ # =====================================================
+ grid_ts = 0
+ cur_t = 0
+ while cur_t < latents.shape[1]:
+ cur_t = max(grid_ts * t_tile_length - t_tile_overlap * grid_ts, 0) + t_tile_length
+ grid_ts += 1
- 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)
+ all_t = latents.shape[1]
+ latents_all_list = []
+ # =====================================================
- # broadcast to batch dimension in a way that's compatible with ONNX/Core ML
- timestep = t.expand(latent_model_input.shape[0])
+ image_rotary_emb = (
+ self._prepare_rotary_positional_embeddings(height, width, t_tile_length, device)
+ if self.transformer.config.use_rotary_positional_embeddings
+ else None
+ )
- # predict noise model_output
- noise_pred = self.transformer(
- hidden_states=latent_model_input,
- encoder_hidden_states=prompt_embeds,
- timestep=timestep,
- image_rotary_emb=image_rotary_emb,
- return_dict=False,
- inpaint_latents=inpaint_latents,
- )[0]
- noise_pred = noise_pred.float()
+ for t_i in range(grid_ts):
+ if t_i < grid_ts - 1:
+ ofs_t = max(t_i * t_tile_length - t_tile_overlap * t_i, 0)
+ if t_i == grid_ts - 1:
+ ofs_t = all_t - t_tile_length
- # perform guidance
- if use_dynamic_cfg:
- self._guidance_scale = 1 + guidance_scale * (
- (1 - math.cos(math.pi * ((num_inference_steps - t.item()) / num_inference_steps) ** 5.0)) / 2
- )
- if do_classifier_free_guidance:
- noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
- noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_text - noise_pred_uncond)
+ input_start_t = ofs_t
+ input_end_t = ofs_t + t_tile_length
- # compute the previous noisy sample x_t -> x_t-1
- if not isinstance(self.scheduler, CogVideoXDPMScheduler):
- latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
+ latents_tile = latents[:, input_start_t:input_end_t,:, :, :]
+ inpaint_latents_tile = inpaint_latents[:, input_start_t:input_end_t, :, :, :]
+
+ latent_model_input_tile = torch.cat([latents_tile] * 2) if do_classifier_free_guidance else latents_tile
+ latent_model_input_tile = self.scheduler.scale_model_input(latent_model_input_tile, t)
+
+ #t_input = t[None].to(device)
+ t_input = t.expand(latent_model_input_tile.shape[0]) # broadcast to batch dimension in a way that's compatible with ONNX/Core ML
+
+ # predict noise model_output
+ noise_pred = self.transformer(
+ hidden_states=latent_model_input_tile,
+ encoder_hidden_states=prompt_embeds,
+ timestep=t_input,
+ image_rotary_emb=image_rotary_emb,
+ return_dict=False,
+ inpaint_latents=inpaint_latents_tile,
+ )[0]
+ noise_pred = noise_pred.float()
+
+ if do_classifier_free_guidance:
+ noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
+ noise_pred = noise_pred_uncond + self._guidance_scale * (noise_pred_text - noise_pred_uncond)
+
+ # compute the previous noisy sample x_t -> x_t-1
+ latents_tile = self.scheduler.step(noise_pred, t, latents_tile.to(self.vae.dtype), **extra_step_kwargs, return_dict=False)[0]
+ latents_all_list.append(latents_tile)
+
+ # ==========================================
+ latents_all = torch.zeros(latents.shape, device=latents.device, dtype=self.vae.dtype)
+ contributors = torch.zeros(latents.shape, device=latents.device, dtype=self.vae.dtype)
+ # Add each tile contribution to overall latents
+ for t_i in range(grid_ts):
+ if t_i < grid_ts - 1:
+ ofs_t = max(t_i * t_tile_length - t_tile_overlap * t_i, 0)
+ if t_i == grid_ts - 1:
+ ofs_t = all_t - t_tile_length
+
+ input_start_t = ofs_t
+ input_end_t = ofs_t + t_tile_length
+
+ latents_all[:, input_start_t:input_end_t,:, :, :] += latents_all_list[t_i] * t_tile_weights
+ contributors[:, input_start_t:input_end_t,:, :, :] += t_tile_weights
+
+ latents_all /= contributors
+
+ latents = latents_all
+
+ if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
+ progress_bar.update()
+ pbar.update(1)
+ # ==========================================
+ elif use_context_schedule:
+
+ 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])
+
+ context_queue = list(context(
+ i, num_inference_steps, latents.shape[1], context_frames, context_stride, context_overlap,
+ ))
+ 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)
+ if self.transformer.config.use_rotary_positional_embeddings
+ else None
+ )
+
+ for c in context_queue:
+ partial_latent_model_input = latent_model_input[:, c, :, :, :]
+ partial_inpaint_latents = inpaint_latents[:, c, :, :, :]
+ partial_inpaint_latents[:, 0, :, :, :] = inpaint_latents[:, 0, :, :, :]
+
+ # 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,
+ inpaint_latents=partial_inpaint_latents,
+ )[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()
+
+ noise_pred /= counter
+ if do_classifier_free_guidance:
+ noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
+ noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_text - noise_pred_uncond)
+
+ # compute the previous noisy sample x_t -> x_t-1
+ if not isinstance(self.scheduler, CogVideoXDPMScheduler):
+ latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
+ else:
+ latents, old_pred_original_sample = self.scheduler.step(
+ noise_pred,
+ old_pred_original_sample,
+ t,
+ timesteps[i - 1] if i > 0 else None,
+ latents,
+ **extra_step_kwargs,
+ return_dict=False,
+ )
+ latents = latents.to(prompt_embeds.dtype)
+
+ # call the callback, if provided
+ if callback_on_step_end is not None:
+ callback_kwargs = {}
+ for k in callback_on_step_end_tensor_inputs:
+ callback_kwargs[k] = locals()[k]
+ callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
+
+ latents = callback_outputs.pop("latents", latents)
+ prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
+ negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds)
+
+ if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
+ progress_bar.update()
+ if comfyui_progressbar:
+ pbar.update(1)
+
else:
- latents, old_pred_original_sample = self.scheduler.step(
- noise_pred,
- old_pred_original_sample,
- t,
- timesteps[i - 1] if i > 0 else None,
- latents,
- **extra_step_kwargs,
+ 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)
+
+ # broadcast to batch dimension in a way that's compatible with ONNX/Core ML
+ timestep = t.expand(latent_model_input.shape[0])
+
+ # predict noise model_output
+ noise_pred = self.transformer(
+ hidden_states=latent_model_input,
+ encoder_hidden_states=prompt_embeds,
+ timestep=timestep,
+ image_rotary_emb=image_rotary_emb,
return_dict=False,
- )
- latents = latents.to(prompt_embeds.dtype)
+ inpaint_latents=inpaint_latents,
+ )[0]
+ noise_pred = noise_pred.float()
- # call the callback, if provided
- if callback_on_step_end is not None:
- callback_kwargs = {}
- for k in callback_on_step_end_tensor_inputs:
- callback_kwargs[k] = locals()[k]
- callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
+ # perform guidance
+ if use_dynamic_cfg:
+ self._guidance_scale = 1 + guidance_scale * (
+ (1 - math.cos(math.pi * ((num_inference_steps - t.item()) / num_inference_steps) ** 5.0)) / 2
+ )
+ if do_classifier_free_guidance:
+ noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
+ noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_text - noise_pred_uncond)
- latents = callback_outputs.pop("latents", latents)
- prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
- negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds)
+ # compute the previous noisy sample x_t -> x_t-1
+ if not isinstance(self.scheduler, CogVideoXDPMScheduler):
+ latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
+ else:
+ latents, old_pred_original_sample = self.scheduler.step(
+ noise_pred,
+ old_pred_original_sample,
+ t,
+ timesteps[i - 1] if i > 0 else None,
+ latents,
+ **extra_step_kwargs,
+ return_dict=False,
+ )
+ latents = latents.to(prompt_embeds.dtype)
- if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
- progress_bar.update()
- if comfyui_progressbar:
- pbar.update(1)
+ # call the callback, if provided
+ if callback_on_step_end is not None:
+ callback_kwargs = {}
+ for k in callback_on_step_end_tensor_inputs:
+ callback_kwargs[k] = locals()[k]
+ callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
+
+ latents = callback_outputs.pop("latents", latents)
+ prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
+ negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds)
+
+ if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
+ progress_bar.update()
+ if comfyui_progressbar:
+ pbar.update(1)
# if output_type == "numpy":
# video = self.decode_latents(latents)
diff --git a/cogvideox_fun/transformer_3d.py b/cogvideox_fun/transformer_3d.py
index 88c8013..8a607b4 100644
--- a/cogvideox_fun/transformer_3d.py
+++ b/cogvideox_fun/transformer_3d.py
@@ -26,15 +26,169 @@ 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
from diffusers.models.attention import Attention, FeedForward
-from diffusers.models.attention_processor import AttentionProcessor, CogVideoXAttnProcessor2_0, FusedCogVideoXAttnProcessor2_0
+from diffusers.models.attention_processor import AttentionProcessor#, CogVideoXAttnProcessor2_0, FusedCogVideoXAttnProcessor2_0
from diffusers.models.embeddings import TimestepEmbedding, Timesteps, get_3d_sincos_pos_embed
from diffusers.models.modeling_outputs import Transformer2DModelOutput
from diffusers.models.modeling_utils import ModelMixin
from diffusers.models.normalization import AdaLayerNorm, CogVideoXLayerNormZero
-
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
+try:
+ from sageattention import sageattn
+ SAGEATTN_IS_AVAVILABLE = True
+ logger.info("Using sageattn")
+except:
+ logger.info("sageattn not found, using sdpa")
+ SAGEATTN_IS_AVAVILABLE = False
+
+class CogVideoXAttnProcessor2_0:
+ r"""
+ Processor for implementing scaled dot-product attention for the CogVideoX model. It applies a rotary embedding on
+ query and key vectors, but does not include spatial normalization.
+ """
+
+ def __init__(self):
+ if not hasattr(F, "scaled_dot_product_attention"):
+ raise ImportError("CogVideoXAttnProcessor requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
+
+ def __call__(
+ self,
+ attn: Attention,
+ hidden_states: torch.Tensor,
+ encoder_hidden_states: torch.Tensor,
+ attention_mask: Optional[torch.Tensor] = None,
+ image_rotary_emb: Optional[torch.Tensor] = None,
+ ) -> torch.Tensor:
+ text_seq_length = encoder_hidden_states.size(1)
+
+ hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
+
+ batch_size, sequence_length, _ = (
+ hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
+ )
+
+ if attention_mask is not None:
+ attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
+ attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1])
+
+ query = attn.to_q(hidden_states)
+ key = attn.to_k(hidden_states)
+ value = attn.to_v(hidden_states)
+
+ inner_dim = key.shape[-1]
+ head_dim = inner_dim // attn.heads
+
+ query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
+ key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
+ value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
+
+ if attn.norm_q is not None:
+ query = attn.norm_q(query)
+ if attn.norm_k is not None:
+ key = attn.norm_k(key)
+
+ # Apply RoPE if needed
+ if image_rotary_emb is not None:
+ from diffusers.models.embeddings import apply_rotary_emb
+
+ query[:, :, text_seq_length:] = apply_rotary_emb(query[:, :, text_seq_length:], image_rotary_emb)
+ if not attn.is_cross_attention:
+ key[:, :, text_seq_length:] = apply_rotary_emb(key[:, :, text_seq_length:], image_rotary_emb)
+
+ 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)
+
+ # linear proj
+ hidden_states = attn.to_out[0](hidden_states)
+ # dropout
+ hidden_states = attn.to_out[1](hidden_states)
+
+ encoder_hidden_states, hidden_states = hidden_states.split(
+ [text_seq_length, hidden_states.size(1) - text_seq_length], dim=1
+ )
+ return hidden_states, encoder_hidden_states
+
+
+class FusedCogVideoXAttnProcessor2_0:
+ r"""
+ Processor for implementing scaled dot-product attention for the CogVideoX model. It applies a rotary embedding on
+ query and key vectors, but does not include spatial normalization.
+ """
+
+ def __init__(self):
+ if not hasattr(F, "scaled_dot_product_attention"):
+ raise ImportError("CogVideoXAttnProcessor requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
+
+ def __call__(
+ self,
+ attn: Attention,
+ hidden_states: torch.Tensor,
+ encoder_hidden_states: torch.Tensor,
+ attention_mask: Optional[torch.Tensor] = None,
+ image_rotary_emb: Optional[torch.Tensor] = None,
+ ) -> torch.Tensor:
+ text_seq_length = encoder_hidden_states.size(1)
+
+ hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
+
+ batch_size, sequence_length, _ = (
+ hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
+ )
+
+ if attention_mask is not None:
+ attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
+ attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1])
+
+ qkv = attn.to_qkv(hidden_states)
+ split_size = qkv.shape[-1] // 3
+ query, key, value = torch.split(qkv, split_size, dim=-1)
+
+ inner_dim = key.shape[-1]
+ head_dim = inner_dim // attn.heads
+
+ query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
+ key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
+ value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
+
+ if attn.norm_q is not None:
+ query = attn.norm_q(query)
+ if attn.norm_k is not None:
+ key = attn.norm_k(key)
+
+ # Apply RoPE if needed
+ if image_rotary_emb is not None:
+ from diffusers.models.embeddings import apply_rotary_emb
+
+ query[:, :, text_seq_length:] = apply_rotary_emb(query[:, :, text_seq_length:], image_rotary_emb)
+ if not attn.is_cross_attention:
+ key[:, :, text_seq_length:] = apply_rotary_emb(key[:, :, text_seq_length:], image_rotary_emb)
+
+ 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)
+
+ # linear proj
+ hidden_states = attn.to_out[0](hidden_states)
+ # dropout
+ hidden_states = attn.to_out[1](hidden_states)
+
+ encoder_hidden_states, hidden_states = hidden_states.split(
+ [text_seq_length, hidden_states.size(1) - text_seq_length], dim=1
+ )
+ return hidden_states, encoder_hidden_states
+
class CogVideoXPatchEmbed(nn.Module):
def __init__(
self,
diff --git a/custom_cogvideox_transformer_3d.py b/custom_cogvideox_transformer_3d.py
new file mode 100644
index 0000000..f2c27fd
--- /dev/null
+++ b/custom_cogvideox_transformer_3d.py
@@ -0,0 +1,641 @@
+# Copyright 2024 The CogVideoX team, Tsinghua University & ZhipuAI and The HuggingFace Team.
+# All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+from typing import Any, Dict, Optional, Tuple, Union
+
+import torch
+from torch import nn
+import torch.nn.functional as F
+
+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
+from diffusers.models.attention import Attention, FeedForward
+from diffusers.models.attention_processor import AttentionProcessor
+from diffusers.models.embeddings import CogVideoXPatchEmbed, TimestepEmbedding, Timesteps
+from diffusers.models.modeling_outputs import Transformer2DModelOutput
+from diffusers.models.modeling_utils import ModelMixin
+from diffusers.models.normalization import AdaLayerNorm, CogVideoXLayerNormZero
+
+
+logger = logging.get_logger(__name__) # pylint: disable=invalid-name
+
+try:
+ from sageattention import sageattn
+ SAGEATTN_IS_AVAVILABLE = True
+ logger.info("Using sageattn")
+except:
+ logger.info("sageattn not found, using sdpa")
+ SAGEATTN_IS_AVAVILABLE = False
+
+class CogVideoXAttnProcessor2_0:
+ r"""
+ Processor for implementing scaled dot-product attention for the CogVideoX model. It applies a rotary embedding on
+ query and key vectors, but does not include spatial normalization.
+ """
+
+ def __init__(self):
+ if not hasattr(F, "scaled_dot_product_attention"):
+ raise ImportError("CogVideoXAttnProcessor requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
+
+ def __call__(
+ self,
+ attn: Attention,
+ hidden_states: torch.Tensor,
+ encoder_hidden_states: torch.Tensor,
+ attention_mask: Optional[torch.Tensor] = None,
+ image_rotary_emb: Optional[torch.Tensor] = None,
+ ) -> torch.Tensor:
+ text_seq_length = encoder_hidden_states.size(1)
+
+ hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
+
+ batch_size, sequence_length, _ = (
+ hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
+ )
+
+ if attention_mask is not None:
+ attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
+ attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1])
+
+ query = attn.to_q(hidden_states)
+ key = attn.to_k(hidden_states)
+ value = attn.to_v(hidden_states)
+
+ inner_dim = key.shape[-1]
+ head_dim = inner_dim // attn.heads
+
+ query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
+ key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
+ value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
+
+ if attn.norm_q is not None:
+ query = attn.norm_q(query)
+ if attn.norm_k is not None:
+ key = attn.norm_k(key)
+
+ # Apply RoPE if needed
+ if image_rotary_emb is not None:
+ from diffusers.models.embeddings import apply_rotary_emb
+
+ query[:, :, text_seq_length:] = apply_rotary_emb(query[:, :, text_seq_length:], image_rotary_emb)
+ if not attn.is_cross_attention:
+ key[:, :, text_seq_length:] = apply_rotary_emb(key[:, :, text_seq_length:], image_rotary_emb)
+
+ 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)
+
+ # linear proj
+ hidden_states = attn.to_out[0](hidden_states)
+ # dropout
+ hidden_states = attn.to_out[1](hidden_states)
+
+ encoder_hidden_states, hidden_states = hidden_states.split(
+ [text_seq_length, hidden_states.size(1) - text_seq_length], dim=1
+ )
+ return hidden_states, encoder_hidden_states
+
+
+class FusedCogVideoXAttnProcessor2_0:
+ r"""
+ Processor for implementing scaled dot-product attention for the CogVideoX model. It applies a rotary embedding on
+ query and key vectors, but does not include spatial normalization.
+ """
+
+ def __init__(self):
+ if not hasattr(F, "scaled_dot_product_attention"):
+ raise ImportError("CogVideoXAttnProcessor requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
+
+ def __call__(
+ self,
+ attn: Attention,
+ hidden_states: torch.Tensor,
+ encoder_hidden_states: torch.Tensor,
+ attention_mask: Optional[torch.Tensor] = None,
+ image_rotary_emb: Optional[torch.Tensor] = None,
+ ) -> torch.Tensor:
+ text_seq_length = encoder_hidden_states.size(1)
+
+ hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
+
+ batch_size, sequence_length, _ = (
+ hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
+ )
+
+ if attention_mask is not None:
+ attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
+ attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1])
+
+ qkv = attn.to_qkv(hidden_states)
+ split_size = qkv.shape[-1] // 3
+ query, key, value = torch.split(qkv, split_size, dim=-1)
+
+ inner_dim = key.shape[-1]
+ head_dim = inner_dim // attn.heads
+
+ query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
+ key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
+ value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
+
+ if attn.norm_q is not None:
+ query = attn.norm_q(query)
+ if attn.norm_k is not None:
+ key = attn.norm_k(key)
+
+ # Apply RoPE if needed
+ if image_rotary_emb is not None:
+ from diffusers.models.embeddings import apply_rotary_emb
+
+ query[:, :, text_seq_length:] = apply_rotary_emb(query[:, :, text_seq_length:], image_rotary_emb)
+ if not attn.is_cross_attention:
+ key[:, :, text_seq_length:] = apply_rotary_emb(key[:, :, text_seq_length:], image_rotary_emb)
+
+ 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)
+
+ # linear proj
+ hidden_states = attn.to_out[0](hidden_states)
+ # dropout
+ hidden_states = attn.to_out[1](hidden_states)
+
+ encoder_hidden_states, hidden_states = hidden_states.split(
+ [text_seq_length, hidden_states.size(1) - text_seq_length], dim=1
+ )
+ return hidden_states, encoder_hidden_states
+
+@maybe_allow_in_graph
+class CogVideoXBlock(nn.Module):
+ r"""
+ Transformer block used in [CogVideoX](https://github.com/THUDM/CogVideo) model.
+
+ Parameters:
+ dim (`int`):
+ The number of channels in the input and output.
+ num_attention_heads (`int`):
+ The number of heads to use for multi-head attention.
+ attention_head_dim (`int`):
+ The number of channels in each head.
+ time_embed_dim (`int`):
+ The number of channels in timestep embedding.
+ dropout (`float`, defaults to `0.0`):
+ The dropout probability to use.
+ activation_fn (`str`, defaults to `"gelu-approximate"`):
+ Activation function to be used in feed-forward.
+ attention_bias (`bool`, defaults to `False`):
+ Whether or not to use bias in attention projection layers.
+ qk_norm (`bool`, defaults to `True`):
+ Whether or not to use normalization after query and key projections in Attention.
+ norm_elementwise_affine (`bool`, defaults to `True`):
+ Whether to use learnable elementwise affine parameters for normalization.
+ norm_eps (`float`, defaults to `1e-5`):
+ Epsilon value for normalization layers.
+ final_dropout (`bool` defaults to `False`):
+ Whether to apply a final dropout after the last feed-forward layer.
+ ff_inner_dim (`int`, *optional*, defaults to `None`):
+ Custom hidden dimension of Feed-forward layer. If not provided, `4 * dim` is used.
+ ff_bias (`bool`, defaults to `True`):
+ Whether or not to use bias in Feed-forward layer.
+ attention_out_bias (`bool`, defaults to `True`):
+ Whether or not to use bias in Attention output projection layer.
+ """
+
+ def __init__(
+ self,
+ dim: int,
+ num_attention_heads: int,
+ attention_head_dim: int,
+ time_embed_dim: int,
+ dropout: float = 0.0,
+ activation_fn: str = "gelu-approximate",
+ attention_bias: bool = False,
+ qk_norm: bool = True,
+ norm_elementwise_affine: bool = True,
+ norm_eps: float = 1e-5,
+ final_dropout: bool = True,
+ ff_inner_dim: Optional[int] = None,
+ ff_bias: bool = True,
+ attention_out_bias: bool = True,
+ ):
+ super().__init__()
+
+ # 1. Self Attention
+ self.norm1 = CogVideoXLayerNormZero(time_embed_dim, dim, norm_elementwise_affine, norm_eps, bias=True)
+
+ self.attn1 = Attention(
+ query_dim=dim,
+ dim_head=attention_head_dim,
+ heads=num_attention_heads,
+ qk_norm="layer_norm" if qk_norm else None,
+ eps=1e-6,
+ bias=attention_bias,
+ out_bias=attention_out_bias,
+ processor=CogVideoXAttnProcessor2_0(),
+ )
+
+ # 2. Feed Forward
+ self.norm2 = CogVideoXLayerNormZero(time_embed_dim, dim, norm_elementwise_affine, norm_eps, bias=True)
+
+ self.ff = FeedForward(
+ dim,
+ dropout=dropout,
+ activation_fn=activation_fn,
+ final_dropout=final_dropout,
+ inner_dim=ff_inner_dim,
+ bias=ff_bias,
+ )
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ encoder_hidden_states: torch.Tensor,
+ temb: torch.Tensor,
+ image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
+ ) -> torch.Tensor:
+ text_seq_length = encoder_hidden_states.size(1)
+
+ # norm & modulate
+ norm_hidden_states, norm_encoder_hidden_states, gate_msa, enc_gate_msa = self.norm1(
+ hidden_states, encoder_hidden_states, temb
+ )
+
+ # attention
+ attn_hidden_states, attn_encoder_hidden_states = self.attn1(
+ hidden_states=norm_hidden_states,
+ encoder_hidden_states=norm_encoder_hidden_states,
+ image_rotary_emb=image_rotary_emb,
+ )
+
+ hidden_states = hidden_states + gate_msa * attn_hidden_states
+ encoder_hidden_states = encoder_hidden_states + enc_gate_msa * attn_encoder_hidden_states
+
+ # norm & modulate
+ norm_hidden_states, norm_encoder_hidden_states, gate_ff, enc_gate_ff = self.norm2(
+ hidden_states, encoder_hidden_states, temb
+ )
+
+ # feed-forward
+ norm_hidden_states = torch.cat([norm_encoder_hidden_states, norm_hidden_states], dim=1)
+ ff_output = self.ff(norm_hidden_states)
+
+ hidden_states = hidden_states + gate_ff * ff_output[:, text_seq_length:]
+ encoder_hidden_states = encoder_hidden_states + enc_gate_ff * ff_output[:, :text_seq_length]
+
+ return hidden_states, encoder_hidden_states
+
+
+class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin):
+ """
+ A Transformer model for video-like data in [CogVideoX](https://github.com/THUDM/CogVideo).
+
+ Parameters:
+ num_attention_heads (`int`, defaults to `30`):
+ The number of heads to use for multi-head attention.
+ attention_head_dim (`int`, defaults to `64`):
+ The number of channels in each head.
+ in_channels (`int`, defaults to `16`):
+ The number of channels in the input.
+ out_channels (`int`, *optional*, defaults to `16`):
+ The number of channels in the output.
+ flip_sin_to_cos (`bool`, defaults to `True`):
+ Whether to flip the sin to cos in the time embedding.
+ time_embed_dim (`int`, defaults to `512`):
+ Output dimension of timestep embeddings.
+ text_embed_dim (`int`, defaults to `4096`):
+ Input dimension of text embeddings from the text encoder.
+ num_layers (`int`, defaults to `30`):
+ The number of layers of Transformer blocks to use.
+ dropout (`float`, defaults to `0.0`):
+ The dropout probability to use.
+ attention_bias (`bool`, defaults to `True`):
+ Whether or not to use bias in the attention projection layers.
+ sample_width (`int`, defaults to `90`):
+ The width of the input latents.
+ sample_height (`int`, defaults to `60`):
+ The height of the input latents.
+ sample_frames (`int`, defaults to `49`):
+ The number of frames in the input latents. Note that this parameter was incorrectly initialized to 49
+ instead of 13 because CogVideoX processed 13 latent frames at once in its default and recommended settings,
+ but cannot be changed to the correct value to ensure backwards compatibility. To create a transformer with
+ K latent frames, the correct value to pass here would be: ((K - 1) * temporal_compression_ratio + 1).
+ patch_size (`int`, defaults to `2`):
+ The size of the patches to use in the patch embedding layer.
+ temporal_compression_ratio (`int`, defaults to `4`):
+ The compression ratio across the temporal dimension. See documentation for `sample_frames`.
+ max_text_seq_length (`int`, defaults to `226`):
+ The maximum sequence length of the input text embeddings.
+ activation_fn (`str`, defaults to `"gelu-approximate"`):
+ Activation function to use in feed-forward.
+ timestep_activation_fn (`str`, defaults to `"silu"`):
+ Activation function to use when generating the timestep embeddings.
+ norm_elementwise_affine (`bool`, defaults to `True`):
+ Whether or not to use elementwise affine in normalization layers.
+ norm_eps (`float`, defaults to `1e-5`):
+ The epsilon value to use in normalization layers.
+ spatial_interpolation_scale (`float`, defaults to `1.875`):
+ Scaling factor to apply in 3D positional embeddings across spatial dimensions.
+ temporal_interpolation_scale (`float`, defaults to `1.0`):
+ Scaling factor to apply in 3D positional embeddings across temporal dimensions.
+ """
+
+ _supports_gradient_checkpointing = True
+
+ @register_to_config
+ def __init__(
+ self,
+ num_attention_heads: int = 30,
+ attention_head_dim: int = 64,
+ in_channels: int = 16,
+ out_channels: Optional[int] = 16,
+ flip_sin_to_cos: bool = True,
+ freq_shift: int = 0,
+ time_embed_dim: int = 512,
+ text_embed_dim: int = 4096,
+ num_layers: int = 30,
+ 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."
+ )
+
+ # 1. Patch embedding
+ self.patch_embed = CogVideoXPatchEmbed(
+ patch_size=patch_size,
+ in_channels=in_channels,
+ embed_dim=inner_dim,
+ text_embed_dim=text_embed_dim,
+ bias=True,
+ sample_width=sample_width,
+ sample_height=sample_height,
+ sample_frames=sample_frames,
+ temporal_compression_ratio=temporal_compression_ratio,
+ max_text_seq_length=max_text_seq_length,
+ 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.norm_final = nn.LayerNorm(inner_dim, norm_eps, norm_elementwise_affine)
+
+ # 4. Output blocks
+ self.norm_out = AdaLayerNorm(
+ embedding_dim=time_embed_dim,
+ output_dim=2 * inner_dim,
+ norm_elementwise_affine=norm_elementwise_affine,
+ norm_eps=norm_eps,
+ chunk_dim=1,
+ )
+ self.proj_out = nn.Linear(inner_dim, patch_size * patch_size * out_channels)
+
+ self.gradient_checkpointing = False
+
+ def _set_gradient_checkpointing(self, module, value=False):
+ self.gradient_checkpointing = value
+
+ @property
+ # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.attn_processors
+ def attn_processors(self) -> Dict[str, AttentionProcessor]:
+ r"""
+ Returns:
+ `dict` of attention processors: A dictionary containing all attention processors used in the model with
+ indexed by its weight name.
+ """
+ # set recursively
+ processors = {}
+
+ def fn_recursive_add_processors(name: str, module: torch.nn.Module, processors: Dict[str, AttentionProcessor]):
+ if hasattr(module, "get_processor"):
+ processors[f"{name}.processor"] = module.get_processor()
+
+ for sub_name, child in module.named_children():
+ fn_recursive_add_processors(f"{name}.{sub_name}", child, processors)
+
+ return processors
+
+ for name, module in self.named_children():
+ fn_recursive_add_processors(name, module, processors)
+
+ return processors
+
+ # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_attn_processor
+ def set_attn_processor(self, processor: Union[AttentionProcessor, Dict[str, AttentionProcessor]]):
+ r"""
+ Sets the attention processor to use to compute attention.
+
+ Parameters:
+ processor (`dict` of `AttentionProcessor` or only `AttentionProcessor`):
+ The instantiated processor class or a dictionary of processor classes that will be set as the processor
+ for **all** `Attention` layers.
+
+ If `processor` is a dict, the key needs to define the path to the corresponding cross attention
+ processor. This is strongly recommended when setting trainable attention processors.
+
+ """
+ count = len(self.attn_processors.keys())
+
+ if isinstance(processor, dict) and len(processor) != count:
+ raise ValueError(
+ f"A dict of processors was passed, but the number of processors {len(processor)} does not match the"
+ f" number of attention layers: {count}. Please make sure to pass {count} processor classes."
+ )
+
+ def fn_recursive_attn_processor(name: str, module: torch.nn.Module, processor):
+ if hasattr(module, "set_processor"):
+ if not isinstance(processor, dict):
+ module.set_processor(processor)
+ else:
+ module.set_processor(processor.pop(f"{name}.processor"))
+
+ for sub_name, child in module.named_children():
+ fn_recursive_attn_processor(f"{name}.{sub_name}", child, processor)
+
+ for name, module in self.named_children():
+ fn_recursive_attn_processor(name, module, processor)
+
+ # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.fuse_qkv_projections with FusedAttnProcessor2_0->FusedCogVideoXAttnProcessor2_0
+ def fuse_qkv_projections(self):
+ """
+ Enables fused QKV projections. For self-attention modules, all projection matrices (i.e., query, key, value)
+ are fused. For cross-attention modules, key and value projection matrices are fused.
+
+
+
+ This API is 🧪 experimental.
+
+
+ """
+ self.original_attn_processors = None
+
+ for _, attn_processor in self.attn_processors.items():
+ if "Added" in str(attn_processor.__class__.__name__):
+ raise ValueError("`fuse_qkv_projections()` is not supported for models having added KV projections.")
+
+ self.original_attn_processors = self.attn_processors
+
+ for module in self.modules():
+ if isinstance(module, Attention):
+ module.fuse_projections(fuse=True)
+
+ self.set_attn_processor(FusedCogVideoXAttnProcessor2_0())
+
+ # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.unfuse_qkv_projections
+ def unfuse_qkv_projections(self):
+ """Disables the fused QKV projection if enabled.
+
+
+
+ This API is 🧪 experimental.
+
+
+
+ """
+ if self.original_attn_processors is not None:
+ self.set_attn_processor(self.original_attn_processors)
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ encoder_hidden_states: torch.Tensor,
+ timestep: Union[int, float, torch.LongTensor],
+ timestep_cond: Optional[torch.Tensor] = None,
+ image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
+ return_dict: bool = True,
+ ):
+ batch_size, num_frames, channels, height, width = hidden_states.shape
+
+ # 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)
+
+ # 2. Patch embedding
+ 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:]
+
+ # 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,
+ )
+
+ if not self.config.use_rotary_positional_embeddings:
+ # CogVideoX-2B
+ hidden_states = self.norm_final(hidden_states)
+ else:
+ # CogVideoX-5B
+ hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
+ hidden_states = self.norm_final(hidden_states)
+ hidden_states = hidden_states[:, text_seq_length:]
+
+ # 4. Final block
+ hidden_states = self.norm_out(hidden_states, temb=emb)
+ hidden_states = self.proj_out(hidden_states)
+
+ # 5. Unpatchify
+ # Note: we use `-1` instead of `channels`:
+ # - It is okay to `channels` use for CogVideoX-2b and CogVideoX-5b (number of input channels is equal to output channels)
+ # - However, for CogVideoX-5b-I2V also takes concatenated input image latents (number of input channels is twice the output channels)
+ p = self.config.patch_size
+ output = hidden_states.reshape(batch_size, num_frames, height // p, width // p, -1, p, p)
+ output = output.permute(0, 1, 4, 2, 5, 3, 6).flatten(5, 6).flatten(3, 4)
+
+ if not return_dict:
+ return (output,)
+ return Transformer2DModelOutput(sample=output)
diff --git a/pipeline_cogvideox.py b/pipeline_cogvideox.py
index 57e2172..6b9e909 100644
--- a/pipeline_cogvideox.py
+++ b/pipeline_cogvideox.py
@@ -20,14 +20,15 @@ import torch
import torch.nn.functional as F
import math
-from diffusers.models import AutoencoderKLCogVideoX, CogVideoXTransformer3DModel
-from diffusers.pipelines.pipeline_utils import DiffusionPipeline
+from diffusers.models import AutoencoderKLCogVideoX#, CogVideoXTransformer3DModel
from diffusers.schedulers import CogVideoXDDIMScheduler, CogVideoXDPMScheduler
from diffusers.utils import logging
from diffusers.utils.torch_utils import randn_tensor
from diffusers.video_processor import VideoProcessor
from diffusers.models.embeddings import get_3d_rotary_pos_embed
+from .custom_cogvideox_transformer_3d import CogVideoXTransformer3DModel
+
from comfy.utils import ProgressBar
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
@@ -164,7 +165,8 @@ class CogVideoXPipeline(VideoSysPipeline):
set_pab_manager(pab_config)
def prepare_latents(
- self, batch_size, num_channels_latents, num_frames, height, width, dtype, device, generator, timesteps, denoise_strength, num_inference_steps, latents=None,
+ self, batch_size, num_channels_latents, num_frames, height, width, dtype, device, generator, timesteps, denoise_strength,
+ num_inference_steps, latents=None, freenoise=True, context_size=None, context_overlap=None
):
shape = (
batch_size,
@@ -178,9 +180,43 @@ class CogVideoXPipeline(VideoSysPipeline):
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
)
- noise = randn_tensor(shape, generator=generator, device=device, dtype=self.vae.dtype)
+ noise = randn_tensor(shape, generator=generator, device=torch.device("cpu"), dtype=self.vae.dtype)
+ if freenoise:
+ print("Applying FreeNoise")
+ # code and comments from AnimateDiff-Evolved by Kosinkadink (https://github.com/Kosinkadink/ComfyUI-AnimateDiff-Evolved)
+ video_length = num_frames // 4
+ delta = context_size - context_overlap
+ for start_idx in range(0, video_length-context_size, delta):
+ # start_idx corresponds to the beginning of a context window
+ # goal: place shuffled in the delta region right after the end of the context window
+ # if space after context window is not enough to place the noise, adjust and finish
+ place_idx = start_idx + context_size
+ # if place_idx is outside the valid indexes, we are already finished
+ if place_idx >= video_length:
+ break
+ end_idx = place_idx - 1
+ #print("video_length:", video_length, "start_idx:", start_idx, "end_idx:", end_idx, "place_idx:", place_idx, "delta:", delta)
+
+ # if there is not enough room to copy delta amount of indexes, copy limited amount and finish
+ if end_idx + delta >= video_length:
+ final_delta = video_length - place_idx
+ # generate list of indexes in final delta region
+ list_idx = torch.tensor(list(range(start_idx,start_idx+final_delta)), device=torch.device("cpu"), dtype=torch.long)
+ # shuffle list
+ list_idx = list_idx[torch.randperm(final_delta, generator=generator)]
+ # apply shuffled indexes
+ noise[:, place_idx:place_idx + final_delta, :, :, :] = noise[:, list_idx, :, :, :]
+ break
+ # otherwise, do normal behavior
+ # generate list of indexes in delta region
+ list_idx = torch.tensor(list(range(start_idx,start_idx+delta)), device=torch.device("cpu"), dtype=torch.long)
+ # shuffle list
+ list_idx = list_idx[torch.randperm(delta, generator=generator)]
+ # apply shuffled indexes
+ #print("place_idx:", place_idx, "delta:", delta, "list_idx:", list_idx)
+ noise[:, place_idx:place_idx + delta, :, :, :] = noise[:, list_idx, :, :, :]
if latents is None:
- latents = noise
+ latents = noise.to(device)
else:
latents = latents.to(device)
timesteps, num_inference_steps = self.get_timesteps(num_inference_steps, denoise_strength, device)
@@ -346,6 +382,11 @@ class CogVideoXPipeline(VideoSysPipeline):
negative_prompt_embeds: Optional[torch.Tensor] = None,
device = torch.device("cuda"),
scheduler_name: str = "DPM",
+ context_schedule: Optional[str] = None,
+ context_frames: Optional[int] = None,
+ context_stride: Optional[int] = None,
+ context_overlap: Optional[int] = None,
+ freenoise: Optional[bool] = True,
):
"""
Function invoked when calling the pipeline for generation.
@@ -448,7 +489,10 @@ class CogVideoXPipeline(VideoSysPipeline):
timesteps,
denoise_strength,
num_inference_steps,
- latents
+ latents,
+ context_size=context_frames,
+ context_overlap=context_overlap,
+ freenoise=freenoise,
)
latents = latents.to(self.vae.dtype)
#print("latents", latents.shape)
@@ -492,22 +536,39 @@ class CogVideoXPipeline(VideoSysPipeline):
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
comfy_pbar = ProgressBar(num_inference_steps)
- # 8. Temporal tiling prep
- if "tiled" in scheduler_name:
+ # 8.5. Temporal tiling prep
+ if context_schedule is not None and context_schedule == "temporal_tiling":
+ t_tile_length = context_frames
+ t_tile_overlap = context_overlap
t_tile_weights = self._gaussian_weights(t_tile_length=t_tile_length, t_batch_size=1).to(latents.device).to(self.vae.dtype)
- temporal_tiling = True
+ use_temporal_tiling = True
print("Temporal tiling enabled")
+ elif context_schedule is not None:
+ if image_cond_latents is not None:
+ raise NotImplementedError("Context schedule not currently supported with image conditioning")
+ print(f"Context schedule enabled: {context_frames} frames, {context_stride} stride, {context_overlap} overlap")
+ use_temporal_tiling = False
+ use_context_schedule = True
+ from .cogvideox_fun.context import get_context_scheduler
+ context = get_context_scheduler(context_schedule)
+
else:
- temporal_tiling = False
- print("Temporal tiling disabled")
- #print("latents.shape", latents.shape)
+ use_temporal_tiling = False
+ use_context_schedule = False
+ print("Temporal tiling and context schedule disabled")
+ # 7. Create rotary embeds if required
+ image_rotary_emb = (
+ self._prepare_rotary_positional_embeddings(height, width, latents.size(1), device)
+ if self.transformer.config.use_rotary_positional_embeddings
+ else None
+ )
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):
if self.interrupt:
continue
- if temporal_tiling and isinstance(self.scheduler, CogVideoXDDIMScheduler):
+ if use_temporal_tiling and isinstance(self.scheduler, CogVideoXDDIMScheduler):
#temporal tiling code based on https://github.com/mayuelala/FollowYourEmoji/blob/main/models/video_pipeline.py
# =====================================================
grid_ts = 0
@@ -533,7 +594,7 @@ class CogVideoXPipeline(VideoSysPipeline):
#latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
image_rotary_emb = (
- self._prepare_rotary_positional_embeddings(height, width, latents.size(1), device, input_start_t, input_end_t)
+ self._prepare_rotary_positional_embeddings(height, width, t_tile_length, device)
if self.transformer.config.use_rotary_positional_embeddings
else None
)
@@ -600,6 +661,79 @@ class CogVideoXPipeline(VideoSysPipeline):
progress_bar.update()
comfy_pbar.update(1)
# ==========================================
+ elif use_context_schedule:
+ 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)
+ 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
+ latent_model_input = torch.cat([latent_model_input, latent_image_input], dim=2)
+
+ # 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,
+ ))
+
+ 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]
+
+ # 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,
+ )[0]
+
+ counter[:, c, :, :, :] += 1
+ noise_pred = noise_pred.float()
+
+ noise_pred /= counter
+ if do_classifier_free_guidance:
+ noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
+ noise_pred = noise_pred_uncond + self._guidance_scale * (noise_pred_text - noise_pred_uncond)
+
+ # compute the previous noisy sample x_t -> x_t-1
+ if not isinstance(self.scheduler, CogVideoXDPMScheduler):
+ latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
+ else:
+ latents, old_pred_original_sample = self.scheduler.step(
+ noise_pred,
+ old_pred_original_sample,
+ t,
+ timesteps[i - 1] if i > 0 else None,
+ latents,
+ **extra_step_kwargs,
+ return_dict=False,
+ )
+ latents = latents.to(prompt_embeds.dtype)
+
+ if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
+ progress_bar.update()
+ comfy_pbar.update(1)
+
else:
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)
@@ -610,7 +744,6 @@ 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])
-
# predict noise model_output
noise_pred = self.transformer(
hidden_states=latent_model_input,