Added LatentKeyframe and TimestepKeyframe nodes, updated Load ControlNet (Advanced) node to accept timestep_keyframe, added functionality to limit ControlNet's influence to one (or more) specific latents in a batch at specified strengths (useful for AnimateDiff)

This commit is contained in:
Jedrzej Kosinski
2023-08-26 07:04:13 -05:00
parent edc0a20815
commit 354a8ce10c
2 changed files with 312 additions and 26 deletions
+117 -12
View File
@@ -7,6 +7,7 @@ import contextlib
import copy
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"))
@@ -18,18 +19,104 @@ import comfy.utils as utils
import comfy.model_management as model_management
import comfy.model_detection as model_detection
ControlNetWeightsType = list[float]
T2IAdapterWeightsType = list[float]
class LatentKeyframe:
def __init__(self, batch_index: int, strength: float) -> None:
self.batch_index = batch_index
self.strength = strength
# always maintain sorted state (by batch_index of LatentKeyframe)
class LatentKeyframeGroup:
def __init__(self) -> None:
self.keyframes: list[LatentKeyframe] = []
def add(self, keyframe: LatentKeyframe) -> None:
added = False
# replace existing keyframe if same batch_index
for i in range(len(self.keyframes)):
if self.keyframes[i].batch_index == keyframe.batch_index:
self.keyframes[i] = keyframe
added = True
break
if not added:
self.keyframes.append(keyframe)
self.keyframes.sort(key=lambda k: k.batch_index)
def get_index(self, index: int) -> LatentKeyframe | None:
try:
return self.keyframes[index]
except IndexError:
return None
def __getitem__(self, index) -> LatentKeyframe:
return self.keyframes[index]
def is_empty(self) -> bool:
return len(self.keyframes) == 0
class TimestepKeyframe:
def __init__(self,
start_percent: float,
control_net_weights: ControlNetWeightsType = None,
t2i_adapter_weights: T2IAdapterWeightsType = None,
latent_keyframes: LatentKeyframeGroup = None) -> None:
self.start_percent = start_percent
self.control_net_weights = control_net_weights
self.t2i_adapter_weights = t2i_adapter_weights
self.latent_keyframes = latent_keyframes
@classmethod
def default(cls) -> 'TimestepKeyframe':
return cls(0.0)
# always maintain sorted state (by start_percent of TimestepKeyFrame)
class TimestepKeyframeGroup:
def __init__(self) -> None:
self.keyframes: list[TimestepKeyframe] = []
self.keyframes.append(TimestepKeyframe.default())
def add(self, keyframe: TimestepKeyframe) -> None:
added = False
# replace existing keyframe if same start_percent
for i in range(len(self.keyframes)):
if self.keyframes[i].start_percent == keyframe.start_percent:
self.keyframes[i] = keyframe
added = True
break
if not added:
self.keyframes.append(keyframe)
self.keyframes.sort(key=lambda k: k.start_percent)
def get_index(self, index: int) -> TimestepKeyframe | None:
try:
return self.keyframes[index]
except IndexError:
return None
def __getitem__(self, index) -> TimestepKeyframe:
return self.keyframes[index]
def is_empty(self) -> bool:
return len(self.keyframes) == 0
# Copied from comfy.sd, weights modified
class ControlNetAdvanced(ControlBase):
def __init__(self, control_model, weights: ControlNetWeightsType, global_average_pooling=False, device=None):
def __init__(self, control_model, timestep_keyframes: TimestepKeyframeGroup, global_average_pooling=False, device=None):
super().__init__(device)
self.control_model = control_model
self.control_model_wrapped = ModelPatcher(self.control_model, load_device=model_management.get_torch_device(), offload_device=model_management.unet_offload_device())
self.weights = weights if weights else [1.0]*13
self.timestep_keyframes = timestep_keyframes if timestep_keyframes else TimestepKeyframeGroup()
self.weights = self.timestep_keyframes.keyframes[0].control_net_weights if self.timestep_keyframes.keyframes[0].control_net_weights else [1.0]*13
self.global_average_pooling = global_average_pooling
def get_control(self, x_noisy, t, cond, batched_number):
@@ -58,6 +145,9 @@ class ControlNetAdvanced(ControlBase):
else:
precision_scope = contextlib.nullcontext
# TODO: select based on progress in diffusion
current_timestep_keyframe = self.timestep_keyframes[0]
with precision_scope(model_management.get_autocast_device(self.device)):
context = torch.cat(cond['c_crossattn'], 1)
y = cond.get('c_adm', None)
@@ -76,6 +166,17 @@ class ControlNetAdvanced(ControlBase):
if self.global_average_pooling:
x = torch.mean(x, dim=(2, 3), keepdim=True).repeat(1, 1, x.shape[2], x.shape[3])
# get batch indeces to zero out, AKA latents that should not be influenced by ControlNet
indeces_to_zero = set(range(x.size()[0]//2))
for keyframe in current_timestep_keyframe.latent_keyframes:
if keyframe.batch_index in indeces_to_zero:
indeces_to_zero.remove(keyframe.batch_index)
# zero them out by multiplying by zero
for batch_index in indeces_to_zero:
x[batch_index] *= 0.0
x[(x.size()[0]//2) + batch_index] *= 0.0
x *= self.strength * self.weights[i]
if x.dtype != output_dtype and not autocast_enabled:
x = x.to(output_dtype)
@@ -90,7 +191,7 @@ class ControlNetAdvanced(ControlBase):
return out
def copy(self):
c = ControlNetAdvanced(self.control_model, self.weights, global_average_pooling=self.global_average_pooling)
c = ControlNetAdvanced(self.control_model, self.timestep_keyframes, global_average_pooling=self.global_average_pooling)
self.copy_to(c)
return c
@@ -100,7 +201,7 @@ class ControlNetAdvanced(ControlBase):
return out
def load_controlnet(ckpt_path, control_net_weights: ControlNetWeightsType=None, t2i_adapter_weights: T2IAdapterWeightsType=None, model=None):
def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroup=None, model=None):
controlnet_data = utils.load_torch_file(ckpt_path, safe_load=True)
if "lora_controlnet" in controlnet_data:
return ControlLora(controlnet_data) # TODO: apply weights to ControlLora
@@ -162,7 +263,7 @@ def load_controlnet(ckpt_path, control_net_weights: ControlNetWeightsType=None,
elif key in controlnet_data:
prefix = ""
else:
net = load_t2i_adapter(controlnet_data)
net = load_t2i_adapter(controlnet_data, timestep_keyframe)
if net is None:
print("error checkpoint does not contain controlnet or t2i adapter data", ckpt_path)
return net
@@ -205,16 +306,20 @@ def load_controlnet(ckpt_path, control_net_weights: ControlNetWeightsType=None,
if ckpt_path.endswith("_shuffle.pth") or ckpt_path.endswith("_shuffle.safetensors") or ckpt_path.endswith("_shuffle_fp16.safetensors"): #TODO: smarter way of enabling global_average_pooling
global_average_pooling = True
control = ControlNetAdvanced(control_model, control_net_weights, global_average_pooling=global_average_pooling)
control = ControlNetAdvanced(control_model, timestep_keyframe, global_average_pooling=global_average_pooling)
return control
# Copied from comfy.sd, weights modified
class T2IAdapterAdvanced(ControlBase):
def __init__(self, t2i_model, weights: T2IAdapterWeightsType, channels_in, device=None):
def __init__(self, t2i_model, timestep_keyframes: TimestepKeyframeGroup, channels_in, device=None):
super().__init__(device)
self.t2i_model = t2i_model
self.weights = weights if weights else [1.0]*4
# TODO: make this actually pull values based on timestep instead of first value
self.timestep_keyframes = timestep_keyframes if timestep_keyframes else TimestepKeyframeGroup()
first_weight = self.timestep_keyframes.keyframes[0].t2i_adapter_weights if self.timestep_keyframes.get_index(0) else None
self.weights = first_weight if first_weight else [1.0]*4
self.channels_in = channels_in
self.control_input = None
@@ -275,12 +380,12 @@ class T2IAdapterAdvanced(ControlBase):
return out
def copy(self):
c = T2IAdapterAdvanced(self.t2i_model, self.weights, self.channels_in)
c = T2IAdapterAdvanced(self.t2i_model, self.timestep_keyframes, self.channels_in)
self.copy_to(c)
return c
def load_t2i_adapter(t2i_data, weights: T2IAdapterWeightsType=None):
def load_t2i_adapter(t2i_data, timestep_keyframes: TimestepKeyframeGroup=None):
keys = t2i_data.keys()
if 'adapter' in keys:
t2i_data = t2i_data['adapter']
@@ -300,4 +405,4 @@ def load_t2i_adapter(t2i_data, weights: T2IAdapterWeightsType=None):
else:
return None
model_ad.load_state_dict(t2i_data)
return T2IAdapterAdvanced(model_ad, weights, cin // 64)
return T2IAdapterAdvanced(model_ad, timestep_keyframes, cin // 64)
+195 -14
View File
@@ -6,7 +6,8 @@ import folder_paths
sys.path.insert(0, os.path.join(os.path.dirname(os.path.realpath(__file__)), "comfy"))
from comfy.sd import ControlBase
from .control import load_controlnet, ControlNetWeightsType, T2IAdapterWeightsType
from .control import load_controlnet, ControlNetWeightsType, T2IAdapterWeightsType,\
LatentKeyframe, LatentKeyframeGroup, TimestepKeyframe, TimestepKeyframeGroup
class ScaledSoftControlNetWeights:
@@ -22,7 +23,7 @@ class ScaledSoftControlNetWeights:
RETURN_TYPES = ("CONTROL_NET_WEIGHTS", )
FUNCTION = "load_weights"
CATEGORY = "controlnet weights"
CATEGORY = "adv-controlnet/weights"
def load_weights(self, base_multiplier, flip_weights):
weights = [(base_multiplier ** float(12 - i)) for i in range(13)]
@@ -56,7 +57,7 @@ class SoftControlNetWeights:
RETURN_TYPES = ("CONTROL_NET_WEIGHTS", )
FUNCTION = "load_weights"
CATEGORY = "controlnet weights"
CATEGORY = "adv-controlnet/weights"
def load_weights(self, weight_00, weight_01, weight_02, weight_03, weight_04, weight_05, weight_06,
weight_07, weight_08, weight_09, weight_10, weight_11, weight_12, flip_weights):
@@ -92,7 +93,7 @@ class CustomControlNetWeights:
RETURN_TYPES = ("CONTROL_NET_WEIGHTS", )
FUNCTION = "load_weights"
CATEGORY = "controlnet weights"
CATEGORY = "adv-controlnet/weights"
def load_weights(self, weight_00, weight_01, weight_02, weight_03, weight_04, weight_05, weight_06,
weight_07, weight_08, weight_09, weight_10, weight_11, weight_12, flip_weights):
@@ -119,7 +120,7 @@ class SoftT2IAdapterWeights:
RETURN_TYPES = ("T2I_ADAPTER_WEIGHTS", )
FUNCTION = "load_weights"
CATEGORY = "controlnet weights"
CATEGORY = "adv-controlnet/weights"
def load_weights(self, weight_00, weight_01, weight_02, weight_03, flip_weights):
weights = [weight_00, weight_01, weight_02, weight_03]
@@ -144,7 +145,7 @@ class CustomT2IAdapterWeights:
RETURN_TYPES = ("T2I_ADAPTER_WEIGHTS", )
FUNCTION = "load_weights"
CATEGORY = "controlnet weights"
CATEGORY = "adv-controlnet/weights"
def load_weights(self, weight_00, weight_01, weight_02, weight_03, flip_weights):
weights = [weight_00, weight_01, weight_02, weight_03]
@@ -153,6 +154,68 @@ class CustomT2IAdapterWeights:
return (weights,)
class TimestepKeyframeNode:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}, ),
},
"optional": {
"control_net_weights": ("CONTROL_NET_WEIGHTS", ),
"t2i_adapter_weights": ("T2I_ADAPTER_WEIGHTS", ),
"latent_keyframe": ("LATENT_KEYFRAME", ),
"prev_timestep_keyframe": ("TIMESTEP_KEYFRAME", ),
}
}
RETURN_TYPES = ("TIMESTEP_KEYFRAME", )
FUNCTION = "load_keyframe"
CATEGORY = "adv-controlnet/keyframes"
def load_keyframe(self,
start_percent: float,
control_net_weights: ControlNetWeightsType=None,
t2i_adapter_weights: T2IAdapterWeightsType=None,
latent_keyframe: LatentKeyframeGroup=None,
prev_timestep_keyframe: TimestepKeyframeGroup=None):
if not prev_timestep_keyframe:
prev_timestep_keyframe = TimestepKeyframeGroup()
keyframe = TimestepKeyframe(start_percent, control_net_weights, t2i_adapter_weights, latent_keyframe)
prev_timestep_keyframe.add(keyframe)
return (prev_timestep_keyframe,)
class LatentKeyframeNode:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"batch_index": ("INT", {"default": 0, "min": -1000, "max": 1000, "step": 1}),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.00001}, ),
},
"optional": {
"prev_latent_keyframe": ("LATENT_KEYFRAME", ),
}
}
RETURN_TYPES = ("LATENT_KEYFRAME", )
FUNCTION = "load_keyframe"
CATEGORY = "adv-controlnet/keyframes"
def load_keyframe(self,
batch_index: int,
strength: float,
prev_latent_keyframe: LatentKeyframeGroup=None):
if not prev_latent_keyframe:
prev_latent_keyframe = LatentKeyframeGroup()
keyframe = LatentKeyframe(batch_index, strength)
prev_latent_keyframe.add(keyframe)
return (prev_latent_keyframe,)
class ControlNetLoaderAdvanced:
@classmethod
def INPUT_TYPES(s):
@@ -161,19 +224,18 @@ class ControlNetLoaderAdvanced:
"control_net_name": (folder_paths.get_filename_list("controlnet"), ),
},
"optional": {
"control_net_weights": ("CONTROL_NET_WEIGHTS", ),
"t2i_adapter_weights": ("T2I_ADAPTER_WEIGHTS", ),
"timestep_keyframe": ("TIMESTEP_KEYFRAME", ),
}
}
RETURN_TYPES = ("CONTROL_NET", )
FUNCTION = "load_controlnet"
CATEGORY = "loaders"
CATEGORY = "adv-controlnet/loaders"
def load_controlnet(self, control_net_name, control_net_weights: ControlNetWeightsType=None, t2i_adapter_weights: T2IAdapterWeightsType=None):
def load_controlnet(self, control_net_name, timestep_keyframe: TimestepKeyframeGroup=None):
controlnet_path = folder_paths.get_full_path("controlnet", control_net_name)
controlnet = load_controlnet(controlnet_path, control_net_weights, t2i_adapter_weights)
controlnet = load_controlnet(controlnet_path, timestep_keyframe)
return (controlnet,)
@@ -194,16 +256,130 @@ class DiffControlNetLoaderAdvanced:
RETURN_TYPES = ("CONTROL_NET", )
FUNCTION = "load_controlnet"
CATEGORY = "loaders"
CATEGORY = "adv-controlnet/loaders"
def load_controlnet(self, control_net_name, control_net_weights, t2i_adapter_weights, model):
def load_controlnet(self, control_net_name, timestep_keyframe: TimestepKeyframeGroup, model):
controlnet_path = folder_paths.get_full_path("controlnet", control_net_name)
controlnet = load_controlnet(controlnet_path, control_net_weights, t2i_adapter_weights, model)
controlnet = load_controlnet(controlnet_path, timestep_keyframe, model)
return (controlnet,)
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])
# NODE MAPPING
NODE_CLASS_MAPPINGS = {
# Keyframes
"TimestepKeyframe": TimestepKeyframeNode,
"LatentKeyframe": LatentKeyframeNode,
# Conditioning
# "ControlNetApplyPartialBatch": ControlNetApplyPartialBatch,
# Loaders
"ControlNetLoaderAdvanced": ControlNetLoaderAdvanced,
"DiffControlNetLoaderAdvanced": DiffControlNetLoaderAdvanced,
@@ -216,6 +392,11 @@ NODE_CLASS_MAPPINGS = {
}
NODE_DISPLAY_NAME_MAPPINGS = {
# Keyframes
"TimestepKeyframe": "Timestep Keyframe",
"LatentKeyframe": "Latent Keyframe",
# Conditioning
# "ControlNetApplyPartialBatch": "Apply ControlNet (Partial Batch)",
# Loaders
"ControlNetLoaderAdvanced": "Load ControlNet Model (Advanced)",
"DiffControlNetLoaderAdvanced": "Load ControlNet Model (diff Advanced)",