Updated to support sliding window usage by other extensions, added scaffolding code to support masks in controlnet, added logger for future use

This commit is contained in:
Jedrzej Kosinski
2023-09-18 11:23:01 -05:00
parent 3cc9679b57
commit 15e438100f
3 changed files with 255 additions and 131 deletions
+98 -10
View File
@@ -9,14 +9,12 @@ import inspect
from ldm.modules.diffusionmodules.util import timestep_embedding
sys.path.insert(0, os.path.join(os.path.dirname(os.path.realpath(__file__)), "comfy"))
from comfy.cldm import cldm
from comfy.model_patcher import ModelPatcher
from comfy.controlnet import ControlBase, ControlNet, T2IAdapter, broadcast_image_to, ControlLora
import comfy.t2i_adapter as t2i_adapter
import comfy.utils as utils
import comfy.utils
import comfy.model_management as model_management
import comfy.model_detection as model_detection
@@ -174,15 +172,61 @@ class ControlNetAdvanced(ControlNet):
self.weights = self.timestep_keyframes.keyframes[0].control_net_weights if self.timestep_keyframes.keyframes[0].control_net_weights else [1.0]*13
# mask for which parts of controlnet output to keep
self.cond_hint_mask = None
# actual index values
self.sub_idxs = None
self.full_latent_length = 0
self.context_length = 0
# override control_merge
self.control_merge = control_merge_inject.__get__(self, type(self))
def set_cond_hint_mask(self, mask_hint):
self.cond_hint_mask = mask_hint
return self
def get_control(self, x_noisy, t, cond, batched_number):
# need to reference t and batched_number later
self.t = t
self.batched_number = batched_number
# TODO: choose TimestepKeyframe based on t
return super().get_control(x_noisy, t, cond, batched_number)
if self.sub_idxs is not None:
# perform special version of get_control
return self.sliding_get_control(x_noisy, t, cond, batched_number)
else:
return super().get_control(x_noisy, t, cond, batched_number)
def sliding_get_control(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)
if self.timestep_range is not None:
if t[0] > self.timestep_range[0] or t[0] < self.timestep_range[1]:
if control_prev is not None:
return control_prev
else:
return None
output_dtype = x_noisy.dtype
# TODO: change this to not require cond_hint upscaling every step
if self.sub_idxs is not None or self.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
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(self.control_model.dtype).to(self.device)
# if self.cond_hint length matches real latent count, need to subdivide it
if self.cond_hint.size(0) == self.full_latent_length:
self.cond_hint = self.cond_hint[self.sub_idxs]
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)
context = cond['c_crossattn']
y = cond.get('c_adm', None)
if y is not None:
y = y.to(self.control_model.dtype)
control = self.control_model(x=x_noisy.to(self.control_model.dtype), hint=self.cond_hint, timesteps=t, context=context.to(self.control_model.dtype), y=y)
return self.control_merge(None, control, control_prev, output_dtype)
def apply_advanced_strengths_and_masks(self, x, current_timestep_keyframe: TimestepKeyframe, batched_number: int):
if current_timestep_keyframe.latent_keyframes is not None:
@@ -191,12 +235,28 @@ class ControlNetAdvanced(ControlNet):
latent_count = x.size(0)//batched_number
indeces_to_zero = set(range(latent_count))
mapped_indeces = None
# if expecting subdivision, will need to translate between subset and actual idx values
if self.sub_idxs:
mapped_indeces = {}
for i, actual in enumerate(self.sub_idxs):
mapped_indeces[actual] = i
for keyframe in current_timestep_keyframe.latent_keyframes:
if keyframe.batch_index in indeces_to_zero:
indeces_to_zero.remove(keyframe.batch_index)
real_index = keyframe.batch_index
# if not mapping indeces, what you see is what you get
if mapped_indeces is None:
if real_index in indeces_to_zero:
indeces_to_zero.remove(keyframe.batch_index)
# otherwise, see if batch_index is even included in this set of latents
else:
real_index = mapped_indeces.get(keyframe.batch_index, None)
if real_index is None:
continue
indeces_to_zero.remove(real_index)
# apply strength for each batched cond/uncond
for b in range(batched_number):
x[(latent_count*b)+keyframe.batch_index] *= keyframe.strength
x[(latent_count*b)+real_index] *= keyframe.strength
# zero them out by multiplying by zero
for batch_index in indeces_to_zero:
@@ -213,6 +273,12 @@ class ControlNetAdvanced(ControlNet):
out = super().get_models()
out.append(self.control_model_wrapped)
return out
def cleanup(self):
super().cleanup()
self.sub_idxs = None
self.full_latent_length = 0
self.context_length = 0
class T2IAdapterAdvanced(T2IAdapter):
@@ -224,6 +290,10 @@ class T2IAdapterAdvanced(T2IAdapter):
self.weights = first_weight if first_weight else [1.0]*12
# mask for which parts of controlnet output to keep
self.cond_hint_mask = None
# actual index values
self.sub_idxs = None
self.full_latent_length = 0
self.context_length = 0
# override control_merge
self.control_merge = control_merge_inject.__get__(self, type(self))
@@ -232,7 +302,19 @@ class T2IAdapterAdvanced(T2IAdapter):
self.t = t
self.batched_number = batched_number
# TODO: choose TimestepKeyframe based on t
return super().get_control(x_noisy, t, cond, batched_number)
try:
# if sub indexes present, replace original hint with subsection
if self.sub_idxs is not None:
full_cond_hint_original = self.cond_hint_original
del self.cond_hint
self.cond_hint = None
self.cond_hint_original = full_cond_hint_original[self.sub_idxs]
return super().get_control(x_noisy, t, cond, batched_number)
finally:
if self.sub_idxs is not None:
# replace original cond hint
self.cond_hint_original = full_cond_hint_original
del full_cond_hint_original
def apply_advanced_strengths_and_masks(self, x, current_timestep_keyframe: TimestepKeyframe, batched_number: int):
# For now, do nothing; need to figure out LatentKeyframe control is even possible for T2I Adapters
@@ -242,10 +324,16 @@ class T2IAdapterAdvanced(T2IAdapter):
c = T2IAdapterAdvanced(self.t2i_model, self.timestep_keyframes, self.channels_in)
self.copy_to(c)
return c
def cleanup(self):
super().cleanup()
self.sub_idxs = None
self.full_latent_length = 0
self.context_length = 0
def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroup=None, model=None):
controlnet_data = utils.load_torch_file(ckpt_path, safe_load=True)
controlnet_data = comfy.utils.load_torch_file(ckpt_path, safe_load=True)
if "lora_controlnet" in controlnet_data:
return ControlLora(controlnet_data) # TODO: apply weights to ControlLora
@@ -253,7 +341,7 @@ def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroup=None, mo
if "controlnet_cond_embedding.conv_in.weight" in controlnet_data: #diffusers format
use_fp16 = model_management.should_use_fp16()
controlnet_config = model_detection.unet_config_from_diffusers_unet(controlnet_data, use_fp16)
diffusers_keys = utils.unet_to_diffusers(controlnet_config)
diffusers_keys = comfy.utils.unet_to_diffusers(controlnet_config)
diffusers_keys["controlnet_mid_block.weight"] = "middle_block_out.0.weight"
diffusers_keys["controlnet_mid_block.bias"] = "middle_block_out.0.bias"
+36
View File
@@ -0,0 +1,36 @@
import sys
import copy
import logging
class ColoredFormatter(logging.Formatter):
COLORS = {
"DEBUG": "\033[0;36m", # CYAN
"INFO": "\033[0;32m", # GREEN
"WARNING": "\033[0;33m", # YELLOW
"ERROR": "\033[0;31m", # RED
"CRITICAL": "\033[0;37;41m", # WHITE ON RED
"RESET": "\033[0m", # RESET COLOR
}
def format(self, record):
colored_record = copy.copy(record)
levelname = colored_record.levelname
seq = self.COLORS.get(levelname, self.COLORS["RESET"])
colored_record.levelname = f"{seq}{levelname}{self.COLORS['RESET']}"
return super().format(colored_record)
# Create a new logger
logger = logging.getLogger("Advanced-ControlNet")
logger.propagate = False
# Add handler if we don't have one.
if not logger.handlers:
handler = logging.StreamHandler(sys.stdout)
handler.setFormatter(ColoredFormatter("[%(name)s] - %(levelname)s - %(message)s"))
logger.addHandler(handler)
# Configure logger
loglevel = logging.INFO
logger.setLevel(loglevel)
+121 -121
View File
@@ -8,11 +8,9 @@ from PIL import Image, ImageOps
import folder_paths
sys.path.insert(0, os.path.join(os.path.dirname(os.path.realpath(__file__)), "comfy"))
from .control import load_controlnet, ControlNetWeightsType, T2IAdapterWeightsType,\
from .control import ControlNetAdvanced, T2IAdapterAdvanced, load_controlnet, ControlNetWeightsType, T2IAdapterWeightsType,\
LatentKeyframe, LatentKeyframeGroup, TimestepKeyframe, TimestepKeyframeGroup
from .logger import logger
def get_properly_arranged_t2i_weights(initial_weights: list[float]):
new_weights = []
@@ -231,6 +229,106 @@ class LatentKeyframeNode:
return (prev_latent_keyframe,)
class LatentKeyframeGroupNode:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"index_strengths": ("STRING", {"multiline": True, "default": ""}),
},
"optional": {
"prev_latent_keyframe": ("LATENT_KEYFRAME", ),
"latent_image_opt": ("LATENT", ),
}
}
RETURN_TYPES = ("LATENT_KEYFRAME", )
FUNCTION = "load_keyframes"
CATEGORY = "adv-controlnet/keyframes"
def validate_index(self, index: int, latent_count: int = 0, is_range: bool = False, allow_negative = False) -> int:
# if part of range, do nothing
if is_range:
return index
# otherwise, validate index
# validate not out of range - only when latent_count is passed in
if latent_count > 0 and index > latent_count-1:
raise IndexError(f"Index '{index}' out of range for the total {latent_count} latents.")
# if negative, validate not out of range
if index < 0:
if not allow_negative:
raise IndexError(f"Negative indeces not allowed, but was {index}.")
conv_index = latent_count+index
if conv_index < 0:
raise IndexError(f"Index '{index}', converted to '{conv_index}' out of range for the total {latent_count} latents.")
index = conv_index
return index
def convert_to_index_int(self, raw_index: str, latent_count: int = 0, is_range: bool = False, allow_negative = False) -> int:
try:
return self.validate_index(int(raw_index), latent_count=latent_count, is_range=is_range, allow_negative=allow_negative)
except ValueError as e:
raise ValueError(f"index '{raw_index}' must be an integer.", e)
def convert_to_latent_keyframes(self, latent_indeces: str, latent_count: int) -> set[LatentKeyframe]:
if not latent_indeces:
return set()
all_indeces = [i for i in range(0, latent_count)]
allow_negative = latent_count > 0
chosen_indeces = set()
# parse string - allow positive ints, negative ints, and ranges separated by ':'
groups = latent_indeces.split(",")
groups = [g.strip() for g in groups]
for g in groups:
# parse strengths - default to 1.0 if no strength given
strength = 1.0
if '=' in g:
g, strength_str = g.split("=", 1)
g = g.strip()
try:
strength = float(strength_str.strip())
except ValueError as e:
raise ValueError(f"strength '{strength_str}' must be a float.", e)
if strength < 0:
raise ValueError(f"Strength '{strength}' cannot be negative.")
# parse range of indeces (e.g. 2:16)
if ':' in g:
index_range = g.split(":", 1)
index_range = [r.strip() for r in index_range]
start_index = self.convert_to_index_int(index_range[0], latent_count=latent_count, is_range=True, allow_negative=allow_negative)
end_index = self.convert_to_index_int(index_range[1], latent_count=latent_count, is_range=True, allow_negative=allow_negative)
for i in all_indeces[start_index:end_index]:
chosen_indeces.add(LatentKeyframe(i, strength))
# parse individual indeces
else:
chosen_indeces.add(LatentKeyframe(self.convert_to_index_int(g, latent_count=latent_count, allow_negative=allow_negative), strength))
return chosen_indeces
def load_keyframes(self,
index_strengths: str,
prev_latent_keyframe: LatentKeyframeGroup=None,
latent_image_opt=None):
if not prev_latent_keyframe:
prev_latent_keyframe = LatentKeyframeGroup()
curr_latent_keyframe = LatentKeyframeGroup()
latent_count = -1
if latent_image_opt:
latent_count = latent_image_opt['samples'].size()[0]
latent_keyframes = self.convert_to_latent_keyframes(index_strengths, latent_count=latent_count)
for latent_keyframe in latent_keyframes:
curr_latent_keyframe.add(latent_keyframe)
for latent_keyframe in prev_latent_keyframe.keyframes:
curr_latent_keyframe.add(latent_keyframe)
return (curr_latent_keyframe,)
class ControlNetLoaderAdvanced:
@classmethod
def INPUT_TYPES(s):
@@ -288,10 +386,12 @@ class ControlNetApplyAdvanced_AdvControlNet:
"negative": ("CONDITIONING", ),
"control_net": ("CONTROL_NET", ),
"image": ("IMAGE", ),
"mask_opt": ("MASK", ),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001})
},
"optional": {
"mask_opt": ("MASK", ),
}
}
@@ -301,10 +401,12 @@ class ControlNetApplyAdvanced_AdvControlNet:
CATEGORY = "adv-controlnet/loaders/conditioning"
def apply_controlnet(self, positive, negative, control_net, image, mask_opt, strength, start_percent, end_percent):
def apply_controlnet(self, positive, negative, control_net, image, strength, start_percent, end_percent, mask_opt=None):
if strength == 0:
return (positive, negative)
if mask_opt is not None:
mask_hint = mask_opt.movedim(-1,1)
control_hint = image.movedim(-1,1)
cnets = {}
@@ -319,6 +421,12 @@ class ControlNetApplyAdvanced_AdvControlNet:
c_net = cnets[prev_cnet]
else:
c_net = control_net.copy().set_cond_hint(control_hint, strength, (1.0 - start_percent, 1.0 - end_percent))
# TODO: finish mask implemention, does nothing right now
if mask_opt is not None:
if isinstance(c_net, ControlNetAdvanced) or isinstance(c_net, T2IAdapterAdvanced):
c_net.set_cond_hint_mask(mask_hint)
else:
logger
c_net.set_previous_controlnet(prev_cnet)
cnets[prev_cnet] = c_net
@@ -330,115 +438,6 @@ class ControlNetApplyAdvanced_AdvControlNet:
return (out[0], out[1])
class ControlNetApplyPartialBatch: # NOT USED: was used for a different test, has useful index parsing code though
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"positive": ("CONDITIONING", ),
"negative": ("CONDITIONING", ),
"control_net": ("CONTROL_NET", ),
"image": ("IMAGE", ),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001})
},
"optional": {
"latent_image": ("LATENT", ),
"latent_indeces": ("STRING", {"default": ""}),
}
}
RETURN_TYPES = ("CONDITIONING","CONDITIONING")
RETURN_NAMES = ("positive", "negative")
FUNCTION = "apply_controlnet"
CATEGORY = "adv-controlnet/conditioning"
def validate_index(self, index: int, latent_count: int, is_range: bool = False) -> int:
# if part of range, do nothing
if is_range:
return index
# otherwise, validate index
# validate not out of range
if index > latent_count-1:
raise IndexError(f"Index '{index}' out of range for the total {latent_count} latents.")
# if negative, validate not out of range
if index < 0:
conv_index = latent_count+index
if conv_index < 0:
raise IndexError(f"Index '{index}', converted to '{conv_index}' out of range for the total {latent_count} latents.")
index = conv_index
return index
def convert_to_index_int(self, raw_index: str, is_range: bool = False) -> int:
try:
return self.validate_index(int(raw_index), is_range=is_range)
except ValueError as e:
raise ValueError(f"index '{raw_index}' must be an integer.", e)
def convert_to_indeces(self, latent_indeces: str, latent_count: int) -> set[int]:
if not latent_indeces:
return set()
all_indeces = [i for i in range(0, latent_count)]
chosen_indeces = set()
# parse string - allow positive ints, negative ints, and ranges separated by ':'
groups = latent_indeces.split(",")
groups = [g.strip() for g in groups]
for g in groups:
# parse range of indeces (e.g. 2:16)
if ':' in g:
index_range = g.split(":", 1)
index_range = [r.strip() for r in index_range]
start_index = self.convert_to_index_int(index_range[0], is_range=True)
end_index = self.convert_to_index_int(index_range[1], is_range=True)
for i in all_indeces[start_index, end_index]:
chosen_indeces.add(i)
# parse individual indeces
else:
chosen_indeces.add(self.convert_to_index_int(g))
return chosen_indeces
def apply_controlnet(self, positive, negative, control_net, image, strength, start_percent, end_percent, latent_image=None, latent_indeces: str=None):
if strength == 0:
return (positive, negative)
latent_count = 1
if latent_image:
latent_count = latent_image['samples'].size()[0]
indeces_to_apply = self.convert_to_indeces(latent_indeces, latent_count)
control_hint = image.movedim(-1,1)
cnets = {}
evaluating_positive = True
out = []
for conditioning in [positive, negative]:
c = []
if evaluating_positive and latent_count > 1:
# should copy positive conditioning to match latent_count
if len(conditioning) < latent_count:
pass
for t in conditioning:
d = t[1].copy()
prev_cnet = d.get('control', None)
if prev_cnet in cnets:
c_net = cnets[prev_cnet]
else:
c_net = control_net.copy().set_cond_hint(control_hint, strength, (1.0 - start_percent, 1.0 - end_percent))
c_net.set_previous_controlnet(prev_cnet)
cnets[prev_cnet] = c_net
d['control'] = c_net
d['control_apply_to_uncond'] = False
n = [t[0], d]
c.append(n)
evaluating_positive = False
out.append(c)
return (out[0], out[1])
class LoadImagesFromDirectory:
@classmethod
def INPUT_TYPES(s):
@@ -447,7 +446,8 @@ class LoadImagesFromDirectory:
"directory": ("STRING", {"default": ""}),
},
"optional": {
"image_load_cap": ("INT", {"default": 0, "min": 0, "step": 1})
"image_load_cap": ("INT", {"default": 0, "min": 0, "step": 1}),
"start_index": ("INT", {"default": 0, "min": 0, "step": 1}),
}
}
@@ -456,7 +456,7 @@ class LoadImagesFromDirectory:
CATEGORY = "adv-controlnet/image"
def load_images(self, directory: str, image_load_cap: int = 0):
def load_images(self, directory: str, image_load_cap: int = 0, start_index: int = 0):
if not os.path.isdir(directory):
raise FileNotFoundError(f"Directory '{directory} cannot be found.'")
dir_files = os.listdir(directory)
@@ -465,6 +465,8 @@ class LoadImagesFromDirectory:
dir_files = sorted(dir_files)
dir_files = [os.path.join(directory, x) for x in dir_files]
# start at start_index
dir_files = dir_files[start_index:]
images = []
masks = []
@@ -506,8 +508,7 @@ NODE_CLASS_MAPPINGS = {
# Keyframes
"TimestepKeyframe": TimestepKeyframeNode,
"LatentKeyframe": LatentKeyframeNode,
# Conditioning
# "ControlNetApplyPartialBatch": ControlNetApplyPartialBatch,
"LatentKeyframeGroup": LatentKeyframeGroupNode,
# Loaders
"ControlNetLoaderAdvanced": ControlNetLoaderAdvanced,
"DiffControlNetLoaderAdvanced": DiffControlNetLoaderAdvanced,
@@ -525,8 +526,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
# Keyframes
"TimestepKeyframe": "Timestep Keyframe",
"LatentKeyframe": "Latent Keyframe",
# Conditioning
# "ControlNetApplyPartialBatch": "Apply ControlNet (Partial Batch)",
"LatentKeyframeGroup": "Latent Keyframe Group",
# Loaders
"ControlNetLoaderAdvanced": "Load ControlNet Model (Advanced)",
"DiffControlNetLoaderAdvanced": "Load ControlNet Model (diff Advanced)",