Make T2IAdapter work with uncond_multiplier (results are not great when < 1.0, but match up exactly with auto1111 results)

This commit is contained in:
Jedrzej Kosinski
2024-05-16 17:40:46 -05:00
parent db0bf14dac
commit 4ed14e9ae0
2 changed files with 20 additions and 2 deletions
+17
View File
@@ -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]]
+3 -2
View File
@@ -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