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
This commit is contained in:
+9
-10
@@ -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]]
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user