Fixed SparseCtrl not working with lowvram

This commit is contained in:
Jedrzej Kosinski
2024-04-04 14:35:28 -05:00
parent 0e02a2669a
commit 0556fd2265
3 changed files with 71 additions and 25 deletions
+6 -5
View File
@@ -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()
+52 -12
View File
@@ -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
+13 -8
View File
@@ -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