diff --git a/adv_control/control.py b/adv_control/control.py index 0226793..dfca88f 100644 --- a/adv_control/control.py +++ b/adv_control/control.py @@ -10,7 +10,7 @@ import comfy.controlnet as comfy_cn from comfy.controlnet import ControlBase, ControlNet, ControlLora, T2IAdapter, broadcast_image_to from comfy.model_patcher import ModelPatcher -from .control_sparsectrl import SparseControlNet, SparseCtrlMotionWrapper, SparseMethod, SparseSettings, SparseSpreadMethod, PreprocSparseRGBWrapper +from .control_sparsectrl import SparseModelPatcher, SparseControlNet, SparseCtrlMotionWrapper, SparseMethod, SparseSettings, SparseSpreadMethod, PreprocSparseRGBWrapper from .control_lllite import LLLiteModule, LLLitePatch from .control_svd import svd_unet_config_from_diffusers_unet, SVDControlNet, svd_unet_to_diffusers from .utils import (AdvancedControlBase, TimestepKeyframeGroup, LatentKeyframeGroup, ControlWeightType, ControlWeights, WeightTypeException, @@ -238,6 +238,7 @@ class SVDControlNetAdvanced(ControlNetAdvanced): class SparseCtrlAdvanced(ControlNetAdvanced): def __init__(self, control_model, timestep_keyframes: TimestepKeyframeGroup, sparse_settings: SparseSettings=None, global_average_pooling=False, device=None, load_device=None, manual_cast_dtype=None): super().__init__(control_model=control_model, timestep_keyframes=timestep_keyframes, global_average_pooling=global_average_pooling, device=device, load_device=load_device, manual_cast_dtype=manual_cast_dtype) + self.control_model_wrapped = SparseModelPatcher(self.control_model, load_device=load_device, offload_device=comfy.model_management.unet_offload_device()) self.add_compatible_weight(ControlWeightType.SPARSECTRL) self.control_model: SparseControlNet = self.control_model # does nothing except help with IDE hints self.sparse_settings = sparse_settings if sparse_settings is not None else SparseSettings.default() @@ -331,10 +332,10 @@ class SparseCtrlAdvanced(ControlNetAdvanced): raise ValueError("Any model besides RGB SparseCtrl should NOT have its images go through the RGB SparseCtrl preprocessor.") self.cond_hint_original = self.cond_hint_original.condhint self.latent_format = model.latent_format # LatentFormat object, used to process_in latent cond hint - if self.control_model.motion_holder is not None: - self.control_model.motion_holder.motion_wrapper.reset() - self.control_model.motion_holder.motion_wrapper.set_strength(self.sparse_settings.motion_strength) - self.control_model.motion_holder.motion_wrapper.set_scale_multiplier(self.sparse_settings.motion_scale) + if self.control_model.motion_wrapper is not None: + self.control_model.motion_wrapper.reset() + self.control_model.motion_wrapper.set_strength(self.sparse_settings.motion_strength) + self.control_model.motion_wrapper.set_scale_multiplier(self.sparse_settings.motion_scale) def cleanup_advanced(self): super().cleanup_advanced() diff --git a/adv_control/control_sparsectrl.py b/adv_control/control_sparsectrl.py index 5885ed5..2c91f68 100644 --- a/adv_control/control_sparsectrl.py +++ b/adv_control/control_sparsectrl.py @@ -23,6 +23,7 @@ from comfy.ldm.modules.attention import SpatialTransformer from comfy.ldm.modules.attention import attention_basic, attention_pytorch, attention_split, attention_sub_quad, default from comfy.ldm.modules.attention import FeedForward, SpatialTransformer from comfy.ldm.modules.diffusionmodules.openaimodel import TimestepEmbedSequential, ResBlock, Downsample +from comfy.model_patcher import ModelPatcher from comfy.controlnet import broadcast_image_to from comfy.utils import repeat_to_batch_size import comfy.ops @@ -57,11 +58,11 @@ class SparseControlNet(ControlNetCLDM): self.input_hint_block = TimestepEmbedSequential( zero_module(operations.conv_nd(self.dims, hint_channels, self.model_channels, 3, padding=1, dtype=self.dtype, device=device)), ) - self.motion_holder: MotionWrapperHolder = None + self.motion_wrapper: SparseCtrlMotionWrapper = None def set_actual_length(self, actual_length: int, full_length: int): - if self.motion_holder is not None: - self.motion_holder.motion_wrapper.set_video_length(video_length=actual_length, full_length=full_length) + if self.motion_wrapper is not None: + self.motion_wrapper.set_video_length(video_length=actual_length, full_length=full_length) def forward(self, x: Tensor, hint: Tensor, timesteps, context, y=None, **kwargs): t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False).to(x.dtype) @@ -94,6 +95,50 @@ class SparseControlNet(ControlNetCLDM): return outs +class SparseModelPatcher(ModelPatcher): + def __init__(self, *args, **kwargs): + self.model: SparseControlNet + super().__init__(*args, **kwargs) + + def patch_model(self, device_to=None, patch_weights=True): + if patch_weights: + patched_model = super().patch_model(device_to) + else: + patched_model = super().patch_model(device_to, patch_weights) + try: + self.model.motion_wrapper.to(device=device_to) + except Exception: + raise + return patched_model + + def unpatch_model(self, device_to=None, unpatch_weights=True): + try: + self.model.motion_wrapper.to(device=device_to) + except Exception: + pass + if unpatch_weights: + return super().unpatch_model(device_to) + else: + return super().unpatch_model(device_to, unpatch_weights) + + def clone(self): + # normal ModelPatcher clone actions + n = SparseModelPatcher(self.model, self.load_device, self.offload_device, self.size, self.current_device, weight_inplace_update=self.weight_inplace_update) + n.patches = {} + for k in self.patches: + n.patches[k] = self.patches[k][:] + if hasattr(n, "patches_uuid"): + self.patches_uuid = n.patches_uuid + + n.object_patches = self.object_patches.copy() + n.model_options = copy.deepcopy(self.model_options) + n.model_keys = self.model_keys + if hasattr(n, "backup"): + self.backup = n.backup + if hasattr(n, "object_patches_backup"): + self.object_patches_backup = n.object_patches_backup + + class PreprocSparseRGBWrapper: error_msg = "Invalid use of RGB SparseCtrl output. The output of RGB SparseCtrl preprocessor is NOT a usual image, but a latent pretending to be an image - you must connect the output directly to an Apply ControlNet node (advanced or otherwise). It cannot be used for anything else that accepts IMAGE input." def __init__(self, condhint: Tensor): @@ -270,11 +315,6 @@ def get_position_encoding_max_len(mm_state_dict: dict[str, Tensor], mm_name: str raise ValueError(f"No pos_encoder.pe found in SparseCtrl state_dict - {mm_name} is not a valid SparseCtrl model!") -class MotionWrapperHolder: - def __init__(self, motion_wrapper: 'SparseCtrlMotionWrapper'): - self.motion_wrapper = motion_wrapper - - class SparseCtrlMotionWrapper(nn.Module): def __init__(self, mm_state_dict: dict[str, Tensor]): super().__init__() @@ -293,14 +333,14 @@ class SparseCtrlMotionWrapper(nn.Module): self.up_blocks.append(MotionModule(c, temporal_position_encoding_max_len=self.encoding_max_len, block_type=BlockType.UP)) if has_mid_block(mm_state_dict): self.mid_block = MotionModule(1280, temporal_position_encoding_max_len=self.encoding_max_len, block_type=BlockType.MID) - + def inject(self, unet: SparseControlNet): # inject input (down) blocks self._inject(unet.input_blocks, self.down_blocks) # inject mid block, if present if self.mid_block is not None: self._inject([unet.middle_block], [self.mid_block]) - unet.motion_holder = MotionWrapperHolder(self) + unet.motion_wrapper = self def _inject(self, unet_blocks: nn.ModuleList, mm_blocks: nn.ModuleList): # Rules for injection: @@ -342,8 +382,8 @@ class SparseCtrlMotionWrapper(nn.Module): self._eject(unet.input_blocks) # remove from middle block (encapsulate in list to make compatible) self._eject([unet.middle_block]) - del unet.motion_holder - unet.motion_holder = None + del unet.motion_wrapper + unet.motion_wrapper = None def _eject(self, unet_blocks: nn.ModuleList): # eject all VanillaTemporalModule objects from all blocks diff --git a/adv_control/utils.py b/adv_control/utils.py index f825048..50587cf 100644 --- a/adv_control/utils.py +++ b/adv_control/utils.py @@ -2,7 +2,7 @@ from copy import deepcopy from typing import Callable, Union import torch from torch import Tensor -import torch.nn.functional as F +import torch.nn.functional import math import comfy.ops @@ -271,14 +271,19 @@ class AbstractPreprocWrapper: # depending on model, AnimateDiff may inject into GroupNorm, so make sure GroupNorm will be clean class disable_weight_init_clean_groupnorm(comfy.ops.disable_weight_init): class GroupNorm(comfy.ops.disable_weight_init.GroupNorm): - def forward(self, input: Tensor) -> Tensor: - return F.group_norm( - input, self.num_groups, self.weight, self.bias, self.eps) + def forward_comfy_cast_weights(self, input): + weight, bias = comfy.ops.cast_bias_weight(self, input) + return torch.nn.functional.group_norm(input, self.num_groups, weight, bias, self.eps) + + def forward(self, input): + if self.comfy_cast_weights: + return self.forward_comfy_cast_weights(input) + else: + return torch.nn.functional.group_norm(input, self.num_groups, self.weight, self.bias, self.eps) + class manual_cast_clean_groupnorm(comfy.ops.manual_cast): - class GroupNorm(comfy.ops.manual_cast.GroupNorm): - def forward(self, input: Tensor) -> Tensor: - return F.group_norm( - input, self.num_groups, self.weight, self.bias, self.eps) + class GroupNorm(disable_weight_init_clean_groupnorm.GroupNorm): + comfy_cast_weights = True # adapted from comfy/sample.py