Merge PR #65 - ControlLLLite size edge case fix
Fixed ControlLLLite edge cases for certain latent sizes
This commit is contained in:
+48
-11
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -4,7 +4,7 @@ import torch
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image, ImageOps
|
||||
from .utils import ControlWeights, LatentKeyframeGroup, TimestepKeyframeGroup, TimestepKeyframe
|
||||
from .utils import ControlWeights, LatentKeyframeGroup, TimestepKeyframeGroup, TimestepKeyframe, BIGMAX
|
||||
from .logger import logger
|
||||
|
||||
|
||||
@@ -16,8 +16,8 @@ class LoadImagesFromDirectory:
|
||||
"directory": ("STRING", {"default": ""}),
|
||||
},
|
||||
"optional": {
|
||||
"image_load_cap": ("INT", {"default": 0, "min": 0, "step": 1}),
|
||||
"start_index": ("INT", {"default": 0, "min": 0, "step": 1}),
|
||||
"image_load_cap": ("INT", {"default": 0, "min": 0, "max": BIGMAX, "step": 1}),
|
||||
"start_index": ("INT", {"default": 0, "min": 0, "max": BIGMAX, "step": 1}),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@ from typing import Union
|
||||
import numpy as np
|
||||
from collections.abc import Iterable
|
||||
|
||||
from .utils import LatentKeyframe, LatentKeyframeGroup
|
||||
from .utils import LatentKeyframe, LatentKeyframeGroup, BIGMIN, BIGMAX
|
||||
from .utils import StrengthInterpolation as SI
|
||||
from .logger import logger
|
||||
|
||||
@@ -12,7 +12,7 @@ class LatentKeyframeNode:
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"batch_index": ("INT", {"default": 0, "min": -1000, "max": 1000, "step": 1}),
|
||||
"batch_index": ("INT", {"default": 0, "min": BIGMIN, "max": BIGMAX, "step": 1}),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
},
|
||||
"optional": {
|
||||
@@ -163,8 +163,8 @@ class LatentKeyframeInterpolationNode:
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"batch_index_from": ("INT", {"default": 0, "min": -10000, "max": 10000, "step": 1}),
|
||||
"batch_index_to_excl": ("INT", {"default": 0, "min": -10000, "max": 10000, "step": 1}),
|
||||
"batch_index_from": ("INT", {"default": 0, "min": BIGMIN, "max": BIGMAX, "step": 1}),
|
||||
"batch_index_to_excl": ("INT", {"default": 0, "min": BIGMIN, "max": BIGMAX, "step": 1}),
|
||||
"strength_from": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"strength_to": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"interpolation": ([SI.LINEAR, SI.EASE_IN, SI.EASE_OUT, SI.EASE_IN_OUT], ),
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
from copy import deepcopy
|
||||
from typing import Callable, Union
|
||||
import torch
|
||||
from torch import Tensor
|
||||
@@ -10,6 +11,9 @@ from comfy.model_patcher import ModelPatcher
|
||||
|
||||
from .logger import logger
|
||||
|
||||
BIGMIN = -(2**63-1)
|
||||
BIGMAX = (2**63-1)
|
||||
|
||||
def load_torch_file_with_dict_factory(controlnet_data: dict[str, Tensor], orig_load_torch_file: Callable):
|
||||
def load_torch_file_with_dict(*args, **kwargs):
|
||||
# immediately restore load_torch_file to original version
|
||||
@@ -259,6 +263,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
|
||||
|
||||
Reference in New Issue
Block a user