diff --git a/control.py b/control.py index 8fd8eaa..16a8e8b 100644 --- a/control.py +++ b/control.py @@ -124,8 +124,11 @@ class ControlNetAdvanced(ControlBase): 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 + # mask for which parts of controlnet output to keep + self.cond_hint_mask = None def get_control(self, x_noisy, t, cond, batched_number): + #print(f"$$$$ t len: f{len(t)}") control_prev = None if self.previous_controlnet is not None: control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number) @@ -155,31 +158,22 @@ class ControlNetAdvanced(ControlBase): 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, current_timestep_keyframe) + return self.control_merge(None, control, control_prev, output_dtype, current_timestep_keyframe, batched_number) - def control_merge(self, control_input, control_output, control_prev, output_dtype, current_timestep_keyframe: TimestepKeyframe): + def control_merge(self, control_input, control_output, control_prev, output_dtype, current_timestep_keyframe: TimestepKeyframe, batched_number: int): out = {'input':[], 'middle':[], 'output': []} + #print(f"$$$$ control_input size: {control_input.size() if control_input != None else None}") + #print(f"$$$$ control_output size: {control_output.size() if control_input != None else None}") + if control_input is not None: for i in range(len(control_input)): key = 'input' x = control_input[i] - 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 - 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) - # apply strength - x[keyframe.batch_index] *= keyframe.strength - x[(x.size()[0]//2) + keyframe.batch_index] *= keyframe.strength + #print(f"$$$$ x size: {x.size()}") - # 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 + self.apply_advanced_strengths_and_masks(x, current_timestep_keyframe, batched_number) if x is not None: x *= self.strength * self.weights[i] @@ -197,21 +191,9 @@ class ControlNetAdvanced(ControlBase): index = i x = control_output[i] - 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 - 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) - # apply strength - x[keyframe.batch_index] *= keyframe.strength - x[(x.size()[0]//2) + keyframe.batch_index] *= keyframe.strength + #print(f"$$$$ x size: {x.size()}") - # zero out indeces that should not be affected 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 + self.apply_advanced_strengths_and_masks(x, current_timestep_keyframe, batched_number) if x is not None: if self.global_average_pooling: @@ -236,6 +218,27 @@ class ControlNetAdvanced(ControlBase): o[i] += prev_val return out + def apply_advanced_strengths_and_masks(self, x, current_timestep_keyframe: TimestepKeyframe, batched_number: int): + 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)) + for keyframe in current_timestep_keyframe.latent_keyframes: + if keyframe.batch_index in indeces_to_zero: + indeces_to_zero.remove(keyframe.batch_index) + # apply strength for each batched cond/uncond + for b in range(batched_number): + x[(latent_count*b)+keyframe.batch_index] *= keyframe.strength + + # zero them out by multiplying by zero + for batch_index in indeces_to_zero: + # apply zero for each batched cond/uncond + for b in range(batched_number): + x[(latent_count*b)+batch_index] *= 0.0 + + def copy(self): c = ControlNetAdvanced(self.control_model, self.timestep_keyframes, global_average_pooling=self.global_average_pooling) self.copy_to(c) diff --git a/nodes.py b/nodes.py index e0b22ee..4b1f18d 100644 --- a/nodes.py +++ b/nodes.py @@ -269,6 +269,57 @@ class DiffControlNetLoaderAdvanced: return (controlnet,) +class ControlNetApplyAdvanced_AdvControlNet: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "positive": ("CONDITIONING", ), + "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}) + } + } + + RETURN_TYPES = ("CONDITIONING","CONDITIONING") + RETURN_NAMES = ("positive", "negative") + FUNCTION = "apply_controlnet" + + CATEGORY = "adv-controlnet/loaders/conditioning" + + def apply_controlnet(self, positive, negative, control_net, image, mask_opt, strength, start_percent, end_percent): + if strength == 0: + return (positive, negative) + + control_hint = image.movedim(-1,1) + cnets = {} + + out = [] + for conditioning in [positive, negative]: + c = [] + 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) + out.append(c) + 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):