From 1e2d9a54916d19020b3ac744ff29473e5da9f1e5 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Wed, 29 Nov 2023 03:36:36 -0600 Subject: [PATCH] Added masks to TimestepKeyframes, fixed duplication bug between runs with LatentKeyframes and TimestepKeyframes --- control/control.py | 81 ++++++++++++++++++++++++++++++-- control/latent_keyframe_nodes.py | 8 ++++ control/nodes.py | 7 ++- 3 files changed, 92 insertions(+), 4 deletions(-) diff --git a/control/control.py b/control/control.py index f70900d..180c9e9 100644 --- a/control/control.py +++ b/control/control.py @@ -122,6 +122,12 @@ class LatentKeyframeGroup: def is_empty(self) -> bool: return len(self.keyframes) == 0 + def clone(self) -> 'LatentKeyframeGroup': + cloned = LatentKeyframeGroup() + for tk in self.keyframes: + cloned.add(tk) + return cloned + class TimestepKeyframe: def __init__(self, @@ -132,7 +138,8 @@ class TimestepKeyframe: latent_keyframes: LatentKeyframeGroup = None, null_latent_kf_strength: float = 0.0, inherit_missing: bool = True, - guarantee_usage: bool = False) -> None: + guarantee_usage: bool = True, + mask_hint_orig: Tensor = None) -> None: self.start_percent = start_percent self.start_t = 999999999.9 self.strength = strength @@ -142,6 +149,7 @@ class TimestepKeyframe: self.null_latent_kf_strength = null_latent_kf_strength self.inherit_missing = inherit_missing self.guarantee_usage = guarantee_usage + self.mask_hint_orig = mask_hint_orig def has_control_weights(self): return self.control_weights is not None @@ -149,6 +157,9 @@ class TimestepKeyframe: def has_latent_keyframes(self): return self.latent_keyframes is not None + def has_mask_hint(self): + return self.mask_hint_orig is not None + @classmethod def default(cls) -> 'TimestepKeyframe': @@ -191,6 +202,12 @@ class TimestepKeyframeGroup: def is_empty(self) -> bool: return len(self.keyframes) == 0 + def clone(self) -> 'TimestepKeyframeGroup': + cloned = TimestepKeyframeGroup() + for tk in self.keyframes: + cloned.add(tk) + return cloned + @classmethod def default(cls, keyframe: TimestepKeyframe) -> 'TimestepKeyframeGroup': group = cls() @@ -209,6 +226,9 @@ class AdvancedControlBase: # mask for which parts of controlnet output to keep self.mask_cond_hint_original = None self.mask_cond_hint = None + self.tk_mask_cond_hint_original = None + self.tk_mask_cond_hint = None + self.weight_mask_cond_hint = None # actual index values self.sub_idxs = None self.full_latent_length = 0 @@ -267,6 +287,11 @@ class AdvancedControlBase: self.latent_keyframes = self.current_timestep_keyframe.latent_keyframes elif not self.current_timestep_keyframe.inherit_missing: self.latent_keyframes = None + if self.current_timestep_keyframe.has_mask_hint(): + self.tk_mask_cond_hint_original = self.current_timestep_keyframe.mask_hint_orig + elif not self.current_timestep_keyframe.inherit_missing: + del self.tk_mask_cond_hint_original + self.tk_mask_cond_hint_original = None # if guarantee_usage, stop searching for other TKs if self.current_timestep_keyframe.guarantee_usage: break @@ -365,11 +390,13 @@ class AdvancedControlBase: # 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 - # apply masks + # apply masks, resizing mask to required dims 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 + if self.tk_mask_cond_hint is not None: + masks = prepare_mask_batch(self.tk_mask_cond_hint, x.shape) + x[:] = x[:] * masks # apply timestep keyframe strengths if self.current_timestep_keyframe.strength != 1.0: x[:] *= self.current_timestep_keyframe.strength @@ -444,6 +471,41 @@ class AdvancedControlBase: if dtype is None: dtype = x_noisy.dtype self.mask_cond_hint = self.mask_cond_hint.to(dtype=dtype).to(self.device) + # prepare other masks + self.prepare_tk_mask_cond_hint(x_noisy, t, cond, batched_number, dtype) + + def prepare_tk_mask_cond_hint(self, x_noisy: Tensor, t, cond, batched_number, dtype=None): + return self._prepare_mask("tk_mask_cond_hint", self.current_timestep_keyframe.mask_hint_orig, x_noisy, t, cond, batched_number, dtype) + + def prepare_weight_mask_cond_hint(self, x_noisy: Tensor, t, cond, batched_number, dtype=None): + return self._prepare_mask("weight_mask_cond_hint", self.current_timestep_keyframe.mask_hint_orig, x_noisy, t, cond, batched_number, dtype) + + def _prepare_mask(self, attr_name, orig_mask: Tensor, x_noisy: Tensor, t, cond, batched_number, dtype=None): + if orig_mask is not None: + out_mask = getattr(self, attr_name) + if self.sub_idxs is not None or out_mask is None or x_noisy.shape[2] * 8 != out_mask.shape[1] or x_noisy.shape[3] * 8 != out_mask.shape[2]: + self._reset_attr(attr_name) + del out_mask + # TODO: perform upscale on only the sub_idxs masks at a time instead of all to conserve RAM + # resize mask and match batch count + out_mask = prepare_mask_batch(orig_mask, x_noisy.shape, multiplier=8) + actual_latent_length = x_noisy.shape[0] // batched_number + out_mask = comfy.utils.repeat_to_batch_size(out_mask, actual_latent_length if self.sub_idxs is None else self.full_latent_length) + if self.sub_idxs is not None: + out_mask = out_mask[self.sub_idxs] + # make cond_hint_mask length match x_noise + if x_noisy.shape[0] != out_mask.shape[0]: + out_mask = broadcast_image_to(out_mask, x_noisy.shape[0], batched_number) + # default dtype to be same as x_noisy + if dtype is None: + dtype = x_noisy.dtype + setattr(self, attr_name, out_mask.to(dtype=dtype).to(self.device)) + del out_mask + + def _reset_attr(self, attr_name, new_value=None): + if hasattr(self, attr_name): + delattr(self, attr_name) + setattr(self, attr_name, new_value) def cleanup_inject(self): self.base.cleanup() @@ -461,6 +523,19 @@ class AdvancedControlBase: self.current_timestep_keyframe = None self.next_timestep_keyframe = None self.current_timestep_index = -1 + # clear mask hints + if self.mask_cond_hint is not None: + del self.mask_cond_hint + self.mask_cond_hint = None + if self.tk_mask_cond_hint_original is not None: + del self.tk_mask_cond_hint_original + self.tk_mask_cond_hint_original = None + if self.tk_mask_cond_hint is not None: + del self.tk_mask_cond_hint + self.tk_mask_cond_hint = None + if self.weight_mask_cond_hint is not None: + del self.weight_mask_cond_hint + self.weight_mask_cond_hint = None def copy_to_advanced(self, copied: 'AdvancedControlBase'): copied.mask_cond_hint_original = self.mask_cond_hint_original diff --git a/control/latent_keyframe_nodes.py b/control/latent_keyframe_nodes.py index 8152921..d9b5856 100644 --- a/control/latent_keyframe_nodes.py +++ b/control/latent_keyframe_nodes.py @@ -31,6 +31,8 @@ class LatentKeyframeNode: prev_latent_keyframe: LatentKeyframeGroup=None): if not prev_latent_keyframe: prev_latent_keyframe = LatentKeyframeGroup() + else: + prev_latent_keyframe = prev_latent_keyframe.clone() keyframe = LatentKeyframe(batch_index, strength) prev_latent_keyframe.add(keyframe) return (prev_latent_keyframe,) @@ -126,6 +128,8 @@ class LatentKeyframeGroupNode: print_keyframes=False): if not prev_latent_keyframe: prev_latent_keyframe = LatentKeyframeGroup() + else: + prev_latent_keyframe = prev_latent_keyframe.clone() curr_latent_keyframe = LatentKeyframeGroup() latent_count = -1 @@ -185,6 +189,8 @@ class LatentKeyframeInterpolationNode: if not prev_latent_keyframe: prev_latent_keyframe = LatentKeyframeGroup() + else: + prev_latent_keyframe = prev_latent_keyframe.clone() curr_latent_keyframe = LatentKeyframeGroup() steps = batch_index_to_excl - batch_index_from @@ -236,6 +242,8 @@ class LatentKeyframeBatchedGroupNode: def load_keyframe(self, float_strengths: Union[float, list[float]], prev_latent_keyframe: LatentKeyframeGroup=None, print_keyframes=False): if not prev_latent_keyframe: prev_latent_keyframe = LatentKeyframeGroup() + else: + prev_latent_keyframe = prev_latent_keyframe.clone() curr_latent_keyframe = LatentKeyframeGroup() # if received a normal float input, do nothing diff --git a/control/nodes.py b/control/nodes.py index e3803b3..3d03faf 100644 --- a/control/nodes.py +++ b/control/nodes.py @@ -27,6 +27,7 @@ class TimestepKeyframeNode: "null_latent_kf_strength": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), "inherit_missing": ("BOOLEAN", {"default": True}, ), "guarantee_usage": ("BOOLEAN", {"default": True}, ), + "mask_optional": ("MASK", ), #"interpolation": ([SI.LINEAR, SI.EASE_IN, SI.EASE_OUT, SI.EASE_IN_OUT, SI.NONE], {"default": SI.NONE}, ), } } @@ -46,11 +47,15 @@ class TimestepKeyframeNode: null_latent_kf_strength: float=0.0, inherit_missing=True, guarantee_usage=True, + mask_optional=None, interpolation: str=SI.NONE,): if not prev_timestep_keyframe: prev_timestep_keyframe = TimestepKeyframeGroup() + else: + prev_timestep_keyframe = prev_timestep_keyframe.clone() keyframe = TimestepKeyframe(start_percent=start_percent, strength=strength, interpolation=interpolation, null_latent_kf_strength=null_latent_kf_strength, - control_weights=control_net_weights, latent_keyframes=latent_keyframe, inherit_missing=inherit_missing, guarantee_usage=guarantee_usage) + control_weights=control_net_weights, latent_keyframes=latent_keyframe, inherit_missing=inherit_missing, guarantee_usage=guarantee_usage, + mask_hint_orig=mask_optional) prev_timestep_keyframe.add(keyframe) return (prev_timestep_keyframe,)