From 96d957e9ec99ec365d6056c61058e286f6b3c703 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Thu, 27 Jun 2024 17:43:21 -0500 Subject: [PATCH] Fixed get_calc_pow to work the same as before (only current difference would be with SDXL T2IAdapter models with custom/soft weights), automatically resize T2IAdapter control tensors to match batch_size to make my life easier --- adv_control/control.py | 19 +++++++++---------- adv_control/utils.py | 13 ++++++++----- 2 files changed, 17 insertions(+), 15 deletions(-) diff --git a/adv_control/control.py b/adv_control/control.py index ac1efc5..bf2429f 100644 --- a/adv_control/control.py +++ b/adv_control/control.py @@ -101,15 +101,14 @@ class T2IAdapterAdvanced(T2IAdapter, AdvancedControlBase): AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.t2iadapter()) def control_merge_inject(self, control: dict[str, list[Tensor]], control_prev, output_dtype): - # if has uncond multiplier, need to make sure control shapes are the same batch size as expected - if self.weights.has_uncond_multiplier or self.weights.has_uncond_mask: - for key in control: - control_current = control[key] - for i in range(len(control_current)): - x = control_current[i] - if x is not None: - if x.size(0) < self.batch_size: - control_current[i] = x.repeat(self.batched_number, 1, 1, 1)[:self.batch_size] + # match batch_size + # TODO: make this more efficient by modifying the cached self.control_input val instead of doing this every step + for key in control: + control_current = control[key] + for i in range(len(control_current)): + x = control_current[i] + if x is not None and x.size(0) == 1 and x.size(0) != self.batch_size: + control_current[i] = x.repeat(self.batch_size, 1, 1, 1)[:self.batch_size] return AdvancedControlBase.control_merge_inject(self, control, control_prev, output_dtype) def get_universal_weights(self) -> ControlWeights: @@ -119,7 +118,7 @@ class T2IAdapterAdvanced(T2IAdapter, AdvancedControlBase): raw_weights.reverse() # need to reverse to match recent ComfyUI changes return self.weights.copy_with_new_weights(raw_weights) - def get_calc_pow(self, idx: int, layers: int) -> int: + def get_calc_pow(self, idx: int, control: dict[str, list[Tensor]], key: str) -> int: # match how T2IAdapterAdvanced deals with universal weights indeces = [7 - i for i in range(8)] indeces = [indeces[-8], indeces[-3], indeces[-2], indeces[-1]] diff --git a/adv_control/utils.py b/adv_control/utils.py index 98c8676..f16ac35 100644 --- a/adv_control/utils.py +++ b/adv_control/utils.py @@ -753,7 +753,11 @@ class AdvancedControlBase: return self.weights.get(idx=idx, control=control, key=key) def get_calc_pow(self, idx: int, control: dict[str, list[Tensor]], key: str) -> int: - return (len(control[key])-1)-idx + c_len = len(control[key])-1 + if key == "output": + if "middle" in control: + c_len += len(control["middle"]) + return c_len-idx def calc_latent_keyframe_mults(self, x: Tensor, batched_number: int) -> Tensor: # apply strengths, and get batch indeces to null out @@ -815,15 +819,14 @@ class AdvancedControlBase: if self.weights.has_uncond_mask: pass - x_len = x.size(0) # mainly to account for how ComfyUI T2IAdapter works when only one condhint is provided if self.latent_keyframes is not None: - x[:] = x[:] * self.calc_latent_keyframe_mults(x=x, batched_number=batched_number)[:x_len] + x[:] = x[:] * self.calc_latent_keyframe_mults(x=x, batched_number=batched_number) # apply masks, resizing mask to required dims if self.mask_cond_hint is not None: - masks = prepare_mask_batch(self.mask_cond_hint, x.shape)[:x_len] + 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_len] + 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: