diff --git a/adv_control/control.py b/adv_control/control.py index b9253de..235bb02 100644 --- a/adv_control/control.py +++ b/adv_control/control.py @@ -8,14 +8,15 @@ import comfy.utils import comfy.model_management import comfy.model_detection import comfy.controlnet as comfy_cn -from comfy.controlnet import ControlBase, ControlNet, ControlLora, T2IAdapter, broadcast_image_to +from comfy.controlnet import ControlBase, ControlNet, ControlLora, T2IAdapter from comfy.model_patcher import ModelPatcher from .control_sparsectrl import SparseModelPatcher, SparseControlNet, SparseCtrlMotionWrapper, SparseMethod, SparseSettings, SparseSpreadMethod, PreprocSparseRGBWrapper from .control_lllite import LLLiteModule, LLLitePatch from .control_svd import svd_unet_config_from_diffusers_unet, SVDControlNet, svd_unet_to_diffusers from .utils import (AdvancedControlBase, TimestepKeyframeGroup, LatentKeyframeGroup, ControlWeightType, ControlWeights, WeightTypeException, - manual_cast_clean_groupnorm, disable_weight_init_clean_groupnorm, prepare_mask_batch, get_properly_arranged_t2i_weights, load_torch_file_with_dict_factory) + manual_cast_clean_groupnorm, disable_weight_init_clean_groupnorm, prepare_mask_batch, get_properly_arranged_t2i_weights, load_torch_file_with_dict_factory, + broadcast_image_to_extend, extend_to_batch_size) from .logger import logger @@ -56,12 +57,15 @@ class ControlNetAdvanced(ControlNet, AdvancedControlBase): del self.cond_hint self.cond_hint = None # if self.cond_hint_original length greater or equal to real latent count, subdivide it before scaling - if self.sub_idxs is not None and self.cond_hint_original.size(0) >= self.full_latent_length: - self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original[self.sub_idxs], x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(self.device) + if self.sub_idxs is not None: + actual_cond_hint_orig = self.cond_hint_original + if self.cond_hint_original.size(0) < self.full_latent_length: + actual_cond_hint_orig = extend_to_batch_size(tensor=actual_cond_hint_orig, batch_size=self.full_latent_length) + self.cond_hint = comfy.utils.common_upscale(actual_cond_hint_orig[self.sub_idxs], x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(self.device) else: self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original, x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(self.device) if x_noisy.shape[0] != self.cond_hint.shape[0]: - self.cond_hint = broadcast_image_to(self.cond_hint, x_noisy.shape[0], batched_number) + self.cond_hint = broadcast_image_to_extend(self.cond_hint, x_noisy.shape[0], batched_number) # prepare mask_cond_hint self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number, dtype=dtype) @@ -98,7 +102,7 @@ class T2IAdapterAdvanced(T2IAdapter, AdvancedControlBase): 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 self.weights.has_uncond_multiplier or self.weights.has_uncond_mask: if control_input is not None: for i in range(len(control_input)): x = control_input[i] @@ -132,9 +136,12 @@ class T2IAdapterAdvanced(T2IAdapter, AdvancedControlBase): if self.sub_idxs is not None: # cond hints full_cond_hint_original = self.cond_hint_original + actual_cond_hint_orig = full_cond_hint_original del self.cond_hint self.cond_hint = None - self.cond_hint_original = full_cond_hint_original[self.sub_idxs] + if full_cond_hint_original.size(0) < self.full_latent_length: + actual_cond_hint_orig = extend_to_batch_size(tensor=full_cond_hint_original, batch_size=full_cond_hint_original.size(0)) + self.cond_hint_original = actual_cond_hint_orig[self.sub_idxs] # mask hints self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number) return super().get_control(x_noisy, t, cond, batched_number) @@ -222,12 +229,15 @@ class SVDControlNetAdvanced(ControlNetAdvanced): del self.cond_hint self.cond_hint = None # if self.cond_hint_original length greater or equal to real latent count, subdivide it before scaling - if self.sub_idxs is not None and self.cond_hint_original.size(0) >= self.full_latent_length: - self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original[self.sub_idxs], x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(self.device) + if self.sub_idxs is not None: + actual_cond_hint_orig = self.cond_hint_original + if self.cond_hint_original.size(0) < self.full_latent_length: + actual_cond_hint_orig = extend_to_batch_size(tensor=actual_cond_hint_orig, batch_size=self.full_latent_length) + self.cond_hint = comfy.utils.common_upscale(actual_cond_hint_orig[self.sub_idxs], x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(self.device) else: self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original, x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(self.device) if x_noisy.shape[0] != self.cond_hint.shape[0]: - self.cond_hint = broadcast_image_to(self.cond_hint, x_noisy.shape[0], batched_number) + self.cond_hint = broadcast_image_to_extend(self.cond_hint, x_noisy.shape[0], batched_number) # prepare mask_cond_hint self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number, dtype=dtype) @@ -328,7 +338,7 @@ class SparseCtrlAdvanced(ControlNetAdvanced): del cond_mask # make cond_hint match x_noisy batch if x_noisy.shape[0] != self.cond_hint.shape[0]: - self.cond_hint = broadcast_image_to(self.cond_hint, x_noisy.shape[0], batched_number) + self.cond_hint = broadcast_image_to_extend(self.cond_hint, x_noisy.shape[0], batched_number) # prepare mask_cond_hint self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number, dtype=dtype) @@ -412,12 +422,15 @@ class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase): del self.cond_hint self.cond_hint = None # if self.cond_hint_original length greater or equal to real latent count, subdivide it before scaling - if self.sub_idxs is not None and self.cond_hint_original.size(0) >= self.full_latent_length: - self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original[self.sub_idxs], x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(self.device) + if self.sub_idxs is not None: + actual_cond_hint_orig = self.cond_hint_original + if self.cond_hint_original.size(0) < self.full_latent_length: + actual_cond_hint_orig = extend_to_batch_size(tensor=actual_cond_hint_orig, batch_size=self.full_latent_length) + self.cond_hint = comfy.utils.common_upscale(actual_cond_hint_orig[self.sub_idxs], x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(self.device) else: self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original, x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(self.device) if x_noisy.shape[0] != self.cond_hint.shape[0]: - self.cond_hint = broadcast_image_to(self.cond_hint, x_noisy.shape[0], batched_number) + self.cond_hint = broadcast_image_to_extend(self.cond_hint, x_noisy.shape[0], batched_number) # some special logic here compared to other controlnets: # * The cond_emb in attn patches will divide latent dims by 2 or 4, integer # * Due to this loss, the cond_emb will become smaller than x input if latent dims are not divisble by 2 or 4 diff --git a/adv_control/control_reference.py b/adv_control/control_reference.py index 815d651..24f1485 100644 --- a/adv_control/control_reference.py +++ b/adv_control/control_reference.py @@ -14,7 +14,7 @@ from comfy.ldm.modules.diffusionmodules import openaimodel from .logger import logger from .utils import (AdvancedControlBase, ControlWeights, TimestepKeyframeGroup, AbstractPreprocWrapper, - deepcopy_with_sharing, prepare_mask_batch, broadcast_image_to_full) + deepcopy_with_sharing, prepare_mask_batch, broadcast_image_to_extend) def refcn_sample_factory(orig_comfy_sample: Callable, is_custom=False) -> Callable: @@ -326,7 +326,7 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): self.cond_hint_original, x_noisy.shape[3], x_noisy.shape[2], 'nearest-exact', "center").to(dtype).to(self.device) if x_noisy.shape[0] != self.cond_hint.shape[0]: - self.cond_hint = broadcast_image_to_full(self.cond_hint, x_noisy.shape[0], batched_number, except_one=False) + self.cond_hint = broadcast_image_to_extend(self.cond_hint, x_noisy.shape[0], batched_number, except_one=False) # noise cond_hint based on sigma (current step) self.cond_hint = self.latent_format.process_in(self.cond_hint) self.cond_hint = ref_noise_latents(self.cond_hint, sigma=t, noise=None) diff --git a/adv_control/control_sparsectrl.py b/adv_control/control_sparsectrl.py index 58b4d85..c50c7c7 100644 --- a/adv_control/control_sparsectrl.py +++ b/adv_control/control_sparsectrl.py @@ -25,12 +25,11 @@ from comfy.ldm.modules.attention import attention_basic, attention_pytorch, atte from comfy.ldm.modules.attention import FeedForward, SpatialTransformer from comfy.ldm.modules.diffusionmodules.openaimodel import TimestepEmbedSequential, ResBlock, Downsample from comfy.model_patcher import ModelPatcher -from comfy.controlnet import broadcast_image_to -from comfy.utils import repeat_to_batch_size import comfy.ops import comfy.model_management -from .utils import TimestepKeyframeGroup, disable_weight_init_clean_groupnorm, prepare_mask_batch +from .utils import (TimestepKeyframeGroup, disable_weight_init_clean_groupnorm, + prepare_mask_batch, broadcast_image_to_extend, extend_to_batch_size) # until xformers bug is fixed, do not use xformers for VersatileAttention! TODO: change this when fix is out @@ -560,8 +559,8 @@ class VanillaTemporalModule(nn.Module): return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask) elif math.isclose(self.strength, 0.0): return input_tensor - elif self.strength > 1.0: - return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask)*self.strength + # elif self.strength > 1.0: + # return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask)*self.strength else: return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask)*self.strength + input_tensor*(1.0-self.strength) @@ -673,10 +672,10 @@ class TemporalTransformer3DModel(nn.Module): # otherwise, calculate temp mask self.prev_hidden_states_batch = batch mask = prepare_mask_batch(self.raw_scale_mask, shape=(self.full_length, 1, height, width)) - mask = repeat_to_batch_size(mask, self.full_length) + mask = extend_to_batch_size(mask, self.full_length) # if mask not the same amount length as full length, make it match if self.full_length != mask.shape[0]: - mask = broadcast_image_to(mask, self.full_length, 1) + mask = broadcast_image_to_extend(mask, self.full_length, 1) # reshape mask to attention K shape (h*w, latent_count, 1) batch, channel, height, width = mask.shape # first, perform same operations as on hidden_states, diff --git a/adv_control/utils.py b/adv_control/utils.py index 9e0beda..a7e17dd 100644 --- a/adv_control/utils.py +++ b/adv_control/utils.py @@ -12,7 +12,7 @@ import comfy.sample import comfy.samplers import comfy.model_base -from comfy.controlnet import ControlBase, broadcast_image_to +from comfy.controlnet import ControlBase from comfy.model_patcher import ModelPatcher from .logger import logger @@ -152,7 +152,7 @@ class ControlWeightType: class ControlWeights: def __init__(self, weight_type: str, base_multiplier: float=1.0, flip_weights: bool=False, weights: list[float]=None, weight_mask: Tensor=None, - uncond_multiplier=1.0): + uncond_multiplier=1.0, uncond_mask: Tensor=None): self.weight_type = weight_type self.base_multiplier = base_multiplier self.flip_weights = flip_weights @@ -162,6 +162,8 @@ class ControlWeights: self.weight_mask = weight_mask self.uncond_multiplier = float(uncond_multiplier) self.has_uncond_multiplier = not math.isclose(self.uncond_multiplier, 1.0) + self.uncond_mask = uncond_mask if uncond_mask is not None else 1.0 + self.has_uncond_mask = uncond_mask is not None def get(self, idx: int, default=1.0) -> Union[float, Tensor]: # if weights is not none, return index @@ -433,8 +435,15 @@ def normalize_min_max(x: Tensor, new_min = 0.0, new_max = 1.0): def linear_conversion(x, x_min=0.0, x_max=1.0, new_min=0.0, new_max=1.0): return (((x - x_min)/(x_max - x_min)) * (new_max - new_min)) + new_min +def extend_to_batch_size(tensor: Tensor, batch_size: int): + if tensor.shape[0] > batch_size: + return tensor[:batch_size] + elif tensor.shape[0] < batch_size: + remainder = batch_size-tensor.shape[0] + return torch.cat([tensor] + [tensor[-1:]]*remainder, dim=0) + return tensor -def broadcast_image_to_full(tensor, target_batch_size, batched_number, except_one=True): +def broadcast_image_to_extend(tensor, target_batch_size, batched_number, except_one=True): current_batch_size = tensor.shape[0] #print(current_batch_size, target_batch_size) if except_one and current_batch_size == 1: @@ -444,7 +453,7 @@ def broadcast_image_to_full(tensor, target_batch_size, batched_number, except_on tensor = tensor[:per_batch] if per_batch > tensor.shape[0]: - tensor = torch.cat([tensor] * (per_batch // tensor.shape[0]) + [tensor[:(per_batch % tensor.shape[0])]], dim=0) + tensor = extend_to_batch_size(tensor=tensor, batch_size=per_batch) current_batch_size = tensor.shape[0] if current_batch_size == target_batch_size: @@ -772,6 +781,8 @@ class AdvancedControlBase: # if uncond, set to weight's uncond_multiplier if cond_type == 1: x[actual_length*idx:actual_length*(idx+1)] *= self.weights.uncond_multiplier + if self.weights.has_uncond_mask: + pass if self.latent_keyframes is not None: x[:] = x[:] * self.calc_latent_keyframe_mults(x=x, batched_number=batched_number) @@ -860,12 +871,12 @@ class AdvancedControlBase: # resize mask and match batch count out_mask = prepare_mask_batch(orig_mask, x_noisy.shape, multiplier=multiplier) 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) + out_mask = extend_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) + out_mask = broadcast_image_to_extend(out_mask, x_noisy.shape[0], batched_number) # default dtype to be same as x_noisy if dtype is None: dtype = x_noisy.dtype