From dff552141b91db1a8273b2bbd20798cb7de0976e Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Wed, 18 Oct 2023 08:01:42 -0500 Subject: [PATCH 1/2] Added Apply Advanced ControlNet node that supports masks --- control/control.py | 91 +++++++++++++++++++++++++++++++++++++--------- control/nodes.py | 29 ++++++++------- 2 files changed, 89 insertions(+), 31 deletions(-) diff --git a/control/control.py b/control/control.py index 47303e2..362df23 100644 --- a/control/control.py +++ b/control/control.py @@ -1,9 +1,11 @@ from typing import Union +from torch import Tensor import torch import comfy.utils import comfy.controlnet as comfy_cn from comfy.controlnet import ControlNet, T2IAdapter, broadcast_image_to +from sample import prepare_mask ControlNetWeightsType = list[float] T2IAdapterWeightsType = list[float] @@ -50,11 +52,13 @@ class TimestepKeyframe: start_percent: float = 0.0, control_net_weights: ControlNetWeightsType = None, t2i_adapter_weights: T2IAdapterWeightsType = None, - latent_keyframes: LatentKeyframeGroup = None) -> None: + latent_keyframes: LatentKeyframeGroup = None, + default_latent_strength: float = 0.0) -> 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 + self.default_latent_strength = default_latent_strength @classmethod @@ -153,12 +157,14 @@ def control_merge_inject(self, control_input, control_output, control_prev, outp class ControlNetAdvanced(ControlNet): def __init__(self, control_model, timestep_keyframes: TimestepKeyframeGroup, global_average_pooling=False, device=None): super().__init__(control_model=control_model, global_average_pooling=global_average_pooling, device=device) + # initialize timestep_keyframes self.timestep_keyframes = timestep_keyframes if timestep_keyframes else TimestepKeyframeGroup() self.current_timestep_keyframe = self.timestep_keyframes.keyframes[0] # initialize weights 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 + self.mask_cond_hint_original = None + self.mask_cond_hint = None # actual index values self.sub_idxs = None self.full_latent_length = 0 @@ -167,7 +173,7 @@ class ControlNetAdvanced(ControlNet): self.control_merge = control_merge_inject.__get__(self, type(self)) def set_cond_hint_mask(self, mask_hint): - self.cond_hint_mask = mask_hint + self.mask_cond_hint_original = mask_hint return self def get_control(self, x_noisy, t, cond, batched_number): @@ -175,13 +181,14 @@ class ControlNetAdvanced(ControlNet): self.t = t self.batched_number = batched_number # TODO: choose TimestepKeyframe based on t + return self.sliding_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): + def sliding_get_control(self, x_noisy: Tensor, 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) @@ -195,8 +202,9 @@ class ControlNetAdvanced(ControlNet): output_dtype = x_noisy.dtype + # make cond_hint appropriate dimensions # 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.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 @@ -204,10 +212,27 @@ class ControlNetAdvanced(ControlNet): # 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) + # make mask appropriate dimensions, if present + if self.mask_cond_hint_original is not None: + if self.sub_idxs is not None or self.mask_cond_hint is None or x_noisy.shape[2] * 8 != self.mask_cond_hint.shape[1] or x_noisy.shape[3] * 8 != self.mask_cond_hint.shape[2]: + if self.mask_cond_hint is not None: + del self.mask_cond_hint + self.mask_cond_hint = None + # resize mask and match batch count + self.mask_cond_hint = prepare_mask_batch(self.mask_cond_hint_original, x_noisy.shape, multiplier=8) + actual_latent_length = x_noisy.shape[0] // batched_number + self.mask_cond_hint = comfy.utils.repeat_to_batch_size(self.mask_cond_hint, actual_latent_length if self.sub_idxs is None else self.full_latent_length) + if self.sub_idxs is not None: + self.mask_cond_hint = self.mask_cond_hint[self.sub_idxs] + # make cond_hint_mask length match x_noise + if x_noisy.shape[0] != self.mask_cond_hint.shape[0]: + self.mask_cond_hint = broadcast_image_to(self.mask_cond_hint, x_noisy.shape[0], batched_number) + self.mask_cond_hint = self.mask_cond_hint.to(self.control_model.dtype).to(self.device) + #self.cond_hint = self.cond_hint * self.mask_cond_hint + context = cond['c_crossattn'] y = cond.get('c_adm', None) if y is not None: @@ -215,13 +240,12 @@ class ControlNetAdvanced(ControlNet): 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): + def apply_advanced_strengths_and_masks(self, x: Tensor, current_timestep_keyframe: TimestepKeyframe, batched_number: int): + # apply strengths, and get batch indeces to default out + # AKA latents that should not be influenced by ControlNet if current_timestep_keyframe.latent_keyframes is not None: - # apply strengths, and get batch indeces to zero out - # AKA latents that should not be influenced by ControlNet latent_count = x.size(0)//batched_number - - indeces_to_zero = set(range(latent_count)) + indeces_to_default = set(range(latent_count)) mapped_indeces = None # if expecting subdivision, will need to translate between subset and actual idx values if self.sub_idxs: @@ -236,24 +260,31 @@ class ControlNetAdvanced(ControlNet): # 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(real_index) + if real_index in indeces_to_default: + indeces_to_default.remove(real_index) # otherwise, see if batch_index is even included in this set of latents else: real_index = mapped_indeces.get(real_index, None) if real_index is None: continue - indeces_to_zero.remove(real_index) + indeces_to_default.remove(real_index) # 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 - # zero them out by multiplying by zero - for batch_index in indeces_to_zero: - # apply zero for each batched cond/uncond + # default them out by multiplying by default_latent_strength + for batch_index in indeces_to_default: + # apply default for each batched cond/uncond for b in range(batched_number): - x[(latent_count*b)+batch_index] = 0.0 + x[(latent_count*b)+batch_index] = x[(latent_count*b)+batch_index] * current_timestep_keyframe.default_latent_strength + # apply masks + if self.mask_cond_hint is not None: + # first, resize mask to required dims + masks = prepare_mask_batch(self.mask_cond_hint, x.shape) + x[:] = x[:] * masks + + def copy(self): c = ControlNetAdvanced(self.control_model, self.timestep_keyframes, global_average_pooling=self.global_average_pooling) @@ -328,3 +359,27 @@ def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroup=None, mo return T2IAdapterAdvanced(control.t2i_model, timestep_keyframe, control.channels_in) # otherwise, leave it be - probably a ControlLora for SDXL (no support for advanced stuff yet from here) return control + + +def is_advanced_controlnet(input_object): + return isinstance(input_object, ControlNetAdvanced) or isinstance(input_object, T2IAdapterAdvanced) + + +# adapted from comfy/sample.py +def prepare_mask_batch(mask: Tensor, shape: Tensor, multiplier: int=1, match_dim1=False): + mask = mask.clone() + mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), size=(shape[2]*multiplier, shape[3]*multiplier), mode="bilinear") + #mask = comfy.utils.repeat_to_batch_size(mask, shape[0]) + #noise_mask = noise_mask.round() + if match_dim1: + mask = torch.cat([mask] * shape[1], dim=1) + #noise_mask = torch.cat([noise_mask] * shape[1], dim=1) + #noise_mask = noise_mask.to(device) + return mask + +def prepare_mask_batch_old(mask, shape, device): + mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), size=(shape[2], shape[3]), mode="bilinear") + mask = mask.round() + mask = torch.cat([mask] * shape[1], dim=1) + mask = mask.to(device) + return mask diff --git a/control/nodes.py b/control/nodes.py index 36289e3..b43f78f 100644 --- a/control/nodes.py +++ b/control/nodes.py @@ -3,7 +3,7 @@ import numpy as np import folder_paths from .control import ControlNetAdvanced, T2IAdapterAdvanced, load_controlnet, ControlNetWeightsType, T2IAdapterWeightsType,\ - LatentKeyframeGroup, TimestepKeyframe, TimestepKeyframeGroup + LatentKeyframeGroup, TimestepKeyframe, TimestepKeyframeGroup, is_advanced_controlnet from .weight_nodes import ScaledSoftControlNetWeights, SoftControlNetWeights, CustomControlNetWeights, \ SoftT2IAdapterWeights, CustomT2IAdapterWeights from .latent_keyframe_nodes import LatentKeyframeGroupNode, LatentKeyframeInterpolationNode, LatentKeyframeBatchedGroupNode, LatentKeyframeNode @@ -91,7 +91,7 @@ class DiffControlNetLoaderAdvanced: return (controlnet,) -class ControlNetApplyAdvanced_AdvControlNet: +class AdvancedControlNetApply: @classmethod def INPUT_TYPES(s): return { @@ -105,7 +105,7 @@ class ControlNetApplyAdvanced_AdvControlNet: "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}) }, "optional": { - "mask_opt": ("MASK", ), + "mask_optional": ("MASK", ), } } @@ -113,14 +113,12 @@ class ControlNetApplyAdvanced_AdvControlNet: RETURN_NAMES = ("positive", "negative") FUNCTION = "apply_controlnet" - CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/loaders/conditioning" + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/conditioning" - def apply_controlnet(self, positive, negative, control_net, image, strength, start_percent, end_percent, mask_opt=None): + def apply_controlnet(self, positive, negative, control_net, image, strength, start_percent, end_percent, mask_optional=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 = {} @@ -135,12 +133,13 @@ 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 + # set cond hint mask + if mask_optional is not None: + if is_advanced_controlnet(c_net): + # if not in the form of a batch, make it so + if len(mask_optional.shape) < 3: + mask_optional = mask_optional.unsqueeze(0) + c_net.set_cond_hint_mask(mask_optional) c_net.set_previous_controlnet(prev_cnet) cnets[prev_cnet] = c_net @@ -163,6 +162,8 @@ NODE_CLASS_MAPPINGS = { # Loaders "ControlNetLoaderAdvanced": ControlNetLoaderAdvanced, "DiffControlNetLoaderAdvanced": DiffControlNetLoaderAdvanced, + # Conditioning + "ACN_AdvancedControlNetApply": AdvancedControlNetApply, # Weights "ScaledSoftControlNetWeights": ScaledSoftControlNetWeights, "SoftControlNetWeights": SoftControlNetWeights, @@ -183,6 +184,8 @@ NODE_DISPLAY_NAME_MAPPINGS = { # Loaders "ControlNetLoaderAdvanced": "Load ControlNet Model (Advanced) 🛂🅐🅒🅝", "DiffControlNetLoaderAdvanced": "Load ControlNet Model (diff Advanced) 🛂🅐🅒🅝", + # Conditioning + "ACN_AdvancedControlNetApply": "Apply Advanced ControlNet 🛂🅐🅒🅝", # Weights "ScaledSoftControlNetWeights": "Scaled Soft ControlNet Weights 🛂🅐🅒🅝", "SoftControlNetWeights": "Soft ControlNet Weights 🛂🅐🅒🅝", From fd4c2154263c86cce8e7707e9cd6e987ad1d7c27 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Wed, 18 Oct 2023 09:00:48 -0500 Subject: [PATCH 2/2] Cleaned up code for ControlNet masking --- control/control.py | 26 +++++--------------------- 1 file changed, 5 insertions(+), 21 deletions(-) diff --git a/control/control.py b/control/control.py index 362df23..0508ae4 100644 --- a/control/control.py +++ b/control/control.py @@ -5,7 +5,7 @@ import torch import comfy.utils import comfy.controlnet as comfy_cn from comfy.controlnet import ControlNet, T2IAdapter, broadcast_image_to -from sample import prepare_mask + ControlNetWeightsType = list[float] T2IAdapterWeightsType = list[float] @@ -181,12 +181,8 @@ class ControlNetAdvanced(ControlNet): self.t = t self.batched_number = batched_number # TODO: choose TimestepKeyframe based on t + # perform special version of get_control that supports sliding context and masks return self.sliding_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: Tensor, t, cond, batched_number): control_prev = None @@ -203,7 +199,7 @@ class ControlNetAdvanced(ControlNet): output_dtype = x_noisy.dtype # make cond_hint appropriate dimensions - # TODO: change this to not require cond_hint upscaling every step + # TODO: change this to not require cond_hint upscaling every step when self.sub_idxs are present 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 @@ -231,7 +227,6 @@ class ControlNetAdvanced(ControlNet): if x_noisy.shape[0] != self.mask_cond_hint.shape[0]: self.mask_cond_hint = broadcast_image_to(self.mask_cond_hint, x_noisy.shape[0], batched_number) self.mask_cond_hint = self.mask_cond_hint.to(self.control_model.dtype).to(self.device) - #self.cond_hint = self.cond_hint * self.mask_cond_hint context = cond['c_crossattn'] y = cond.get('c_adm', None) @@ -284,8 +279,6 @@ class ControlNetAdvanced(ControlNet): masks = prepare_mask_batch(self.mask_cond_hint, x.shape) x[:] = x[:] * masks - - def copy(self): c = ControlNetAdvanced(self.control_model, self.timestep_keyframes, global_average_pooling=self.global_average_pooling) self.copy_to(c) @@ -335,6 +328,7 @@ class T2IAdapterAdvanced(T2IAdapter): 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 + # TODO: support masks return def copy(self): @@ -358,6 +352,7 @@ def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroup=None, mo elif isinstance(control, T2IAdapter): return T2IAdapterAdvanced(control.t2i_model, timestep_keyframe, control.channels_in) # otherwise, leave it be - probably a ControlLora for SDXL (no support for advanced stuff yet from here) + # TODO add ControlLoraAdvanced return control @@ -369,17 +364,6 @@ def is_advanced_controlnet(input_object): def prepare_mask_batch(mask: Tensor, shape: Tensor, multiplier: int=1, match_dim1=False): mask = mask.clone() mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), size=(shape[2]*multiplier, shape[3]*multiplier), mode="bilinear") - #mask = comfy.utils.repeat_to_batch_size(mask, shape[0]) - #noise_mask = noise_mask.round() if match_dim1: mask = torch.cat([mask] * shape[1], dim=1) - #noise_mask = torch.cat([noise_mask] * shape[1], dim=1) - #noise_mask = noise_mask.to(device) - return mask - -def prepare_mask_batch_old(mask, shape, device): - mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), size=(shape[2], shape[3]), mode="bilinear") - mask = mask.round() - mask = torch.cat([mask] * shape[1], dim=1) - mask = mask.to(device) return mask