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