Added full ControlLLLite support, refactored some code to eliminate code repetition and support controls that require model patching
This commit is contained in:
+117
-58
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user