Fixed ControlLLLite error with certain latent sizes

This commit is contained in:
Jedrzej Kosinski
2024-02-06 01:19:07 -06:00
parent b2cd17ffde
commit 77e92f5f8f
3 changed files with 121 additions and 18 deletions
+48 -11
View File
@@ -284,14 +284,18 @@ class SparseCtrlAdvanced(ControlNetAdvanced):
class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase):
# This ControlNet is more of an attention patch than a traditional controlnet
def __init__(self, patch: LLLitePatch, timestep_keyframes: TimestepKeyframeGroup, device=None):
def __init__(self, patch_attn1: LLLitePatch, patch_attn2: LLLitePatch, timestep_keyframes: TimestepKeyframeGroup, device=None):
super().__init__(device)
AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controllllite(), require_model=True)
self.patch = patch.clone_with_control(self)
self.patch_attn1 = patch_attn1.clone_with_control(self)
self.patch_attn2 = patch_attn2.clone_with_control(self)
self.latent_dims_div2 = None
self.latent_dims_div4 = None
def patch_model(self, model: ModelPatcher):
model.set_model_attn1_patch(self.patch)
model.set_model_attn2_patch(self.patch)
model.set_model_attn1_patch(self.patch_attn1)
model.set_model_attn2_patch(self.patch_attn2)
def set_cond_hint(self, *args, **kwargs):
to_return = super().set_cond_hint(*args, **kwargs)
@@ -301,7 +305,9 @@ class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase):
def pre_run_advanced(self, *args, **kwargs):
AdvancedControlBase.pre_run_advanced(self, *args, **kwargs)
self.patch.set_control(self)
#logger.error(f"in cn: {id(self.patch_attn1)},{id(self.patch_attn2)}")
self.patch_attn1.set_control(self)
self.patch_attn2.set_control(self)
#logger.warn(f"in pre_run_advanced: {id(self)}")
def get_control_advanced(self, x_noisy: Tensor, t, cond, batched_number: int):
@@ -327,6 +333,31 @@ class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase):
self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original, x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(self.device)
if x_noisy.shape[0] != self.cond_hint.shape[0]:
self.cond_hint = broadcast_image_to(self.cond_hint, x_noisy.shape[0], batched_number)
# some special logic here compared to other controlnets:
# * The cond_emb in attn patches will divide latent dims by 2 or 4, integer
# * Due to this loss, the cond_emb will become smaller than x input if latent dims are not divisble by 2 or 4
divisible_by_2_h = x_noisy.shape[2]%2==0
divisible_by_2_w = x_noisy.shape[3]%2==0
if not (divisible_by_2_h and divisible_by_2_w):
#logger.warn(f"{x_noisy.shape} not divisible by 2!")
new_h = (x_noisy.shape[2]//2)*2
new_w = (x_noisy.shape[3]//2)*2
if not divisible_by_2_h:
new_h += 2
if not divisible_by_2_w:
new_w += 2
self.latent_dims_div2 = (new_h, new_w)
divisible_by_4_h = x_noisy.shape[2]%4==0
divisible_by_4_w = x_noisy.shape[3]%4==0
if not (divisible_by_4_h and divisible_by_4_w):
#logger.warn(f"{x_noisy.shape} not divisible by 4!")
new_h = (x_noisy.shape[2]//4)*4
new_w = (x_noisy.shape[3]//4)*4
if not divisible_by_4_h:
new_h += 4
if not divisible_by_4_w:
new_w += 4
self.latent_dims_div4 = (new_h, new_w)
# prepare mask
self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number)
# done preparing; model patches will take care of everything now.
@@ -335,21 +366,26 @@ class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase):
def cleanup_advanced(self):
super().cleanup_advanced()
self.patch.cleanup()
self.patch_attn1.cleanup()
self.patch_attn2.cleanup()
self.latent_dims_div2 = None
self.latent_dims_div4 = None
def copy(self):
c = ControlLLLiteAdvanced(self.patch, self.timestep_keyframes)
c = ControlLLLiteAdvanced(self.patch_attn1, self.patch_attn2, self.timestep_keyframes)
self.copy_to(c)
self.copy_to_advanced(c)
return c
# deepcopy needs to properly keep track of objects to work between model.clone calls!
def __deepcopy__(self, *args, **kwargs):
return self
# def __deepcopy__(self, *args, **kwargs):
# self.cleanup_advanced()
# return self
# def get_models(self):
# # get_models is called once at the start of every KSampler run - use to reset already_patched status
# out = super().get_models()
# logger.error(f"in get_models! {id(self)}")
# return out
@@ -602,6 +638,7 @@ def load_controllllite(ckpt_path: str, controlnet_data: dict[str, Tensor]=None,
#logger.info(f"loaded {ckpt_path} successfully, {len(modules)} modules")
patch = LLLitePatch(modules=modules)
control = ControlLLLiteAdvanced(patch=patch, timestep_keyframes=timestep_keyframe)
patch_attn1 = LLLitePatch(modules=modules, patch_type=LLLitePatch.ATTN1)
patch_attn2 = LLLitePatch(modules=modules, patch_type=LLLitePatch.ATTN2)
control = ControlLLLiteAdvanced(patch_attn1=patch_attn1, patch_attn2=patch_attn2, timestep_keyframes=timestep_keyframe)
return control
+34 -7
View File
@@ -11,7 +11,7 @@ import comfy.utils
from comfy.controlnet import ControlBase
from .logger import logger
from .utils import AdvancedControlBase, prepare_mask_batch
from .utils import AdvancedControlBase, deepcopy_with_sharing, prepare_mask_batch
def extra_options_to_module_prefix(extra_options):
@@ -38,12 +38,16 @@ def extra_options_to_module_prefix(extra_options):
class LLLitePatch:
def __init__(self, modules: dict[str, 'LLLiteModule'], control: Union[AdvancedControlBase, ControlBase]=None):
ATTN1 = "attn1"
ATTN2 = "attn2"
def __init__(self, modules: dict[str, 'LLLiteModule'], patch_type: str, control: Union[AdvancedControlBase, ControlBase]=None):
self.modules = modules
self.control = control
self.patch_type = patch_type
#logger.error(f"create LLLitePatch: {id(self)},{control}")
def __call__(self, q, k, v, extra_options):
#logger.error(f"in __call__: {id(self)}")
# determine if have anything to run
if self.control.timestep_range is not None:
# it turns out comparing single-value tensors to floats is extremely slow
@@ -80,19 +84,34 @@ class LLLitePatch:
def set_control(self, control: Union[AdvancedControlBase, ControlBase]):
self.control = control
#logger.error(f"set control for LLLitePatch: {id(self)},{id(control)}")
#logger.error(f"set control for LLLitePatch: {id(self)}, cn: {id(control)}")
def clone_with_control(self, control: AdvancedControlBase):
#logger.error(f"clone-set control for LLLitePatch: {id(self)},{id(control)}")
return LLLitePatch(self.modules, control)
return LLLitePatch(self.modules, self.patch_type, control)
def cleanup(self):
#del self.control
#self.control = None
#total_cleaned = 0
for module in self.modules.values():
module.cleanup()
# total_cleaned += 1
#logger.info(f"cleaned modules: {total_cleaned}, {id(self)}")
#logger.error(f"cleanup LLLitePatch: {id(self)}")
# make sure deepcopy does not copy control, and deepcopied LLLitePatch should be assigned to control
def __deepcopy__(self, memo):
self.cleanup()
to_return: LLLitePatch = deepcopy_with_sharing(self, shared_attribute_names = ['control'], memo=memo)
#logger.warn(f"patch {id(self)} turned into {id(to_return)}")
try:
if self.patch_type == self.ATTN1:
to_return.control.patch_attn1 = to_return
elif self.patch_type == self.ATTN2:
to_return.control.patch_attn2 = to_return
except Exception:
pass
return to_return
# TODO: use comfy.ops to support fp8 properly
class LLLiteModule(torch.nn.Module):
@@ -159,6 +178,7 @@ class LLLiteModule(torch.nn.Module):
self.prev_sub_idxs = None
def cleanup(self):
del self.cond_emb
self.cond_emb = None
self.cx_shape = None
self.prev_batch = 0
@@ -167,9 +187,15 @@ class LLLiteModule(torch.nn.Module):
def forward(self, x: Tensor, control: Union[AdvancedControlBase, ControlBase]):
mask = None
mask_tk = None
#logger.info(x.shape)
if self.cond_emb is None or control.sub_idxs != self.prev_sub_idxs or x.shape[0] != self.prev_batch:
# print(f"cond_emb is None, {self.name}")
cx = self.conditioning1(control.cond_hint.to(x.device, dtype=x.dtype))
cond_hint = control.cond_hint.to(x.device, dtype=x.dtype)
if control.latent_dims_div2 is not None and x.shape[-1] != 1280:
cond_hint = comfy.utils.common_upscale(cond_hint, control.latent_dims_div2[0] * 8, control.latent_dims_div2[1] * 8, 'nearest-exact', "center").to(x.device, dtype=x.dtype)
elif control.latent_dims_div4 is not None and x.shape[-1] == 1280:
cond_hint = comfy.utils.common_upscale(cond_hint, control.latent_dims_div4[0] * 8, control.latent_dims_div4[1] * 8, 'nearest-exact', "center").to(x.device, dtype=x.dtype)
cx = self.conditioning1(cond_hint)
self.cx_shape = cx.shape
if not self.is_conv2d:
# reshape / b,c,h,w -> b,h*w,c
@@ -211,6 +237,7 @@ class LLLiteModule(torch.nn.Module):
elif mask_tk is not None:
mask = mask * mask_tk
#logger.info(f"cs: {cx.shape}, x: {x.shape}, is_conv2d: {self.is_conv2d}")
cx = torch.cat([cx, self.down(x)], dim=1 if self.is_conv2d else 2)
cx = self.mid(cx)
cx = self.up(cx)
+39
View File
@@ -1,3 +1,4 @@
from copy import deepcopy
from typing import Callable, Union
import torch
from torch import Tensor
@@ -259,6 +260,44 @@ def linear_conversion(x, x_min=0.0, x_max=1.0, new_min=0.0, new_max=1.0):
return (((x - x_min)/(x_max - x_min)) * (new_max - new_min)) + new_min
# from https://stackoverflow.com/a/24621200
def deepcopy_with_sharing(obj, shared_attribute_names, memo=None):
'''
Deepcopy an object, except for a given list of attributes, which should
be shared between the original object and its copy.
obj is some object
shared_attribute_names: A list of strings identifying the attributes that
should be shared between the original and its copy.
memo is the dictionary passed into __deepcopy__. Ignore this argument if
not calling from within __deepcopy__.
'''
assert isinstance(shared_attribute_names, (list, tuple))
shared_attributes = {k: getattr(obj, k) for k in shared_attribute_names}
if hasattr(obj, '__deepcopy__'):
# Do hack to prevent infinite recursion in call to deepcopy
deepcopy_method = obj.__deepcopy__
obj.__deepcopy__ = None
for attr in shared_attribute_names:
del obj.__dict__[attr]
clone = deepcopy(obj)
for attr, val in shared_attributes.items():
setattr(obj, attr, val)
setattr(clone, attr, val)
if hasattr(obj, '__deepcopy__'):
# Undo hack
obj.__deepcopy__ = deepcopy_method
del clone.__deepcopy__
return clone
class WeightTypeException(TypeError):
"Raised when weight not compatible with AdvancedControlBase object"
pass