From 4ed14e9ae05ac59dbe5b5a28db56848019a5ea5c Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Thu, 16 May 2024 17:40:46 -0500 Subject: [PATCH] Make T2IAdapter work with uncond_multiplier (results are not great when < 1.0, but match up exactly with auto1111 results) --- adv_control/control.py | 17 +++++++++++++++++ adv_control/utils.py | 5 +++-- 2 files changed, 20 insertions(+), 2 deletions(-) diff --git a/adv_control/control.py b/adv_control/control.py index 2ea2586..bea5aa7 100644 --- a/adv_control/control.py +++ b/adv_control/control.py @@ -95,6 +95,23 @@ class T2IAdapterAdvanced(T2IAdapter, AdvancedControlBase): super().__init__(t2i_model=t2i_model, channels_in=channels_in, compression_ratio=compression_ratio, upscale_algorithm=upscale_algorithm, device=device) AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.t2iadapter()) + def control_merge_inject(self, control_input, control_output, 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: + if control_input is not None: + for i in range(len(control_input)): + x = control_input[i] + if x is not None: + if x.size(0) < self.batch_size: + control_input[i] = x.repeat(self.batched_number, 1, 1, 1)[:self.batch_size] + if control_output is not None: + for i in range(len(control_output)): + x = control_output[i] + if x is not None: + if x.size(0) < self.batch_size: + control_output[i] = x.repeat(self.batched_number, 1, 1, 1)[:self.batch_size] + return AdvancedControlBase.control_merge_inject(self, control_input, control_output, control_prev, output_dtype) + def get_universal_weights(self) -> ControlWeights: raw_weights = [(self.weights.base_multiplier ** float(7 - i)) for i in range(8)] raw_weights = [raw_weights[-8], raw_weights[-3], raw_weights[-2], raw_weights[-1]] diff --git a/adv_control/utils.py b/adv_control/utils.py index 4257b79..c6c13db 100644 --- a/adv_control/utils.py +++ b/adv_control/utils.py @@ -520,6 +520,7 @@ class AdvancedControlBase: # timesteps self.t: Tensor = None self.batched_number: int = None + self.batch_size: int = 0 # weights + override self.weights: ControlWeights = None self.weights_default: ControlWeights = weights_default @@ -573,6 +574,7 @@ class AdvancedControlBase: def prepare_current_timestep(self, t: Tensor, batched_number: int): self.t = float(t[0]) self.batched_number = batched_number + self.batch_size = len(t) # get current step percent curr_t: float = self.t prev_index = self._current_timestep_index @@ -667,8 +669,6 @@ class AdvancedControlBase: return True def get_control_inject(self, x_noisy, t, cond, batched_number): - if type(batched_number) != IntWithCondOrUncond: - logger.warn(f"not IntWithCondOrUncond! {type(batched_number)}") # prepare timestep and everything related self.prepare_current_timestep(t=t, batched_number=batched_number) # if should not perform any actions for the controlnet, exit without doing any work @@ -869,6 +869,7 @@ class AdvancedControlBase: self.context_length = 0 self.t = None self.batched_number = None + self.batch_size = 0 self.weights = None self.latent_keyframes = None # timestep stuff