Added full ControlLLLite support, refactored some code to eliminate code repetition and support controls that require model patching

This commit is contained in:
Jedrzej Kosinski
2024-01-31 23:19:16 -06:00
parent 9fcc3b664d
commit fbf005f31d
4 changed files with 274 additions and 194 deletions
+117 -58
View File
@@ -8,8 +8,10 @@ import comfy.model_management
import comfy.model_detection
import comfy.controlnet as comfy_cn
from comfy.controlnet import ControlBase, ControlNet, ControlLora, T2IAdapter, broadcast_image_to
from model_patcher import ModelPatcher
from .control_sparsectrl import SparseControlNet, SparseCtrlMotionWrapper, SparseMethod, SparseSettings, SparseSpreadMethod, PreprocSparseRGBWrapper
from .control_lllite import LLLiteModule, LLLitePatch
from .utils import (AdvancedControlBase, TimestepKeyframeGroup, LatentKeyframeGroup, ControlWeightType, ControlWeights, WeightTypeException,
manual_cast_clean_groupnorm, disable_weight_init_clean_groupnorm, prepare_mask_batch, get_properly_arranged_t2i_weights, load_torch_file_with_dict_factory)
from .logger import logger
@@ -106,8 +108,6 @@ class T2IAdapterAdvanced(T2IAdapter, AdvancedControlBase):
return indeces[idx]
def get_control_advanced(self, x_noisy, t, cond, batched_number):
# prepare timestep and everything related
self.prepare_current_timestep(t=t, batched_number=batched_number)
try:
# if sub indexes present, replace original hint with subsection
if self.sub_idxs is not None:
@@ -168,59 +168,6 @@ class ControlLoraAdvanced(ControlLora, AdvancedControlBase):
global_average_pooling=v.global_average_pooling, device=v.device)
class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase):
# This ControlNet is more of an attention patch than a traditional controlnet
# So, the pre_run will be responsible for a lot of the functionality,
# while the usual get_control is mostly used to set some values
def __init__(self, timestep_keyframes: TimestepKeyframeGroup, device=None):
super().__init__(device)
AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controllllite())
self.already_patched = False
def set_cond_hint(self, *args, **kwargs):
super().set_cond_hint(*args, **kwargs)
# cond hint for LLLite needs to be scaled between (-1, 1) instead of (0, 1)
self.cond_hint_original = self.cond_hint_original * 2.0 - 1.0
def pre_run_advanced(self, model, percent_to_timestep_function):
AdvancedControlBase.pre_run_advanced(self, model, percent_to_timestep_function)
logger.info(f"In ControlLLLiteAdvanced pre_run_advanced! {self.already_patched}")
# perform patches if not already patches
if not self.already_patched:
self.already_patched = True
def get_control(self, x_noisy: Tensor, t, cond, batched_number):
logger.info("In ControlLLLiteAdvanced get_control!")
# prepare timestep and everything related
self.prepare_current_timestep(t=t, batched_number=batched_number)
# perform other controlnets
control_prev = None
if self.previous_controlnet is not None:
control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number)
if control_prev is not None:
return control_prev
else:
return None
def get_models(self):
logger.info(f"In ControlLLLiteAdvanced get_models!")
# get_models is called once at the start of every KSampler run - use to reset already_patched status
self.already_patched = False
out = super().get_models()
return out
def copy(self):
c = ControlLLLiteAdvanced(self.timestep_keyframes)
self.copy_to(c)
self.copy_to_advanced(c)
return c
def cleanup(self):
super().cleanup()
self.cleanup_advanced()
self.already_patched = False
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)
@@ -335,6 +282,72 @@ class SparseCtrlAdvanced(ControlNetAdvanced):
return c
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):
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)
def patch_model(self, model: ModelPatcher):
model.set_model_attn1_patch(self.patch)
model.set_model_attn2_patch(self.patch)
def set_cond_hint(self, *args, **kwargs):
to_return = super().set_cond_hint(*args, **kwargs)
# cond hint for LLLite needs to be scaled between (-1, 1) instead of (0, 1)
self.cond_hint_original = self.cond_hint_original * 2.0 - 1.0
return to_return
def pre_run_advanced(self, *args, **kwargs):
AdvancedControlBase.pre_run_advanced(self, *args, **kwargs)
self.patch.control = self
def get_control_advanced(self, x_noisy: Tensor, t, cond, batched_number: int):
# normal ControlNet stuff
control_prev = None
if self.previous_controlnet is not None:
control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number)
if self.timestep_range is not None:
if t[0] > self.timestep_range[0] or t[0] < self.timestep_range[1]:
return control_prev
dtype = x_noisy.dtype
# prepare cond_hint
if self.sub_idxs is not None or self.cond_hint is None or x_noisy.shape[2] * 8 != self.cond_hint.shape[2] or x_noisy.shape[3] * 8 != self.cond_hint.shape[3]:
if self.cond_hint is not None:
del self.cond_hint
self.cond_hint = None
# if self.cond_hint_original length greater or equal to real latent count, subdivide it before scaling
if self.sub_idxs is not None and self.cond_hint_original.size(0) >= self.full_latent_length:
self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original[self.sub_idxs], x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(self.device)
else:
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)
# 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.
# return normal controlnet stuff
return control_prev
def cleanup_advanced(self):
super().cleanup_advanced()
self.patch.cleanup()
def copy(self):
c = ControlLLLiteAdvanced(self.patch, self.timestep_keyframes)
self.copy_to(c)
self.copy_to_advanced(c)
return c
# 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()
# return out
def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroup=None, model=None):
controlnet_data = comfy.utils.load_torch_file(ckpt_path, safe_load=True)
control = None
@@ -358,9 +371,7 @@ def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroup=None, mo
if controlnet_type != ControlWeightType.DEFAULT:
if controlnet_type == ControlWeightType.CONTROLLLLITE:
raise NotImplementedError("ControlLLLite has not been fully implemented yet!")
control = ControlLLLiteAdvanced(timestep_keyframes=timestep_keyframe)
# load Controll
control = load_controllllite(ckpt_path, controlnet_data=controlnet_data, timestep_keyframe=timestep_keyframe)
elif controlnet_type == ControlWeightType.SPARSECTRL:
control = load_sparsectrl(ckpt_path, controlnet_data=controlnet_data, timestep_keyframe=timestep_keyframe, model=model)
# otherwise, load vanilla ControlNet
@@ -542,3 +553,51 @@ def load_sparsectrl(ckpt_path: str, controlnet_data: dict[str, Tensor]=None, tim
control = SparseCtrlAdvanced(control_model, timestep_keyframes=timestep_keyframe, sparse_settings=sparse_settings, global_average_pooling=global_average_pooling, load_device=load_device, manual_cast_dtype=manual_cast_dtype)
return control
def load_controllllite(ckpt_path: str, controlnet_data: dict[str, Tensor]=None, timestep_keyframe: TimestepKeyframeGroup=None):
if controlnet_data is None:
controlnet_data = comfy.utils.load_torch_file(ckpt_path, safe_load=True)
# adapted from https://github.com/kohya-ss/ControlNet-LLLite-ComfyUI
# first, split weights for each module
module_weights = {}
for key, value in controlnet_data.items():
fragments = key.split(".")
module_name = fragments[0]
weight_name = ".".join(fragments[1:])
if module_name not in module_weights:
module_weights[module_name] = {}
module_weights[module_name][weight_name] = value
# next, load each module
modules = {}
for module_name, weights in module_weights.items():
# kohya planned to do something about how these should be chosen, so I'm not touching this
# since I am not familiar with the logic for this
if "conditioning1.4.weight" in weights:
depth = 3
elif weights["conditioning1.2.weight"].shape[-1] == 4:
depth = 2
else:
depth = 1
module = LLLiteModule(
name=module_name,
is_conv2d=weights["down.0.weight"].ndim == 4,
in_dim=weights["down.0.weight"].shape[1],
depth=depth,
cond_emb_dim=weights["conditioning1.0.weight"].shape[0] * 2,
mlp_dim=weights["down.0.weight"].shape[0],
)
# load weights into module
module.load_state_dict(weights)
modules[module_name] = module
if len(modules) == 1:
module.is_first = True
logger.info(f"loaded {ckpt_path} successfully, {len(modules)} modules")
patch = LLLitePatch(modules=modules)
control = ControlLLLiteAdvanced(patch=patch, timestep_keyframes=timestep_keyframe)
return control
+94 -114
View File
@@ -2,10 +2,16 @@
# basically, all the LLLite core code is from there, which I then combined with
# Advanced-ControlNet features and QoL
import math
from typing import Union
from torch import Tensor
import torch
import os
import comfy
import comfy.utils
from comfy.controlnet import ControlBase
from .logger import logger
from .utils import AdvancedControlBase, prepare_mask_batch
def extra_options_to_module_prefix(extra_options):
@@ -27,97 +33,69 @@ def extra_options_to_module_prefix(extra_options):
elif block[0] == "output":
module_pfx = f"lllite_unet_output_blocks_{block[1]}_1_transformer_blocks_{block_index}"
else:
raise Exception("invalid block name")
raise Exception(f"ControlLLLite: invalid block name '{block[0]}'. Expected 'input', 'middle', or 'output'.")
return module_pfx
def load_control_net_lllite_patch(path, cond_image, multiplier, num_steps, start_percent, end_percent):
# calculate start and end step
start_step = math.floor(num_steps * start_percent * 0.01) if start_percent > 0 else 0
end_step = math.floor(num_steps * end_percent * 0.01) if end_percent > 0 else num_steps
class LLLitePatch:
def __init__(self, modules: dict[str, 'LLLiteModule'], control: Union[AdvancedControlBase, ControlBase]=None):
self.modules = modules
self.control = control
def __call__(self, q, k, v, extra_options):
# 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
# a: Tensor = extra_options["sigmas"][0]
if self.control.t > self.control.timestep_range[0] or self.control.t < self.control.timestep_range[1]:
logger.info("Stopping short!!!")
return q, k, v
# load weights
ctrl_sd = comfy.utils.load_torch_file(path, safe_load=True)
module_pfx = extra_options_to_module_prefix(extra_options)
# split each weights for each module
module_weights = {}
for key, value in ctrl_sd.items():
fragments = key.split(".")
module_name = fragments[0]
weight_name = ".".join(fragments[1:])
if module_name not in module_weights:
module_weights[module_name] = {}
module_weights[module_name][weight_name] = value
# load each module
modules = {}
for module_name, weights in module_weights.items():
# kohya planned to do something about how these should be chosen, so I'm not touching this
# since I am not familiar with the logic for this
if "conditioning1.4.weight" in weights:
depth = 3
elif weights["conditioning1.2.weight"].shape[-1] == 4:
depth = 2
is_attn1 = q.shape[-1] == k.shape[-1] # self attention
if is_attn1:
module_pfx = module_pfx + "_attn1"
else:
depth = 1
module_pfx = module_pfx + "_attn2"
module = LLLiteModule(
name=module_name,
is_conv2d=weights["down.0.weight"].ndim == 4,
in_dim=weights["down.0.weight"].shape[1],
depth=depth,
cond_emb_dim=weights["conditioning1.0.weight"].shape[0] * 2,
mlp_dim=weights["down.0.weight"].shape[0],
multiplier=multiplier,
num_steps=num_steps,
start_step=start_step,
end_step=end_step,
)
info = module.load_state_dict(weights)
modules[module_name] = module
if len(modules) == 1:
module.is_first = True
module_pfx_to_q = module_pfx + "_to_q"
module_pfx_to_k = module_pfx + "_to_k"
module_pfx_to_v = module_pfx + "_to_v"
print(f"loaded {path} successfully, {len(modules)} modules")
# if masks present, get masks with same dims as attention
# if q.shape != k.shape or q.shape != v.shape:
# logger.warn(f"mismatch!!! q:{q.shape}, k:{k.shape}, v:{v.shape}")
#logger.warn(f"{q.shape}")
for module in modules.values():
module.set_cond_image(cond_image)
if module_pfx_to_q in self.modules:
q = q + self.modules[module_pfx_to_q](q, self.control)
if module_pfx_to_k in self.modules:
k = k + self.modules[module_pfx_to_k](k, self.control)
if module_pfx_to_v in self.modules:
v = v + self.modules[module_pfx_to_v](v, self.control)
class control_net_lllite_patch:
def __init__(self, modules):
self.modules = modules
return q, k, v
def __call__(self, q, k, v, extra_options):
module_pfx = extra_options_to_module_prefix(extra_options)
def to(self, device):
for d in self.modules.keys():
self.modules[d] = self.modules[d].to(device)
return self
def set_control(self, control: Union[AdvancedControlBase, ControlBase]):
self.control = control
is_attn1 = q.shape[-1] == k.shape[-1] # self attention
if is_attn1:
module_pfx = module_pfx + "_attn1"
else:
module_pfx = module_pfx + "_attn2"
def clone_with_control(self, control: AdvancedControlBase):
return LLLitePatch(self.modules, control)
module_pfx_to_q = module_pfx + "_to_q"
module_pfx_to_k = module_pfx + "_to_k"
module_pfx_to_v = module_pfx + "_to_v"
if module_pfx_to_q in self.modules:
q = q + self.modules[module_pfx_to_q](q)
if module_pfx_to_k in self.modules:
k = k + self.modules[module_pfx_to_k](k)
if module_pfx_to_v in self.modules:
v = v + self.modules[module_pfx_to_v](v)
return q, k, v
def to(self, device):
for d in self.modules.keys():
self.modules[d] = self.modules[d].to(device)
return self
return control_net_lllite_patch(modules)
def cleanup(self):
del self.control
self.control = None
for module in self.modules.values():
module.cleanup()
# TODO: use comfy.ops to support fp8 properly
class LLLiteModule(torch.nn.Module):
def __init__(
self,
@@ -127,18 +105,10 @@ class LLLiteModule(torch.nn.Module):
depth: int,
cond_emb_dim: int,
mlp_dim: int,
multiplier: int,
num_steps: int,
start_step: int,
end_step: int,
):
super().__init__()
self.name = name
self.is_conv2d = is_conv2d
self.multiplier = multiplier
self.num_steps = num_steps
self.start_step = start_step
self.end_step = end_step
self.is_first = False
modules = []
@@ -184,48 +154,47 @@ class LLLiteModule(torch.nn.Module):
)
self.depth = depth
self.cond_image = None
self.cond_emb = None
self.current_step = 0
self.cx_shape = None
self.prev_batch = 0
self.prev_sub_idxs = None
# @torch.inference_mode()
def set_cond_image(self, cond_image):
# print("set_cond_image", self.name)
self.cond_image = cond_image
def cleanup(self):
self.cond_emb = None
self.current_step = 0
self.cx_shape = None
self.prev_batch = 0
self.prev_sub_idxs = None
def forward(self, x):
if self.num_steps > 0:
if self.current_step < self.start_step:
self.current_step += 1
return torch.zeros_like(x)
elif self.current_step >= self.end_step:
if self.is_first and self.current_step == self.end_step:
print(f"end LLLite: step {self.current_step}")
self.current_step += 1
if self.current_step >= self.num_steps:
self.current_step = 0 # reset
return torch.zeros_like(x)
else:
if self.is_first and self.current_step == self.start_step:
print(f"start LLLite: step {self.current_step}")
self.current_step += 1
if self.current_step >= self.num_steps:
self.current_step = 0 # reset
if self.cond_emb is None:
def forward(self, x: Tensor, control: Union[AdvancedControlBase, ControlBase]):
mask = None
mask_tk = None
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(self.cond_image.to(x.device, dtype=x.dtype))
cx = self.conditioning1(control.cond_hint.to(x.device, dtype=x.dtype))
self.cx_shape = cx.shape
if not self.is_conv2d:
# reshape / b,c,h,w -> b,h*w,c
n, c, h, w = cx.shape
cx = cx.view(n, c, h * w).permute(0, 2, 1)
self.cond_emb = cx
# save prev values
self.prev_batch = x.shape[0]
self.prev_sub_idxs = control.sub_idxs
cx: torch.Tensor = self.cond_emb
# print(f"forward {self.name}, {cx.shape}, {x.shape}")
# TODO: make masks work for conv2d (could not find any ControlLLLites at this time that use them)
# create masks
if not self.is_conv2d:
n, c, h, w = self.cx_shape
if control.mask_cond_hint is not None:
mask = prepare_mask_batch(control.mask_cond_hint, (1, 1, h, w)).to(cx.dtype)
mask = mask.view(mask.shape[0], 1, h * w).permute(0, 2, 1)
if control.tk_mask_cond_hint is not None:
mask_tk = prepare_mask_batch(control.mask_cond_hint, (1, 1, h, w)).to(cx.dtype)
mask_tk = mask_tk.view(mask_tk.shape[0], 1, h * w).permute(0, 2, 1)
# x in uncond/cond doubles batch size
if x.shape[0] != cx.shape[0]:
if self.is_conv2d:
@@ -233,8 +202,19 @@ class LLLiteModule(torch.nn.Module):
else:
# print("x.shape[0] != cx.shape[0]", x.shape[0], cx.shape[0])
cx = cx.repeat(x.shape[0] // cx.shape[0], 1, 1)
if mask is not None:
mask = mask.repeat(x.shape[0] // mask.shape[0], 1, 1)
if mask_tk is not None:
mask_tk = mask_tk.repeat(x.shape[0] // mask_tk.shape[0], 1, 1)
if mask is None:
mask = 1.0
elif mask_tk is not None:
mask = mask * mask_tk
cx = torch.cat([cx, self.down(x)], dim=1 if self.is_conv2d else 2)
cx = self.mid(cx)
cx = self.up(cx)
return cx * self.multiplier
if control.latent_keyframes is not None:
cx = cx * control.calc_latent_keyframe_mults(x=cx, batched_number=control.batched_number)
return cx * mask * control.strength * control.current_timestep_keyframe.strength
+16 -5
View File
@@ -2,6 +2,7 @@ import numpy as np
from torch import Tensor
import folder_paths
from comfy.model_patcher import ModelPatcher
from .control import load_controlnet, convert_to_advanced, is_advanced_controlnet
from .utils import ControlWeights, ControlWeightType, LatentKeyframeGroup, TimestepKeyframe, TimestepKeyframeGroup
@@ -136,21 +137,24 @@ class AdvancedControlNetApply:
"timestep_kf": ("TIMESTEP_KEYFRAME", ),
"latent_kf_override": ("LATENT_KEYFRAME", ),
"weights_override": ("CONTROL_NET_WEIGHTS", ),
"model_optional": ("MODEL",),
}
}
RETURN_TYPES = ("CONDITIONING","CONDITIONING")
RETURN_NAMES = ("positive", "negative")
RETURN_TYPES = ("CONDITIONING","CONDITIONING","MODEL",)
RETURN_NAMES = ("positive", "negative", "model_opt")
FUNCTION = "apply_controlnet"
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝"
def apply_controlnet(self, positive, negative, control_net, image, strength, start_percent, end_percent,
mask_optional: Tensor=None,
mask_optional: Tensor=None, model_optional: ModelPatcher=None,
timestep_kf: TimestepKeyframeGroup=None, latent_kf_override: LatentKeyframeGroup=None,
weights_override: ControlWeights=None):
if strength == 0:
return (positive, negative)
return (positive, negative, model_optional)
if model_optional:
model_optional = model_optional.clone()
control_hint = image.movedim(-1,1)
cnets = {}
@@ -168,6 +172,13 @@ class AdvancedControlNetApply:
# copy, convert to advanced if needed, and set cond
c_net = convert_to_advanced(control_net.copy()).set_cond_hint(control_hint, strength, (start_percent, end_percent))
if is_advanced_controlnet(c_net):
# disarm node check
c_net.disarm()
# if model required, verify model is passed in, and if so patch it
if c_net.require_model:
if not model_optional:
raise Exception(f"Type '{type(c_net).__name__}' requires model_optional input, but got None.")
c_net.patch_model(model=model_optional)
# apply optional parameters and overrides, if provided
if timestep_kf is not None:
c_net.set_timestep_keyframes(timestep_kf)
@@ -192,7 +203,7 @@ class AdvancedControlNetApply:
n = [t[0], d]
c.append(n)
out.append(c)
return (out[0], out[1])
return (out[0], out[1], model_optional)
# NODE MAPPING
+47 -17
View File
@@ -2,10 +2,13 @@ from typing import Callable, Union
import torch
from torch import Tensor
import torch.nn.functional as F
import comfy.ops
import comfy.utils
from comfy.controlnet import ControlBase, broadcast_image_to
from comfy.model_patcher import ModelPatcher
from .logger import logger
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):
@@ -262,7 +265,7 @@ class WeightTypeException(TypeError):
class AdvancedControlBase:
def __init__(self, base: ControlBase, timestep_keyframes: TimestepKeyframeGroup, weights_default: ControlWeights):
def __init__(self, base: ControlBase, timestep_keyframes: TimestepKeyframeGroup, weights_default: ControlWeights, require_model=False):
self.base = base
self.compatible_weights = [ControlWeightType.UNIVERSAL]
self.add_compatible_weight(weights_default.weight_type)
@@ -293,6 +296,14 @@ class AdvancedControlBase:
self.control_merge = self.control_merge_inject#.__get__(self, type(self))
self.pre_run = self.pre_run_inject
self.cleanup = self.cleanup_inject
self.set_previous_controlnet = self.set_previous_controlnet_inject
# require model to be passed into Apply Advanced ControlNet 🛂🅐🅒🅝 node
self.require_model = require_model
# disarm - when set to False, used to force usage of Apply Advanced ControlNet 🛂🅐🅒🅝 node (which will set it to True)
self.disarmed = not require_model
def patch_model(self, model: ModelPatcher):
pass
def add_compatible_weight(self, control_weight_type: str):
self.compatible_weights.append(control_weight_type)
@@ -322,10 +333,10 @@ class AdvancedControlBase:
self.latent_keyframes = None
def prepare_current_timestep(self, t: Tensor, batched_number: int):
self.t = t
self.t = float(t[0])
self.batched_number = batched_number
# get current step percent
curr_t: float = t[0]
curr_t: float = self.t
prev_index = self.current_timestep_index
# if has next index, loop through and see if need to switch
if self.timestep_keyframes.has_index(self.current_timestep_index+1):
@@ -395,23 +406,32 @@ class AdvancedControlBase:
# clear variables
self.cleanup_advanced()
def set_previous_controlnet_inject(self, *args, **kwargs):
to_return = self.base.set_previous_controlnet(*args, **kwargs)
if not self.disarmed:
raise Exception(f"Type '{type(self).__name__}' must be used with Apply Advanced ControlNet 🛂🅐🅒🅝 node (with model_optional passed in); otherwise, it will not work.")
return to_return
def disarm(self):
self.disarmed = True
def get_control_inject(self, x_noisy, t, cond, batched_number):
# prepare timestep and everything related
self.prepare_current_timestep(t=t, batched_number=batched_number)
# if should not perform any actions for the controlnet, exit without doing any work
if self.strength == 0.0 or self.current_timestep_keyframe.strength == 0.0:
control_prev = None
if self.previous_controlnet is not None:
control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number)
if control_prev is not None:
return control_prev
else:
return None
return self.default_control_actions(x_noisy, t, cond, batched_number)
# otherwise, perform normal function
return self.get_control_advanced(x_noisy, t, cond, batched_number)
def get_control_advanced(self, x_noisy, t, cond, batched_number):
pass
return self.default_control_actions(x_noisy, t, cond, batched_number)
def default_control_actions(self, x_noisy, t, cond, batched_number):
control_prev = None
if self.previous_controlnet is not None:
control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number)
return control_prev
def calc_weight(self, idx: int, x: Tensor, layers: int) -> Union[float, Tensor]:
if self.weights.weight_mask is not None:
@@ -424,11 +444,12 @@ class AdvancedControlBase:
def get_calc_pow(self, idx: int, layers: int) -> int:
return (layers-1)-idx
def apply_advanced_strengths_and_masks(self, x: Tensor, batched_number: int):
def calc_latent_keyframe_mults(self, x: Tensor, batched_number: int) -> Tensor:
# apply strengths, and get batch indeces to null out
# AKA latents that should not be influenced by ControlNet
if self.latent_keyframes is not None:
latent_count = x.size(0)//batched_number
final_mults = [1.0] * x.shape[0]
if self.latent_keyframes:
latent_count = x.shape[0] // batched_number
indeces_to_null = set(range(latent_count))
mapped_indeces = None
# if expecting subdivision, will need to translate between subset and actual idx values
@@ -459,13 +480,21 @@ class AdvancedControlBase:
# apply strength for each batched cond/uncond
for b in range(batched_number):
x[(latent_count*b)+real_index] = x[(latent_count*b)+real_index] * keyframe.strength
final_mults[(latent_count*b)+real_index] = keyframe.strength
# null them out by multiplying by null_latent_kf_strength
for batch_index in indeces_to_null:
# apply null for each batched cond/uncond
for b in range(batched_number):
x[(latent_count*b)+batch_index] = x[(latent_count*b)+batch_index] * self.current_timestep_keyframe.null_latent_kf_strength
final_mults[(latent_count*b)+batch_index] = self.current_timestep_keyframe.null_latent_kf_strength
# convert final_mults into tensor and match expected dimension count
final_tensor = torch.tensor(final_mults, dtype=x.dtype, device=x.device)
while len(final_tensor.shape) < len(x.shape):
final_tensor = final_tensor.unsqueeze(-1)
return final_tensor
def apply_advanced_strengths_and_masks(self, x: Tensor, batched_number: int):
if self.latent_keyframes is not None:
x[:] = x[:] * self.calc_latent_keyframe_mults(x=x, batched_number=batched_number)
# apply masks, resizing mask to required dims
if self.mask_cond_hint is not None:
masks = prepare_mask_batch(self.mask_cond_hint, x.shape)
@@ -602,3 +631,4 @@ class AdvancedControlBase:
copied.mask_cond_hint_original = self.mask_cond_hint_original
copied.weights_override = self.weights_override
copied.latent_keyframe_override = self.latent_keyframe_override
copied.disarmed = self.disarmed