From b83d6f7213df6c469868c6af204861dc47aeab95 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Mon, 19 Feb 2024 10:33:28 -0600 Subject: [PATCH 01/17] Initial work on Reference CN support --- adv_control/control.py | 55 +--- adv_control/control_reference.py | 545 +++++++++++++++++++++++++++++++ adv_control/nodes.py | 7 + adv_control/nodes_reference.py | 57 ++++ adv_control/utils.py | 75 +++++ 5 files changed, 685 insertions(+), 54 deletions(-) diff --git a/adv_control/control.py b/adv_control/control.py index 261d052..f7db9ae 100644 --- a/adv_control/control.py +++ b/adv_control/control.py @@ -12,6 +12,7 @@ from model_patcher import ModelPatcher from .control_sparsectrl import SparseControlNet, SparseCtrlMotionWrapper, SparseMethod, SparseSettings, SparseSpreadMethod, PreprocSparseRGBWrapper from .control_lllite import LLLiteModule, LLLitePatch +from .control_reference import MachineState, ReferenceOptions, ReferenceAttnPatch 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) @@ -348,60 +349,6 @@ class SparseCtrlAdvanced(ControlNetAdvanced): return c -class ReferenceAdvanced(ControlBase, AdvancedControlBase): - def __init__(self, timestep_keyframes: TimestepKeyframeGroup, device=None): - super().__init__(device) - AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controllllite(), require_model=True) - # TODO: save attn patches here - - def patch_model(self, model: ModelPatcher): - # TODO: do model patching here - pass - - def pre_run_advanced(self, *args, **kwargs): - AdvancedControlBase.pre_run_advanced(self, *args, **kwargs) - # TODO: set control on patches - - def get_control_advanced(self, x_noisy: Tensor, t, cond, batched_number: int): - # normal ControlNet stuff - control_prev = None - if self.previous_controlnet is not None: - control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number) - - if self.timestep_range is not None: - if t[0] > self.timestep_range[0] or t[0] < self.timestep_range[1]: - return control_prev - - dtype = x_noisy.dtype - # prepare cond_hint - if self.sub_idxs is not None or self.cond_hint is None or x_noisy.shape[2] * 8 != self.cond_hint.shape[2] or x_noisy.shape[3] * 8 != self.cond_hint.shape[3]: - if self.cond_hint is not None: - 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) - 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) - # prepare mask - self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number) - # done preparing; model patches will take care of everything now. - # return normal controlnet stuff - return control_prev - - def cleanup_advanced(self): - super().cleanup_advanced() - # TODO: cleanup patches here - - def copy(self): - c = ReferenceAdvanced(self.timestep_keyframes) - self.copy_to(c) - self.copy_to_advanced(c) - return c - - class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase): # This ControlNet is more of an attention patch than a traditional controlnet def __init__(self, patch_attn1: LLLitePatch, patch_attn2: LLLitePatch, timestep_keyframes: TimestepKeyframeGroup, device=None): diff --git a/adv_control/control_reference.py b/adv_control/control_reference.py index e69de29..930d5b2 100644 --- a/adv_control/control_reference.py +++ b/adv_control/control_reference.py @@ -0,0 +1,545 @@ +from typing import Callable, Union + +import math +import torch +from torch import Tensor + +import comfy.model_patcher +import comfy.utils +from comfy.controlnet import ControlBase +from comfy.model_patcher import ModelPatcher +from comfy.ldm.modules.attention import BasicTransformerBlock + +from .logger import logger +from .utils import (AdvancedControlBase, ControlWeights, TimestepKeyframeGroup, AbstractPreprocWrapper, + deepcopy_with_sharing, prepare_mask_batch, broadcast_image_to_full, ddpm_noise_latents, simple_noise_latents) + + +REF_CONTROL_LIST = "ref_control_list" +REF_CONTROL_INFO = "ref_control_info" +REF_MACHINE_STATE = "ref_machine_state" + + +class MachineState: + WRITE = "write" + READ = "read" + STYLEALIGN = "stylealign" + + +class ReferenceType: + ATTN = "reference_attn" + ADAIN = "reference_adain" + ATTN_ADAIN = "reference_attn+adain" + STYLE_ALIGN = "StyleAlign" + + _LIST = [ATTN, ADAIN, ATTN_ADAIN] + + +class ReferenceOptions: + def __init__(self, reference_type: str, style_fidelity: float): + self.reference_type = reference_type + self.original_style_fidelity = style_fidelity + self.style_fidelity = style_fidelity + + def clone(self): + return ReferenceOptions(reference_type=self.reference_type, style_fidelity=self.original_style_fidelity) + + +class ReferencePreprocWrapper(AbstractPreprocWrapper): + error_msg = error_msg = "Invalid use of Reference Preprocess output. The output of RGB SparseCtrl preprocessor is NOT a usual image, but a latent pretending to be an image - you must connect the output directly to an Apply Advanced ControlNet node. It cannot be used for anything else that accepts IMAGE input." + def __init__(self, condhint: Tensor): + super().__init__(condhint) + + +class ReferenceAttnPatch: + def __init__(self, control: 'ReferenceAdvanced'=None): + self.control = control + + # def __call__(self, q: Tensor, k: Tensor, v: Tensor, extra_options: dict): + # # do nothing here - all ReferenceAttnPatch is trying to do is be a + # # ComfyUI-compliant way of tracking the corresponding ControlNet obj + # return q, k, v + + def __call__(self, x: Tensor, extra_options: dict): + # do nothing here - all ReferenceAttnPatch is trying to do is be a + # ComfyUI-compliant way of tracking the corresponding ControlNet obj + return x + + def set_control(self, control: 'ReferenceAdvanced') -> 'ReferenceAttnPatch': + self.control = control + return self + + def cleanup(self): + pass + + # make sure deepcopy does not copy control, and deepcopied patch should be assigned to control + def __deepcopy__(self, memo): + self.cleanup() + to_return: ReferenceAttnPatch = deepcopy_with_sharing(self, shared_attribute_names = ['control'], memo=memo) + #logger.warn(f"patch {id(self)} turned into {id(to_return)}") + try: + to_return.control.patch_attn1 = to_return + except Exception: + pass + return to_return + + +class ReferenceAdvanced(ControlBase, AdvancedControlBase): + def __init__(self, patch_attn1: ReferenceAttnPatch, ref_opts: ReferenceOptions, timestep_keyframes: TimestepKeyframeGroup, device=None): + super().__init__(device) + AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controllllite(), require_model=True) + self.patch_attn1 = patch_attn1.set_control(self) + self.ref_opts = ref_opts + self.order = 0 + self.latent_format = None + + def get_effective_strength(self): + effective_strength = self.strength + if self.current_timestep_keyframe is not None: + effective_strength = effective_strength * self.current_timestep_keyframe.strength + return effective_strength + + def patch_model(self, model: ModelPatcher): + # need to patch model so that control can be found later from it + model.set_model_attn1_output_patch(self.patch_attn1) + #model.set_model_attn1_patch(self.patch_attn1) + # need to add model_options to make patch/unpatch injection know it has to run + if not REF_CONTROL_INFO in model.model_options: + model.model_options[REF_CONTROL_INFO] = 0 + self.order = model.model_options[REF_CONTROL_INFO] + model.model_options[REF_CONTROL_INFO] + + def pre_run_advanced(self, model, percent_to_timestep_function): + AdvancedControlBase.pre_run_advanced(self, model, percent_to_timestep_function) + if type(self.cond_hint_original) == ReferencePreprocWrapper: + self.cond_hint_original = self.cond_hint_original.condhint + self.latent_format = model.latent_format # LatentFormat object, used to process_in latent cond_hint + # SDXL is more sensitive to style_fidelity according to sd-webui-controlnet comments + if type(model).__name__ == "SDXL": + self.ref_opts.style_fidelity = self.ref_opts.style_fidelity ** 3.0 + else: + self.ref_opts.style_fidelity = self.ref_opts.style_fidelity + # set control on patches + self.patch_attn1.set_control(self) + + def get_control_advanced(self, x_noisy: Tensor, t, cond, batched_number: int): + # normal ControlNet stuff + control_prev = None + if self.previous_controlnet is not None: + control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number) + + if self.timestep_range is not None: + if t[0] > self.timestep_range[0] or t[0] < self.timestep_range[1]: + return control_prev + + dtype = x_noisy.dtype + # prepare cond_hint - it is a latent, NOT an image + #if self.sub_idxs is not None or self.cond_hint is None or x_noisy.shape[2] != self.cond_hint.shape[2] or x_noisy.shape[3] != self.cond_hint.shape[3]: + if self.cond_hint is not None: + 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], x_noisy.shape[2], 'nearest-exact', "center").to(dtype).to(self.device) + else: + self.cond_hint = comfy.utils.common_upscale( + 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) + # noise cond_hint based on sigma (current step) + # TODO: how to handle noise? reproducibility is key... + # mess with the order here? + self.cond_hint = ddpm_noise_latents(self.cond_hint, sigma=t[0] / (self.latent_format.scale_factor), noise=None) + self.cond_hint = self.latent_format.process_in(self.cond_hint) + #self.cond_hint = simple_noise_latents(self.cond_hint, sigma=t[0], noise=None) + + # prepare mask + self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number) + # done preparing; model patches will take care of everything now. + # return normal controlnet stuff + return control_prev + + def cleanup_advanced(self): + super().cleanup_advanced() + self.patch_attn1.cleanup() + if self.latent_format is not None: + del self.latent_format + self.latent_format = None + + def copy(self): + c = ReferenceAdvanced(self.patch_attn1, self.ref_opts, self.timestep_keyframes) + c.order = self.order + self.copy_to(c) + self.copy_to_advanced(c) + return c + + +class BankStylesBasicTransformerBlock: + def __init__(self): + self.bank = [] + self.style_cfgs = [] + + def get_avg_style_fidelity(self): + return sum(self.style_cfgs) / float(len(self.style_cfgs)) + + def clean(self): + del self.bank + self.bank = [] + del self.style_cfgs + self.style_cfgs = [] + + +class InjectionBasicTransformerBlockHolder: + def __init__(self, block: BasicTransformerBlock): + self.original_forward = block._forward + self.attn_weight = 1.0 + self.bank_styles: dict[int, BankStylesBasicTransformerBlock] = {} + + def restore(self, block: BasicTransformerBlock): + block._forward = self.original_forward + + def clean(self): + for bank_style in list(self.bank_styles.values()): + bank_style.clean() + self.bank_styles.clear() + + +# def factory_clone_injected_ModelPatcher(orig_clone_ModelPatcher: Callable): +# def clone_injected_ref(self, *args, **kwargs): +# cloned = orig_clone_ModelPatcher(*args, **kwargs) +# if InjectMP.is_injected(self): +# InjectMP.mark_injected(cloned) +# cloned.clone = factory_clone_injected_ModelPatcher(cloned.clone).__get__(cloned, type(cloned)) +# return cloned +# return clone_injected_ref + + +# inject ModelPatcher.clone so that necessary injection will happen when needed +# orig_modelpatcher_clone = comfy.model_patcher.ModelPatcher.clone +# def clone_injection_ref(self: ModelPatcher, *args, **kwargs): +# cloned = orig_modelpatcher_clone(self, *args, **kwargs) +# if InjectMP.is_injected(self): +# InjectMP.mark_injected(cloned) +# return cloned +# comfy.model_patcher.ModelPatcher.clone = clone_injection_ref + + +# inject ModelPatcher.patch_model to apply +orig_modelpatcher_patch_model = comfy.model_patcher.ModelPatcher.patch_model +def patch_model_injection_ref(self: ModelPatcher, *args, **kwargs): + if REF_CONTROL_INFO in self.model_options: + # storage for all Reference-related injections + reference_injections = ReferenceInjections() + # first, handle attn module injection + all_modules = torch_dfs(self.model) + attn_modules: list[RefBasicTransformerBlock] = [] + for module in all_modules: + if isinstance(module, BasicTransformerBlock): + attn_modules.append(module) + attn_modules = [module for module in all_modules if isinstance(module, BasicTransformerBlock)] + attn_modules = sorted(attn_modules, key=lambda x: -x.norm1.normalized_shape[0]) + for i, module in enumerate(attn_modules): + injection_holder = InjectionBasicTransformerBlockHolder(block=module) + injection_holder.attn_weight = float(i) / float(len(attn_modules)) + module._forward = _forward_inject_BasicTransformerBlock.__get__(module, type(module)) + module.injection_holder = injection_holder + reference_injections.attn_modules = attn_modules + # handle diffusion_model forward injection + reference_injections.diffusion_model_orig_forward = self.model.diffusion_model.forward + self.model.diffusion_model.forward = factory_forward_inject_UNetModel(reference_injections).__get__(self.model.diffusion_model, type(self.model.diffusion_model)) + InjectMP.set_injected(self, reference_injections) + to_return = orig_modelpatcher_patch_model(self, *args, **kwargs) + return to_return +comfy.model_patcher.ModelPatcher.patch_model = patch_model_injection_ref + + +orig_modelpatcher_unpatch_model = comfy.model_patcher.ModelPatcher.unpatch_model +def unpatch_model_injection_ref(self: ModelPatcher, *args, **kwargs): + if REF_CONTROL_INFO in self.model_options: + reference_injections: ReferenceInjections = InjectMP.get_injected(self) + # first, restore attn modules + attn_modules: list[RefBasicTransformerBlock] = reference_injections.attn_modules + for module in attn_modules: + module.injection_holder.restore(module) + module.injection_holder.clean() + del module.injection_holder + del attn_modules + # restore diffusion_model forward function + self.model.diffusion_model.forward = reference_injections.diffusion_model_orig_forward.__get__(self.model.diffusion_model, type(self.model.diffusion_model)) + # cleanup + InjectMP.clean_injected(self) + reference_injections.cleanup() + to_return = orig_modelpatcher_unpatch_model(self, *args, **kwargs) + return to_return +comfy.model_patcher.ModelPatcher.unpatch_model = unpatch_model_injection_ref + + +class ReferenceInjections: + def __init__(self, attn_modules: list['RefBasicTransformerBlock']=None): + self.attn_modules = attn_modules if attn_modules else [] + self.diffusion_model_orig_forward: Callable = None + + def clean_module_mem(self): + for attn_module in self.attn_modules: + attn_module.injection_holder.clean() + + def cleanup(self): + del self.attn_modules + self.attn_modules = [] + self.diffusion_model_orig_forward = None + + +class InjectMP: + PARAM_INJECTED_REF = "___injected_ref" + + @staticmethod + def is_injected(model: ModelPatcher): + return getattr(model, InjectMP.PARAM_INJECTED_REF, False) + + def mark_injected(model: ModelPatcher): + setattr(model, InjectMP.PARAM_INJECTED_REF, True) + + def set_injected(model: ModelPatcher, value: ReferenceInjections): + setattr(model, InjectMP.PARAM_INJECTED_REF, value) + + def get_injected(model: ModelPatcher) -> ReferenceInjections: + return getattr(model, InjectMP.PARAM_INJECTED_REF) + + def clean_injected(model: ModelPatcher): + delattr(model, InjectMP.PARAM_INJECTED_REF) + InjectMP.mark_injected(model) + + +def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections): + def forward_inject_UNetModel(self, x: Tensor, *args, **kwargs): + # get control and transformer_options from kwargs + real_args = list(args) + real_kwargs = list(kwargs.keys()) + control = kwargs.get("control", None) + transformer_options = kwargs.get("transformer_options", None) + # look for ReferenceAttnPatch objects to get ReferenceAdvanced objects + # and remove the patch from transformer_options so it won't be ran for no reason + patch_name = "attn1_output_patch" + #patch_name = "attn1_patch" + ref_patches: list[ReferenceAttnPatch] = [] + if "patches" in transformer_options: + if patch_name in transformer_options["patches"]: + patches: list = transformer_options["patches"][patch_name] + remove_idxs = [] + for i in range(len(patches)): + if isinstance(patches[i], ReferenceAttnPatch): + ref_patches.append(patches[i]) + remove_idxs.append(i) + # for i in reversed(remove_idxs): + # patches.pop(i) + # if len(transformer_options["patches"]["attn1_patch"]) == 0: + # transformer_options["patches"].pop("attn1_patch") + ref_controlnets: list[ReferenceAdvanced] = [x.control for x in ref_patches] + # discard any controlnets that should not run + ref_controlnets = [x for x in ref_controlnets if x.should_run()] + ref_controlnets = sorted(ref_controlnets, key=lambda x: x.order) + # if nothing related to reference controlnets, do nothing special + if len(ref_controlnets) == 0: + return reference_injections.diffusion_model_orig_forward(x, *args, **kwargs) + try: + # otherwise, need to handle ref controlnet stuff + for control in ref_controlnets: + transformer_options[REF_MACHINE_STATE] = MachineState.WRITE + transformer_options[REF_CONTROL_LIST] = [control] + # TODO: insert control's cond_hint + reference_injections.diffusion_model_orig_forward(control.cond_hint.to(dtype=x.dtype).to(device=x.device), *args, **kwargs) + transformer_options[REF_MACHINE_STATE] = MachineState.READ + transformer_options[REF_CONTROL_LIST] = ref_controlnets + return reference_injections.diffusion_model_orig_forward(x, *args, **kwargs) + finally: + # make sure banks are cleared no matter what happens - otherwise, RIP VRAM + reference_injections.clean_module_mem() + + return forward_inject_UNetModel + + +# dummy class just to help IDE keep track of injected variables +class RefBasicTransformerBlock(BasicTransformerBlock): + injection_holder: InjectionBasicTransformerBlockHolder = None + +def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Tensor, context: Tensor=None, transformer_options: dict[str]={}): + extra_options = {} + block = transformer_options.get("block", None) + block_index = transformer_options.get("block_index", 0) + transformer_patches = {} + transformer_patches_replace = {} + + for k in transformer_options: + if k == "patches": + transformer_patches = transformer_options[k] + elif k == "patches_replace": + transformer_patches_replace = transformer_options[k] + else: + extra_options[k] = transformer_options[k] + + extra_options["n_heads"] = self.n_heads + extra_options["dim_head"] = self.d_head + + if self.ff_in: + x_skip = x + x = self.ff_in(self.norm_in(x)) + if self.is_res: + x += x_skip + + n = self.norm1(x) + if self.disable_self_attn: + context_attn1 = context + else: + context_attn1 = None + value_attn1 = None + + # Reference CN stuff + # WRITE mode will only have one ReferenceAdvanced, other modes will have all ReferenceAdvanced + ref_controlnets: list[ReferenceAdvanced] = transformer_options.get(REF_CONTROL_LIST, None) + ref_machine_state: str = transformer_options.get(REF_MACHINE_STATE, None) + # if in WRITE mode, save n and style_fidelity + if ref_controlnets and ref_machine_state == MachineState.WRITE: + if ref_controlnets[0].get_effective_strength() > self.injection_holder.attn_weight: + if not self.injection_holder.bank_styles.get(ref_controlnets[0].order, None): + self.injection_holder.bank_styles[ref_controlnets[0].order] = BankStylesBasicTransformerBlock() + bank_style = self.injection_holder.bank_styles[ref_controlnets[0].order] + bank_style.bank.append(n.detach().clone()) + bank_style.style_cfgs.append(ref_controlnets[0].ref_opts.style_fidelity) + # create uc_idx_mask + per_batch = x.shape[0] // len(transformer_options["cond_or_uncond"]) + indiv_conds = [] + for cond_type in transformer_options["cond_or_uncond"]: + indiv_conds.extend([cond_type] * per_batch) + uc_idx_mask = [i for i, x in enumerate(indiv_conds) if x == 0] + + if "attn1_patch" in transformer_patches: + patch = transformer_patches["attn1_patch"] + if context_attn1 is None: + context_attn1 = n + value_attn1 = context_attn1 + for p in patch: + n, context_attn1, value_attn1 = p(n, context_attn1, value_attn1, extra_options) + + if block is not None: + transformer_block = (block[0], block[1], block_index) + else: + transformer_block = None + attn1_replace_patch = transformer_patches_replace.get("attn1", {}) + block_attn1 = transformer_block + if block_attn1 not in attn1_replace_patch: + block_attn1 = block + + if block_attn1 in attn1_replace_patch: + if context_attn1 is None: + context_attn1 = n + value_attn1 = n + n = self.attn1.to_q(n) + # Reference CN READ - use attn1_replace_patch appropriately + if ref_machine_state == MachineState.READ: + bank_styles = self.injection_holder.bank_styles[ref_controlnets[0].order] + style_fidelity = bank_styles.get_avg_style_fidelity() + n_uc = self.attn1.to_out(attn1_replace_patch[block_attn1]( + n, + self.attn1.to_k(torch.cat([context_attn1] + bank_styles.bank, dim=1)), + self.attn1.to_v(torch.cat([value_attn1] + bank_styles.bank, dim=1)), + extra_options)) + n_c = n_uc.clone() + if len(uc_idx_mask) > 0 and not math.isclose(style_fidelity, 0.0): + n_c[uc_idx_mask] = self.attn1.to_out(attn1_replace_patch[block_attn1]( + n[uc_idx_mask], + self.attn1.to_k(context_attn1[uc_idx_mask]), + self.attn1.to_v(value_attn1[uc_idx_mask]), + extra_options)) + n = style_fidelity * n_c + (1.0-style_fidelity) * n_uc + bank_styles.clean() + else: + context_attn1 = self.attn1.to_k(context_attn1) + value_attn1 = self.attn1.to_v(value_attn1) + n = attn1_replace_patch[block_attn1](n, context_attn1, value_attn1, extra_options) + n = self.attn1.to_out(n) + else: + # Reference CN READ - no attn1_replace_patch + if ref_machine_state == MachineState.READ: + if context_attn1 is None: + context_attn1 = n + bank_styles = self.injection_holder.bank_styles[ref_controlnets[0].order] + style_fidelity = bank_styles.get_avg_style_fidelity() + n_uc: Tensor = self.attn1( + n, + context=torch.cat([context_attn1] + bank_styles.bank, dim=1), + value=torch.cat([value_attn1] + bank_styles.bank, dim=1) if value_attn1 is not None else value_attn1) + n_c = n_uc.clone() + if len(uc_idx_mask) > 0 and not math.isclose(style_fidelity, 0.0): + n_c[uc_idx_mask] = self.attn1( + n[uc_idx_mask], + context=context_attn1[uc_idx_mask] if context_attn1 is not None else context_attn1, + value=value_attn1[uc_idx_mask] if value_attn1 is not None else value_attn1) + n = style_fidelity * n_c + (1.0-style_fidelity) * n_uc + bank_styles.clean() + else: + n = self.attn1(n, context=context_attn1, value=value_attn1) + + if "attn1_output_patch" in transformer_patches: + patch = transformer_patches["attn1_output_patch"] + for p in patch: + n = p(n, extra_options) + + x += n + if "middle_patch" in transformer_patches: + patch = transformer_patches["middle_patch"] + for p in patch: + x = p(x, extra_options) + + if self.attn2 is not None: + n = self.norm2(x) + if self.switch_temporal_ca_to_sa: + context_attn2 = n + else: + context_attn2 = context + value_attn2 = None + if "attn2_patch" in transformer_patches: + patch = transformer_patches["attn2_patch"] + value_attn2 = context_attn2 + for p in patch: + n, context_attn2, value_attn2 = p(n, context_attn2, value_attn2, extra_options) + + attn2_replace_patch = transformer_patches_replace.get("attn2", {}) + block_attn2 = transformer_block + if block_attn2 not in attn2_replace_patch: + block_attn2 = block + + if block_attn2 in attn2_replace_patch: + if value_attn2 is None: + value_attn2 = context_attn2 + n = self.attn2.to_q(n) + context_attn2 = self.attn2.to_k(context_attn2) + value_attn2 = self.attn2.to_v(value_attn2) + n = attn2_replace_patch[block_attn2](n, context_attn2, value_attn2, extra_options) + n = self.attn2.to_out(n) + else: + n = self.attn2(n, context=context_attn2, value=value_attn2) + + if "attn2_output_patch" in transformer_patches: + patch = transformer_patches["attn2_output_patch"] + for p in patch: + n = p(n, extra_options) + + x += n + if self.is_res: + x_skip = x + x = self.ff(self.norm3(x)) + if self.is_res: + x += x_skip + + return x + + +# DFS Search for Torch.nn.Module, Written by Lvmin +def torch_dfs(model: torch.nn.Module): + result = [model] + for child in model.children(): + result += torch_dfs(child) + return result diff --git a/adv_control/nodes.py b/adv_control/nodes.py index 1cea1af..ce4c0e7 100644 --- a/adv_control/nodes.py +++ b/adv_control/nodes.py @@ -11,6 +11,7 @@ from .nodes_weight import (DefaultWeights, ScaledSoftMaskedUniversalWeights, Sca SoftT2IAdapterWeights, CustomT2IAdapterWeights) from .nodes_latent_keyframe import LatentKeyframeGroupNode, LatentKeyframeInterpolationNode, LatentKeyframeBatchedGroupNode, LatentKeyframeNode from .nodes_sparsectrl import SparseCtrlMergedLoaderAdvanced, SparseCtrlLoaderAdvanced, SparseIndexMethodNode, SparseSpreadMethodNode, RgbSparseCtrlPreprocessor +from .nodes_reference import ReferenceControlNetNode, ReferencePreprocessorNode from .nodes_loosecontrol import ControlNetLoaderWithLoraAdvanced from .nodes_deprecated import LoadImagesFromDirectory from .logger import logger @@ -233,6 +234,9 @@ NODE_CLASS_MAPPINGS = { "ACN_SparseCtrlMergedLoaderAdvanced": SparseCtrlMergedLoaderAdvanced, "ACN_SparseCtrlIndexMethodNode": SparseIndexMethodNode, "ACN_SparseCtrlSpreadMethodNode": SparseSpreadMethodNode, + # Reference + "ACN_ReferencePreprocessor": ReferencePreprocessorNode, + "ACN_ReferenceControlNet": ReferenceControlNetNode, # LOOSEControl #"ACN_ControlNetLoaderWithLoraAdvanced": ControlNetLoaderWithLoraAdvanced, # Deprecated @@ -265,6 +269,9 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ACN_SparseCtrlMergedLoaderAdvanced": "Load Merged SparseCtrl Model 🛂🅐🅒🅝", "ACN_SparseCtrlIndexMethodNode": "SparseCtrl Index Method 🛂🅐🅒🅝", "ACN_SparseCtrlSpreadMethodNode": "SparseCtrl Spread Method 🛂🅐🅒🅝", + # Reference + "ACN_ReferencePreprocessor": "Reference Preproccessor 🛂🅐🅒🅝", + "ACN_ReferenceControlNet": "Reference ControlNet 🛂🅐🅒🅝", # LOOSEControl #"ACN_ControlNetLoaderWithLoraAdvanced": "Load Adv. ControlNet Model w/ LoRA 🛂🅐🅒🅝", # Deprecated diff --git a/adv_control/nodes_reference.py b/adv_control/nodes_reference.py index e69de29..0730647 100644 --- a/adv_control/nodes_reference.py +++ b/adv_control/nodes_reference.py @@ -0,0 +1,57 @@ +from torch import Tensor + +from nodes import VAEEncode +import comfy.utils + +from .control_reference import ReferenceAdvanced, ReferenceAttnPatch, ReferenceOptions, ReferenceType, ReferencePreprocWrapper + + +# node for ReferenceCN +class ReferenceControlNetNode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "reference_type": (ReferenceType._LIST,), + "style_fidelity": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}) + }, + } + + RETURN_TYPES = ("CONTROL_NET", ) + FUNCTION = "load_controlnet" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/Reference" + + def load_controlnet(self, reference_type: str, style_fidelity: float): + ref_opts = ReferenceOptions(reference_type=reference_type, style_fidelity=style_fidelity) + ref_patch = ReferenceAttnPatch() + controlnet = ReferenceAdvanced(patch_attn1=ref_patch, ref_opts=ref_opts, timestep_keyframes=None) + return (controlnet,) + + +class ReferencePreprocessorNode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE", ), + "vae": ("VAE", ), + "latent_size": ("LATENT", ), + } + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("proc_IMAGE",) + FUNCTION = "preprocess_images" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/Reference/preprocess" + + def preprocess_images(self, vae, image: Tensor, latent_size: Tensor): + # first, resize image to match latents + image = image.movedim(-1,1) + image = comfy.utils.common_upscale(image, latent_size["samples"].shape[3] * 8, latent_size["samples"].shape[2] * 8, 'nearest-exact', "center") + image = image.movedim(1,-1) + # then, vae encode + image = VAEEncode.vae_encode_crop_pixels(image) + encoded = vae.encode(image[:,:,:,:3]) + return (ReferencePreprocWrapper(condhint=encoded),) diff --git a/adv_control/utils.py b/adv_control/utils.py index 1d77707..85a35de 100644 --- a/adv_control/utils.py +++ b/adv_control/utils.py @@ -3,6 +3,7 @@ from typing import Callable, Union import torch from torch import Tensor import torch.nn.functional as F +import math import comfy.ops import comfy.utils @@ -232,6 +233,38 @@ class TimestepKeyframeGroup: return group +class AbstractPreprocWrapper: + error_msg = "Invalid use of [InsertHere] output. The output of [InsertHere] preprocessor is NOT a usual image, but a latent pretending to be an image - you must connect the output directly to an Apply ControlNet node (advanced or otherwise). It cannot be used for anything else that accepts IMAGE input." + def __init__(self, condhint: Tensor): + self.condhint = condhint + + def movedim(self, *args, **kwargs): + return self + + def __getattr__(self, *args, **kwargs): + raise AttributeError(self.error_msg) + + def __setattr__(self, name, value): + if name != "condhint": + raise AttributeError(self.error_msg) + super().__setattr__(name, value) + + def __iter__(self, *args, **kwargs): + raise AttributeError(self.error_msg) + + def __next__(self, *args, **kwargs): + raise AttributeError(self.error_msg) + + def __len__(self, *args, **kwargs): + raise AttributeError(self.error_msg) + + def __getitem__(self, *args, **kwargs): + raise AttributeError(self.error_msg) + + def __setitem__(self, *args, **kwargs): + raise AttributeError(self.error_msg) + + # depending on model, AnimateDiff may inject into GroupNorm, so make sure GroupNorm will be clean class disable_weight_init_clean_groupnorm(comfy.ops.disable_weight_init): class GroupNorm(comfy.ops.disable_weight_init.GroupNorm): @@ -264,6 +297,40 @@ 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 broadcast_image_to_full(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: + return tensor + + per_batch = target_batch_size // batched_number + 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) + + current_batch_size = tensor.shape[0] + if current_batch_size == target_batch_size: + return tensor + else: + return torch.cat([tensor] * batched_number, dim=0) + + +def ddpm_noise_latents(latents: Tensor, sigma: float, noise: Tensor=None): + alpha_cumprod = 1 / ((sigma * sigma) + 1) + sqrt_alpha_prod = alpha_cumprod ** 0.5 + sqrt_one_minus_alpha_prod = (1 - alpha_cumprod) ** 0.5 + if noise is None: + noise = torch.rand_like(latents) + return latents * sqrt_alpha_prod + noise * sqrt_one_minus_alpha_prod + + +def simple_noise_latents(latents: Tensor, sigma: float, noise: Tensor=None): + if noise is None: + noise = torch.rand_like(latents) + return latents + noise * sigma + + # from https://stackoverflow.com/a/24621200 def deepcopy_with_sharing(obj, shared_attribute_names, memo=None): ''' @@ -458,6 +525,14 @@ class AdvancedControlBase: def disarm(self): self.disarmed = True + def should_run(self): + if math.isclose(self.strength, 0.0) or math.isclose(self.current_timestep_keyframe.strength, 0.0): + return False + if self.timestep_range is not None: + if self.t > self.timestep_range[0] or self.t < self.timestep_range[1]: + return False + return True + def get_control_inject(self, x_noisy, t, cond, batched_number): # prepare timestep and everything related self.prepare_current_timestep(t=t, batched_number=batched_number) From f1ef8bb3a02b750c3e0b8f831920718b4228fcc8 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Fri, 1 Mar 2024 02:16:48 -0600 Subject: [PATCH 02/17] Continued debugging of ref cn --- adv_control/control_reference.py | 4 +++- adv_control/nodes_reference.py | 8 ++++++-- adv_control/utils.py | 4 +++- 3 files changed, 12 insertions(+), 4 deletions(-) diff --git a/adv_control/control_reference.py b/adv_control/control_reference.py index 930d5b2..a7869a0 100644 --- a/adv_control/control_reference.py +++ b/adv_control/control_reference.py @@ -24,6 +24,7 @@ class MachineState: WRITE = "write" READ = "read" STYLEALIGN = "stylealign" + TEST = "test" class ReferenceType: @@ -152,8 +153,9 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): # noise cond_hint based on sigma (current step) # TODO: how to handle noise? reproducibility is key... # mess with the order here? - self.cond_hint = ddpm_noise_latents(self.cond_hint, sigma=t[0] / (self.latent_format.scale_factor), noise=None) + # / (self.latent_format.scale_factor) self.cond_hint = self.latent_format.process_in(self.cond_hint) + self.cond_hint = ddpm_noise_latents(self.cond_hint, sigma=t[0], noise=None) #self.cond_hint = simple_noise_latents(self.cond_hint, sigma=t[0], noise=None) # prepare mask diff --git a/adv_control/nodes_reference.py b/adv_control/nodes_reference.py index 0730647..06236d2 100644 --- a/adv_control/nodes_reference.py +++ b/adv_control/nodes_reference.py @@ -2,6 +2,7 @@ from torch import Tensor from nodes import VAEEncode import comfy.utils +from comfy.sd import VAE from .control_reference import ReferenceAdvanced, ReferenceAttnPatch, ReferenceOptions, ReferenceType, ReferencePreprocWrapper @@ -46,12 +47,15 @@ class ReferencePreprocessorNode: CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/Reference/preprocess" - def preprocess_images(self, vae, image: Tensor, latent_size: Tensor): + def preprocess_images(self, vae: VAE, image: Tensor, latent_size: Tensor): # first, resize image to match latents image = image.movedim(-1,1) image = comfy.utils.common_upscale(image, latent_size["samples"].shape[3] * 8, latent_size["samples"].shape[2] * 8, 'nearest-exact', "center") image = image.movedim(1,-1) # then, vae encode - image = VAEEncode.vae_encode_crop_pixels(image) + try: + image = vae.vae_encode_crop_pixels(image) + except Exception: + image = VAEEncode.vae_encode_crop_pixels(image) encoded = vae.encode(image[:,:,:,:3]) return (ReferencePreprocWrapper(condhint=encoded),) diff --git a/adv_control/utils.py b/adv_control/utils.py index 85a35de..ac3cf3d 100644 --- a/adv_control/utils.py +++ b/adv_control/utils.py @@ -321,7 +321,9 @@ def ddpm_noise_latents(latents: Tensor, sigma: float, noise: Tensor=None): sqrt_alpha_prod = alpha_cumprod ** 0.5 sqrt_one_minus_alpha_prod = (1 - alpha_cumprod) ** 0.5 if noise is None: - noise = torch.rand_like(latents) + generator = torch.cuda.manual_seed(0) + noise = torch.empty_like(latents).normal_(generator=generator) + #noise = torch.rand_like(latents) return latents * sqrt_alpha_prod + noise * sqrt_one_minus_alpha_prod From b370f269862cb951a6e22948c1ce4e59e811ac46 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sat, 2 Mar 2024 10:25:27 -0600 Subject: [PATCH 03/17] More progress on ref cn --- adv_control/control_reference.py | 80 +++++++++++++++----------------- adv_control/utils.py | 8 +++- 2 files changed, 45 insertions(+), 43 deletions(-) diff --git a/adv_control/control_reference.py b/adv_control/control_reference.py index a7869a0..7c50092 100644 --- a/adv_control/control_reference.py +++ b/adv_control/control_reference.py @@ -93,6 +93,7 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): self.ref_opts = ref_opts self.order = 0 self.latent_format = None + self.model_sampling_current = None def get_effective_strength(self): effective_strength = self.strength @@ -115,6 +116,7 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): if type(self.cond_hint_original) == ReferencePreprocWrapper: self.cond_hint_original = self.cond_hint_original.condhint self.latent_format = model.latent_format # LatentFormat object, used to process_in latent cond_hint + self.model_sampling_current = model.model_sampling # SDXL is more sensitive to style_fidelity according to sd-webui-controlnet comments if type(model).__name__ == "SDXL": self.ref_opts.style_fidelity = self.ref_opts.style_fidelity ** 3.0 @@ -155,7 +157,10 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): # mess with the order here? # / (self.latent_format.scale_factor) self.cond_hint = self.latent_format.process_in(self.cond_hint) + #self.cond_hint = self.model_sampling_current.calculate_input(t, self.cond_hint) self.cond_hint = ddpm_noise_latents(self.cond_hint, sigma=t[0], noise=None) + timestep = self.model_sampling_current.timestep(t) + #self.cond_hint = ddpm_noise_latents(torch.zeros_like(x_noisy), sigma=t[0], noise=None) #self.cond_hint = simple_noise_latents(self.cond_hint, sigma=t[0], noise=None) # prepare mask @@ -167,9 +172,10 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): def cleanup_advanced(self): super().cleanup_advanced() self.patch_attn1.cleanup() - if self.latent_format is not None: - del self.latent_format - self.latent_format = None + del self.latent_format + self.latent_format = None + del self.model_sampling_current + self.model_sampling_current = None def copy(self): c = ReferenceAdvanced(self.patch_attn1, self.ref_opts, self.timestep_keyframes) @@ -195,8 +201,9 @@ class BankStylesBasicTransformerBlock: class InjectionBasicTransformerBlockHolder: - def __init__(self, block: BasicTransformerBlock): + def __init__(self, block: BasicTransformerBlock, idx=None): self.original_forward = block._forward + self.idx = idx self.attn_weight = 1.0 self.bank_styles: dict[int, BankStylesBasicTransformerBlock] = {} @@ -209,26 +216,6 @@ class InjectionBasicTransformerBlockHolder: self.bank_styles.clear() -# def factory_clone_injected_ModelPatcher(orig_clone_ModelPatcher: Callable): -# def clone_injected_ref(self, *args, **kwargs): -# cloned = orig_clone_ModelPatcher(*args, **kwargs) -# if InjectMP.is_injected(self): -# InjectMP.mark_injected(cloned) -# cloned.clone = factory_clone_injected_ModelPatcher(cloned.clone).__get__(cloned, type(cloned)) -# return cloned -# return clone_injected_ref - - -# inject ModelPatcher.clone so that necessary injection will happen when needed -# orig_modelpatcher_clone = comfy.model_patcher.ModelPatcher.clone -# def clone_injection_ref(self: ModelPatcher, *args, **kwargs): -# cloned = orig_modelpatcher_clone(self, *args, **kwargs) -# if InjectMP.is_injected(self): -# InjectMP.mark_injected(cloned) -# return cloned -# comfy.model_patcher.ModelPatcher.clone = clone_injection_ref - - # inject ModelPatcher.patch_model to apply orig_modelpatcher_patch_model = comfy.model_patcher.ModelPatcher.patch_model def patch_model_injection_ref(self: ModelPatcher, *args, **kwargs): @@ -243,12 +230,16 @@ def patch_model_injection_ref(self: ModelPatcher, *args, **kwargs): attn_modules.append(module) attn_modules = [module for module in all_modules if isinstance(module, BasicTransformerBlock)] attn_modules = sorted(attn_modules, key=lambda x: -x.norm1.normalized_shape[0]) + reference_injections.attn_modules = [] for i, module in enumerate(attn_modules): - injection_holder = InjectionBasicTransformerBlockHolder(block=module) + # if i != 11 and i != 12: + # continue + injection_holder = InjectionBasicTransformerBlockHolder(block=module, idx=i) injection_holder.attn_weight = float(i) / float(len(attn_modules)) module._forward = _forward_inject_BasicTransformerBlock.__get__(module, type(module)) module.injection_holder = injection_holder - reference_injections.attn_modules = attn_modules + reference_injections.attn_modules.append(module) + # handle diffusion_model forward injection reference_injections.diffusion_model_orig_forward = self.model.diffusion_model.forward self.model.diffusion_model.forward = factory_forward_inject_UNetModel(reference_injections).__get__(self.model.diffusion_model, type(self.model.diffusion_model)) @@ -286,9 +277,13 @@ class ReferenceInjections: def clean_module_mem(self): for attn_module in self.attn_modules: - attn_module.injection_holder.clean() + try: + attn_module.injection_holder.clean() + except Exception: + pass def cleanup(self): + self.clean_module_mem() del self.attn_modules self.attn_modules = [] self.diffusion_model_orig_forward = None @@ -323,22 +318,14 @@ def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections): control = kwargs.get("control", None) transformer_options = kwargs.get("transformer_options", None) # look for ReferenceAttnPatch objects to get ReferenceAdvanced objects - # and remove the patch from transformer_options so it won't be ran for no reason patch_name = "attn1_output_patch" - #patch_name = "attn1_patch" ref_patches: list[ReferenceAttnPatch] = [] if "patches" in transformer_options: if patch_name in transformer_options["patches"]: patches: list = transformer_options["patches"][patch_name] - remove_idxs = [] for i in range(len(patches)): if isinstance(patches[i], ReferenceAttnPatch): ref_patches.append(patches[i]) - remove_idxs.append(i) - # for i in reversed(remove_idxs): - # patches.pop(i) - # if len(transformer_options["patches"]["attn1_patch"]) == 0: - # transformer_options["patches"].pop("attn1_patch") ref_controlnets: list[ReferenceAdvanced] = [x.control for x in ref_patches] # discard any controlnets that should not run ref_controlnets = [x for x in ref_controlnets if x.should_run()] @@ -351,7 +338,10 @@ def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections): for control in ref_controlnets: transformer_options[REF_MACHINE_STATE] = MachineState.WRITE transformer_options[REF_CONTROL_LIST] = [control] - # TODO: insert control's cond_hint + # from pathlib import Path + # with open(Path(__file__).parent.parent.parent.parent.parent / "ref_debug" / "ref_xt_noised.pt", "rb") as rfile: + # ref_xt = torch.load(rfile, weights_only=True) + # diffuse cond_hint reference_injections.diffusion_model_orig_forward(control.cond_hint.to(dtype=x.dtype).to(device=x.device), *args, **kwargs) transformer_options[REF_MACHINE_STATE] = MachineState.READ transformer_options[REF_CONTROL_LIST] = ref_controlnets @@ -391,7 +381,7 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten if self.is_res: x += x_skip - n = self.norm1(x) + n: Tensor = self.norm1(x) if self.disable_self_attn: context_attn1 = context else: @@ -410,12 +400,17 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten bank_style = self.injection_holder.bank_styles[ref_controlnets[0].order] bank_style.bank.append(n.detach().clone()) bank_style.style_cfgs.append(ref_controlnets[0].ref_opts.style_fidelity) + # from pathlib import Path + # with open(Path(__file__).parent.parent.parent.parent.parent / "ref_debug" / f"bank_{self.injection_holder.idx}.pt", "rb") as rfile: + # raw_val = torch.load(rfile) + # raw_val[0] = raw_val[0].to(n.dtype).to(n.device) + # bank_style.bank.extend(raw_val) # create uc_idx_mask per_batch = x.shape[0] // len(transformer_options["cond_or_uncond"]) indiv_conds = [] for cond_type in transformer_options["cond_or_uncond"]: indiv_conds.extend([cond_type] * per_batch) - uc_idx_mask = [i for i, x in enumerate(indiv_conds) if x == 0] + uc_idx_mask = [i for i, x in enumerate(indiv_conds) if x == 1] if "attn1_patch" in transformer_patches: patch = transformer_patches["attn1_patch"] @@ -440,7 +435,7 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten value_attn1 = n n = self.attn1.to_q(n) # Reference CN READ - use attn1_replace_patch appropriately - if ref_machine_state == MachineState.READ: + if ref_machine_state == MachineState.READ and self.injection_holder.bank_styles.get(ref_controlnets[0].order, None) is not None: bank_styles = self.injection_holder.bank_styles[ref_controlnets[0].order] style_fidelity = bank_styles.get_avg_style_fidelity() n_uc = self.attn1.to_out(attn1_replace_patch[block_attn1]( @@ -464,7 +459,7 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten n = self.attn1.to_out(n) else: # Reference CN READ - no attn1_replace_patch - if ref_machine_state == MachineState.READ: + if ref_machine_state == MachineState.READ and self.injection_holder.bank_styles.get(ref_controlnets[0].order, None) is not None: if context_attn1 is None: context_attn1 = n bank_styles = self.injection_holder.bank_styles[ref_controlnets[0].order] @@ -472,12 +467,13 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten n_uc: Tensor = self.attn1( n, context=torch.cat([context_attn1] + bank_styles.bank, dim=1), + #context=torch.cat(bank_styles.bank, dim=1), value=torch.cat([value_attn1] + bank_styles.bank, dim=1) if value_attn1 is not None else value_attn1) n_c = n_uc.clone() - if len(uc_idx_mask) > 0 and not math.isclose(style_fidelity, 0.0): + if len(uc_idx_mask) > 0 and style_fidelity > 1e-5:# not math.isclose(style_fidelity, 0.0): n_c[uc_idx_mask] = self.attn1( n[uc_idx_mask], - context=context_attn1[uc_idx_mask] if context_attn1 is not None else context_attn1, + context=n[uc_idx_mask], value=value_attn1[uc_idx_mask] if value_attn1 is not None else value_attn1) n = style_fidelity * n_c + (1.0-style_fidelity) * n_uc bank_styles.clean() diff --git a/adv_control/utils.py b/adv_control/utils.py index ac3cf3d..6bc0d22 100644 --- a/adv_control/utils.py +++ b/adv_control/utils.py @@ -320,11 +320,17 @@ def ddpm_noise_latents(latents: Tensor, sigma: float, noise: Tensor=None): alpha_cumprod = 1 / ((sigma * sigma) + 1) sqrt_alpha_prod = alpha_cumprod ** 0.5 sqrt_one_minus_alpha_prod = (1 - alpha_cumprod) ** 0.5 + #logger.warn(f"sqrt: {sqrt_alpha_prod}, sqrt-1: {sqrt_one_minus_alpha_prod}, t: {sigma}") if noise is None: + # generator = torch.manual_seed(0) + # noise = torch.randn(latents.size(), generator=generator).to(latents.device) generator = torch.cuda.manual_seed(0) noise = torch.empty_like(latents).normal_(generator=generator) + #noise = torch.empty(latents.size()).normal_(generator=generator).to(latents.device) + #return noise #noise = torch.rand_like(latents) - return latents * sqrt_alpha_prod + noise * sqrt_one_minus_alpha_prod + #return latents + return sqrt_alpha_prod * latents + sqrt_one_minus_alpha_prod * noise def simple_noise_latents(latents: Tensor, sigma: float, noise: Tensor=None): From 8be003562df5ebddc50a1120ccf49b7d6e350ee5 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sat, 2 Mar 2024 10:56:11 -0600 Subject: [PATCH 04/17] Make sure to reference context_attn1 --- adv_control/control_reference.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/adv_control/control_reference.py b/adv_control/control_reference.py index 7c50092..0b500c1 100644 --- a/adv_control/control_reference.py +++ b/adv_control/control_reference.py @@ -473,7 +473,7 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten if len(uc_idx_mask) > 0 and style_fidelity > 1e-5:# not math.isclose(style_fidelity, 0.0): n_c[uc_idx_mask] = self.attn1( n[uc_idx_mask], - context=n[uc_idx_mask], + context=context_attn1[uc_idx_mask], value=value_attn1[uc_idx_mask] if value_attn1 is not None else value_attn1) n = style_fidelity * n_c + (1.0-style_fidelity) * n_uc bank_styles.clean() From 7fbd553c666340eb5088f6b2216373f7b4d9b0f4 Mon Sep 17 00:00:00 2001 From: Harel Cain Date: Sun, 3 Mar 2024 13:37:26 +0200 Subject: [PATCH 05/17] Fix break caused by change to ComfyUI --- adv_control/control_svd.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/adv_control/control_svd.py b/adv_control/control_svd.py index 6a7de04..9458f6a 100644 --- a/adv_control/control_svd.py +++ b/adv_control/control_svd.py @@ -305,7 +305,7 @@ class SVDControlNet(nn.Module): cond = kwargs["cond"] num_video_frames = cond["num_video_frames"] - image_only_indicator = cond["image_only_indicator"] + image_only_indicator = cond.get("image_only_indicator", None) time_context = cond.get("time_context", None) del cond From dc3f773bbfb1957414f3dd076f10aa4061440708 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sun, 3 Mar 2024 08:41:55 -0600 Subject: [PATCH 06/17] Revert "Fix break caused by change to ComfyUI" --- adv_control/control_svd.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/adv_control/control_svd.py b/adv_control/control_svd.py index 9458f6a..6a7de04 100644 --- a/adv_control/control_svd.py +++ b/adv_control/control_svd.py @@ -305,7 +305,7 @@ class SVDControlNet(nn.Module): cond = kwargs["cond"] num_video_frames = cond["num_video_frames"] - image_only_indicator = cond.get("image_only_indicator", None) + image_only_indicator = cond["image_only_indicator"] time_context = cond.get("time_context", None) del cond From 3acf378866681d2de8dcd12b790703cb42de98e0 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Thu, 7 Mar 2024 03:52:54 -0600 Subject: [PATCH 07/17] Started debugging ref cn --- adv_control/control_reference.py | 22 +++++++++++++++------- 1 file changed, 15 insertions(+), 7 deletions(-) diff --git a/adv_control/control_reference.py b/adv_control/control_reference.py index 0b500c1..75c130e 100644 --- a/adv_control/control_reference.py +++ b/adv_control/control_reference.py @@ -18,6 +18,8 @@ from .utils import (AdvancedControlBase, ControlWeights, TimestepKeyframeGroup, REF_CONTROL_LIST = "ref_control_list" REF_CONTROL_INFO = "ref_control_info" REF_MACHINE_STATE = "ref_machine_state" +REF_COND_IDXS = "ref_cond_idxs" +REF_UNCOND_IDXS = "ref_uncond_idxs" class MachineState: @@ -334,6 +336,13 @@ def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections): if len(ref_controlnets) == 0: return reference_injections.diffusion_model_orig_forward(x, *args, **kwargs) try: + # assign cond and uncond idxs + per_batch = x.shape[0] // len(transformer_options["cond_or_uncond"]) + indiv_conds = [] + for cond_type in transformer_options["cond_or_uncond"]: + indiv_conds.extend([cond_type] * per_batch) + transformer_options[REF_UNCOND_IDXS] = [i for i, x in enumerate(indiv_conds) if x == 1] + transformer_options[REF_COND_IDXS] = [i for i, x in enumerate(indiv_conds) if x == 0] # otherwise, need to handle ref controlnet stuff for control in ref_controlnets: transformer_options[REF_MACHINE_STATE] = MachineState.WRITE @@ -342,6 +351,8 @@ def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections): # with open(Path(__file__).parent.parent.parent.parent.parent / "ref_debug" / "ref_xt_noised.pt", "rb") as rfile: # ref_xt = torch.load(rfile, weights_only=True) # diffuse cond_hint + + reference_injections.diffusion_model_orig_forward(control.cond_hint.to(dtype=x.dtype).to(device=x.device), *args, **kwargs) transformer_options[REF_MACHINE_STATE] = MachineState.READ transformer_options[REF_CONTROL_LIST] = ref_controlnets @@ -389,6 +400,8 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten value_attn1 = None # Reference CN stuff + uc_idx_mask = transformer_options[REF_UNCOND_IDXS] + c_idx_mask = transformer_options[REF_COND_IDXS] # WRITE mode will only have one ReferenceAdvanced, other modes will have all ReferenceAdvanced ref_controlnets: list[ReferenceAdvanced] = transformer_options.get(REF_CONTROL_LIST, None) ref_machine_state: str = transformer_options.get(REF_MACHINE_STATE, None) @@ -405,12 +418,6 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten # raw_val = torch.load(rfile) # raw_val[0] = raw_val[0].to(n.dtype).to(n.device) # bank_style.bank.extend(raw_val) - # create uc_idx_mask - per_batch = x.shape[0] // len(transformer_options["cond_or_uncond"]) - indiv_conds = [] - for cond_type in transformer_options["cond_or_uncond"]: - indiv_conds.extend([cond_type] * per_batch) - uc_idx_mask = [i for i, x in enumerate(indiv_conds) if x == 1] if "attn1_patch" in transformer_patches: patch = transformer_patches["attn1_patch"] @@ -466,7 +473,8 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten style_fidelity = bank_styles.get_avg_style_fidelity() n_uc: Tensor = self.attn1( n, - context=torch.cat([context_attn1] + bank_styles.bank, dim=1), + #context=torch.cat([context_attn1] + bank_styles.bank, dim=1), + context=torch.cat(bank_styles.bank + [context_attn1], dim=1), #context=torch.cat(bank_styles.bank, dim=1), value=torch.cat([value_attn1] + bank_styles.bank, dim=1) if value_attn1 is not None else value_attn1) n_c = n_uc.clone() From 7678232a80b3b83909555e65365ddaf45f405907 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 19 Mar 2024 10:38:52 -0500 Subject: [PATCH 08/17] More refcn debugging and fix for refcn issue for ancestral samplers while testing --- adv_control/control_reference.py | 13 +++++++++---- adv_control/utils.py | 12 ++++++++---- 2 files changed, 17 insertions(+), 8 deletions(-) diff --git a/adv_control/control_reference.py b/adv_control/control_reference.py index 75c130e..9290362 100644 --- a/adv_control/control_reference.py +++ b/adv_control/control_reference.py @@ -313,6 +313,9 @@ class InjectMP: def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections): + def forward_inject_UNetModel_test(self, x: Tensor, *args, **kwargs): + return reference_injections.diffusion_model_orig_forward(x, *args, **kwargs) + def forward_inject_UNetModel(self, x: Tensor, *args, **kwargs): # get control and transformer_options from kwargs real_args = list(args) @@ -354,6 +357,7 @@ def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections): reference_injections.diffusion_model_orig_forward(control.cond_hint.to(dtype=x.dtype).to(device=x.device), *args, **kwargs) + #reference_injections.diffusion_model_orig_forward(x, *args, **kwargs) transformer_options[REF_MACHINE_STATE] = MachineState.READ transformer_options[REF_CONTROL_LIST] = ref_controlnets return reference_injections.diffusion_model_orig_forward(x, *args, **kwargs) @@ -361,6 +365,7 @@ def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections): # make sure banks are cleared no matter what happens - otherwise, RIP VRAM reference_injections.clean_module_mem() + #return forward_inject_UNetModel_test return forward_inject_UNetModel @@ -400,8 +405,8 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten value_attn1 = None # Reference CN stuff - uc_idx_mask = transformer_options[REF_UNCOND_IDXS] - c_idx_mask = transformer_options[REF_COND_IDXS] + uc_idx_mask = transformer_options.get(REF_UNCOND_IDXS, []) + c_idx_mask = transformer_options.get(REF_COND_IDXS, []) # WRITE mode will only have one ReferenceAdvanced, other modes will have all ReferenceAdvanced ref_controlnets: list[ReferenceAdvanced] = transformer_options.get(REF_CONTROL_LIST, None) ref_machine_state: str = transformer_options.get(REF_MACHINE_STATE, None) @@ -473,8 +478,8 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten style_fidelity = bank_styles.get_avg_style_fidelity() n_uc: Tensor = self.attn1( n, - #context=torch.cat([context_attn1] + bank_styles.bank, dim=1), - context=torch.cat(bank_styles.bank + [context_attn1], dim=1), + context=torch.cat([context_attn1] + bank_styles.bank, dim=1), + #context=torch.cat(bank_styles.bank + [context_attn1], dim=1), #context=torch.cat(bank_styles.bank, dim=1), value=torch.cat([value_attn1] + bank_styles.bank, dim=1) if value_attn1 is not None else value_attn1) n_c = n_uc.clone() diff --git a/adv_control/utils.py b/adv_control/utils.py index 6bc0d22..9d6cce8 100644 --- a/adv_control/utils.py +++ b/adv_control/utils.py @@ -322,13 +322,17 @@ def ddpm_noise_latents(latents: Tensor, sigma: float, noise: Tensor=None): sqrt_one_minus_alpha_prod = (1 - alpha_cumprod) ** 0.5 #logger.warn(f"sqrt: {sqrt_alpha_prod}, sqrt-1: {sqrt_one_minus_alpha_prod}, t: {sigma}") if noise is None: - # generator = torch.manual_seed(0) - # noise = torch.randn(latents.size(), generator=generator).to(latents.device) - generator = torch.cuda.manual_seed(0) - noise = torch.empty_like(latents).normal_(generator=generator) + #noise = torch.randn(latents.size()).to(latents.device) + #generator = torch.manual_seed(0) + generator = torch.Generator(device="cuda") + generator.manual_seed(0) + #noise = torch.randn(latents.size(), generator=generator).to(latents.device) + #generator = torch.cuda.manual_seed(0) + noise = torch.empty_like(latents).normal_(generator=generator).to(latents.device) #noise = torch.empty(latents.size()).normal_(generator=generator).to(latents.device) #return noise #noise = torch.rand_like(latents) + #return None #return latents return sqrt_alpha_prod * latents + sqrt_one_minus_alpha_prod * noise From 111632917c7dbdcb0074e4ad55c90561135003ae Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sat, 30 Mar 2024 04:02:44 -0500 Subject: [PATCH 09/17] Cleaned up ref cn so I can add more features --- adv_control/control_reference.py | 29 ++++------------------------- adv_control/utils.py | 24 +++++++++--------------- 2 files changed, 13 insertions(+), 40 deletions(-) diff --git a/adv_control/control_reference.py b/adv_control/control_reference.py index 9290362..25f98ac 100644 --- a/adv_control/control_reference.py +++ b/adv_control/control_reference.py @@ -155,16 +155,9 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): 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) # noise cond_hint based on sigma (current step) - # TODO: how to handle noise? reproducibility is key... - # mess with the order here? - # / (self.latent_format.scale_factor) self.cond_hint = self.latent_format.process_in(self.cond_hint) - #self.cond_hint = self.model_sampling_current.calculate_input(t, self.cond_hint) - self.cond_hint = ddpm_noise_latents(self.cond_hint, sigma=t[0], noise=None) + self.cond_hint = ddpm_noise_latents(self.cond_hint, sigma=t, noise=None) timestep = self.model_sampling_current.timestep(t) - #self.cond_hint = ddpm_noise_latents(torch.zeros_like(x_noisy), sigma=t[0], noise=None) - #self.cond_hint = simple_noise_latents(self.cond_hint, sigma=t[0], noise=None) - # prepare mask self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number) # done preparing; model patches will take care of everything now. @@ -313,9 +306,6 @@ class InjectMP: def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections): - def forward_inject_UNetModel_test(self, x: Tensor, *args, **kwargs): - return reference_injections.diffusion_model_orig_forward(x, *args, **kwargs) - def forward_inject_UNetModel(self, x: Tensor, *args, **kwargs): # get control and transformer_options from kwargs real_args = list(args) @@ -350,12 +340,8 @@ def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections): for control in ref_controlnets: transformer_options[REF_MACHINE_STATE] = MachineState.WRITE transformer_options[REF_CONTROL_LIST] = [control] - # from pathlib import Path - # with open(Path(__file__).parent.parent.parent.parent.parent / "ref_debug" / "ref_xt_noised.pt", "rb") as rfile: - # ref_xt = torch.load(rfile, weights_only=True) - # diffuse cond_hint - + # TODO: handle masks - apply x to locations where masked out reference_injections.diffusion_model_orig_forward(control.cond_hint.to(dtype=x.dtype).to(device=x.device), *args, **kwargs) #reference_injections.diffusion_model_orig_forward(x, *args, **kwargs) transformer_options[REF_MACHINE_STATE] = MachineState.READ @@ -365,7 +351,6 @@ def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections): # make sure banks are cleared no matter what happens - otherwise, RIP VRAM reference_injections.clean_module_mem() - #return forward_inject_UNetModel_test return forward_inject_UNetModel @@ -418,11 +403,6 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten bank_style = self.injection_holder.bank_styles[ref_controlnets[0].order] bank_style.bank.append(n.detach().clone()) bank_style.style_cfgs.append(ref_controlnets[0].ref_opts.style_fidelity) - # from pathlib import Path - # with open(Path(__file__).parent.parent.parent.parent.parent / "ref_debug" / f"bank_{self.injection_holder.idx}.pt", "rb") as rfile: - # raw_val = torch.load(rfile) - # raw_val[0] = raw_val[0].to(n.dtype).to(n.device) - # bank_style.bank.extend(raw_val) if "attn1_patch" in transformer_patches: patch = transformer_patches["attn1_patch"] @@ -447,6 +427,7 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten value_attn1 = n n = self.attn1.to_q(n) # Reference CN READ - use attn1_replace_patch appropriately + # TODO: test this with a dummy attn1_replace_patch if ref_machine_state == MachineState.READ and self.injection_holder.bank_styles.get(ref_controlnets[0].order, None) is not None: bank_styles = self.injection_holder.bank_styles[ref_controlnets[0].order] style_fidelity = bank_styles.get_avg_style_fidelity() @@ -479,11 +460,9 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten n_uc: Tensor = self.attn1( n, context=torch.cat([context_attn1] + bank_styles.bank, dim=1), - #context=torch.cat(bank_styles.bank + [context_attn1], dim=1), - #context=torch.cat(bank_styles.bank, dim=1), value=torch.cat([value_attn1] + bank_styles.bank, dim=1) if value_attn1 is not None else value_attn1) n_c = n_uc.clone() - if len(uc_idx_mask) > 0 and style_fidelity > 1e-5:# not math.isclose(style_fidelity, 0.0): + if len(uc_idx_mask) > 0 and not math.isclose(style_fidelity, 0.0): n_c[uc_idx_mask] = self.attn1( n[uc_idx_mask], context=context_attn1[uc_idx_mask], diff --git a/adv_control/utils.py b/adv_control/utils.py index 9d6cce8..d153c43 100644 --- a/adv_control/utils.py +++ b/adv_control/utils.py @@ -316,24 +316,18 @@ def broadcast_image_to_full(tensor, target_batch_size, batched_number, except_on return torch.cat([tensor] * batched_number, dim=0) -def ddpm_noise_latents(latents: Tensor, sigma: float, noise: Tensor=None): +def ddpm_noise_latents(latents: Tensor, sigma: Tensor, noise: Tensor=None): + sigma = sigma.unsqueeze(-1).unsqueeze(-1).unsqueeze(-1) alpha_cumprod = 1 / ((sigma * sigma) + 1) sqrt_alpha_prod = alpha_cumprod ** 0.5 - sqrt_one_minus_alpha_prod = (1 - alpha_cumprod) ** 0.5 - #logger.warn(f"sqrt: {sqrt_alpha_prod}, sqrt-1: {sqrt_one_minus_alpha_prod}, t: {sigma}") + sqrt_one_minus_alpha_prod = (1. - alpha_cumprod) ** 0.5 if noise is None: - #noise = torch.randn(latents.size()).to(latents.device) - #generator = torch.manual_seed(0) - generator = torch.Generator(device="cuda") - generator.manual_seed(0) - #noise = torch.randn(latents.size(), generator=generator).to(latents.device) - #generator = torch.cuda.manual_seed(0) - noise = torch.empty_like(latents).normal_(generator=generator).to(latents.device) - #noise = torch.empty(latents.size()).normal_(generator=generator).to(latents.device) - #return noise - #noise = torch.rand_like(latents) - #return None - #return latents + # generator = torch.Generator(device="cuda") + # generator.manual_seed(0) + # generator = torch.Generator() + # generator.manual_seed(0) + # noise = torch.randn(latents.size(), generator=generator).to(latents.device) + noise = torch.randn_like(latents).to(latents.device) return sqrt_alpha_prod * latents + sqrt_one_minus_alpha_prod * noise From 25892b7b7dd85cb7cc67c5ff33321a8b869de5aa Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sat, 30 Mar 2024 05:26:09 -0500 Subject: [PATCH 10/17] Added proper strength control to ReferenceCN, replaced old strength with ref_weight --- adv_control/control_reference.py | 60 ++++++++++++++++++++------------ adv_control/nodes_reference.py | 7 ++-- 2 files changed, 42 insertions(+), 25 deletions(-) diff --git a/adv_control/control_reference.py b/adv_control/control_reference.py index 25f98ac..bc5d9fc 100644 --- a/adv_control/control_reference.py +++ b/adv_control/control_reference.py @@ -39,13 +39,14 @@ class ReferenceType: class ReferenceOptions: - def __init__(self, reference_type: str, style_fidelity: float): + def __init__(self, reference_type: str, style_fidelity: float, ref_weight: float): self.reference_type = reference_type self.original_style_fidelity = style_fidelity self.style_fidelity = style_fidelity + self.ref_weight = ref_weight def clone(self): - return ReferenceOptions(reference_type=self.reference_type, style_fidelity=self.original_style_fidelity) + return ReferenceOptions(reference_type=self.reference_type, style_fidelity=self.original_style_fidelity, ref_weight=self.ref_weight) class ReferencePreprocWrapper(AbstractPreprocWrapper): @@ -96,6 +97,7 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): self.order = 0 self.latent_format = None self.model_sampling_current = None + self.should_apply_effective_strength = True def get_effective_strength(self): effective_strength = self.strength @@ -103,6 +105,9 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): effective_strength = effective_strength * self.current_timestep_keyframe.strength return effective_strength + #def should_apply_effective_strength(self): + # return not (math.isclose(self.strength, 1.0) and math.is_close(self.current_timestep_keyframe.strength, 1.0)) + def patch_model(self, model: ModelPatcher): # need to patch model so that control can be found later from it model.set_model_attn1_output_patch(self.patch_attn1) @@ -158,6 +163,7 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): self.cond_hint = self.latent_format.process_in(self.cond_hint) self.cond_hint = ddpm_noise_latents(self.cond_hint, sigma=t, noise=None) timestep = self.model_sampling_current.timestep(t) + self.should_apply_effective_strength = not (math.isclose(self.strength, 1.0) and math.isclose(self.current_timestep_keyframe.strength, 1.0)) # prepare mask self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number) # done preparing; model patches will take care of everything now. @@ -184,6 +190,7 @@ class BankStylesBasicTransformerBlock: def __init__(self): self.bank = [] self.style_cfgs = [] + self.cn_idx: list[int] = [] def get_avg_style_fidelity(self): return sum(self.style_cfgs) / float(len(self.style_cfgs)) @@ -193,6 +200,8 @@ class BankStylesBasicTransformerBlock: self.bank = [] del self.style_cfgs self.style_cfgs = [] + del self.cn_idx + self.cn_idx = [] class InjectionBasicTransformerBlockHolder: @@ -200,15 +209,13 @@ class InjectionBasicTransformerBlockHolder: self.original_forward = block._forward self.idx = idx self.attn_weight = 1.0 - self.bank_styles: dict[int, BankStylesBasicTransformerBlock] = {} + self.bank_styles = BankStylesBasicTransformerBlock() def restore(self, block: BasicTransformerBlock): block._forward = self.original_forward def clean(self): - for bank_style in list(self.bank_styles.values()): - bank_style.clean() - self.bank_styles.clear() + self.bank_styles.clean() # inject ModelPatcher.patch_model to apply @@ -330,7 +337,8 @@ def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections): return reference_injections.diffusion_model_orig_forward(x, *args, **kwargs) try: # assign cond and uncond idxs - per_batch = x.shape[0] // len(transformer_options["cond_or_uncond"]) + batched_number = len(transformer_options["cond_or_uncond"]) + per_batch = x.shape[0] // batched_number indiv_conds = [] for cond_type in transformer_options["cond_or_uncond"]: indiv_conds.extend([cond_type] * per_batch) @@ -341,7 +349,11 @@ def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections): transformer_options[REF_MACHINE_STATE] = MachineState.WRITE transformer_options[REF_CONTROL_LIST] = [control] - # TODO: handle masks - apply x to locations where masked out + # handle masks - apply x to unmasked + #strength_mask = torch.ones_like(x, dtype=x.dtype) * control.strength + #control.apply_advanced_strengths_and_masks(x=strength_mask, batched_number=batched_number) + #real_cond_hint = control.cond_hint * strength_mask + x * (1 - strength_mask) + reference_injections.diffusion_model_orig_forward(control.cond_hint.to(dtype=x.dtype).to(device=x.device), *args, **kwargs) #reference_injections.diffusion_model_orig_forward(x, *args, **kwargs) transformer_options[REF_MACHINE_STATE] = MachineState.READ @@ -397,12 +409,10 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten ref_machine_state: str = transformer_options.get(REF_MACHINE_STATE, None) # if in WRITE mode, save n and style_fidelity if ref_controlnets and ref_machine_state == MachineState.WRITE: - if ref_controlnets[0].get_effective_strength() > self.injection_holder.attn_weight: - if not self.injection_holder.bank_styles.get(ref_controlnets[0].order, None): - self.injection_holder.bank_styles[ref_controlnets[0].order] = BankStylesBasicTransformerBlock() - bank_style = self.injection_holder.bank_styles[ref_controlnets[0].order] - bank_style.bank.append(n.detach().clone()) - bank_style.style_cfgs.append(ref_controlnets[0].ref_opts.style_fidelity) + if ref_controlnets[0].ref_opts.ref_weight > self.injection_holder.attn_weight: + self.injection_holder.bank_styles.bank.append(n.detach().clone()) + self.injection_holder.bank_styles.style_cfgs.append(ref_controlnets[0].ref_opts.style_fidelity) + self.injection_holder.bank_styles.cn_idx.append(ref_controlnets[0].order) if "attn1_patch" in transformer_patches: patch = transformer_patches["attn1_patch"] @@ -428,13 +438,14 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten n = self.attn1.to_q(n) # Reference CN READ - use attn1_replace_patch appropriately # TODO: test this with a dummy attn1_replace_patch - if ref_machine_state == MachineState.READ and self.injection_holder.bank_styles.get(ref_controlnets[0].order, None) is not None: - bank_styles = self.injection_holder.bank_styles[ref_controlnets[0].order] + if ref_machine_state == MachineState.READ and len(self.injection_holder.bank_styles.bank) > 0: + bank_styles = self.injection_holder.bank_styles style_fidelity = bank_styles.get_avg_style_fidelity() + real_bank = bank_styles.bank.copy() n_uc = self.attn1.to_out(attn1_replace_patch[block_attn1]( n, - self.attn1.to_k(torch.cat([context_attn1] + bank_styles.bank, dim=1)), - self.attn1.to_v(torch.cat([value_attn1] + bank_styles.bank, dim=1)), + self.attn1.to_k(torch.cat([context_attn1] + real_bank, dim=1)), + self.attn1.to_v(torch.cat([value_attn1] + real_bank, dim=1)), extra_options)) n_c = n_uc.clone() if len(uc_idx_mask) > 0 and not math.isclose(style_fidelity, 0.0): @@ -452,15 +463,20 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten n = self.attn1.to_out(n) else: # Reference CN READ - no attn1_replace_patch - if ref_machine_state == MachineState.READ and self.injection_holder.bank_styles.get(ref_controlnets[0].order, None) is not None: + if ref_machine_state == MachineState.READ and len(self.injection_holder.bank_styles.bank) > 0: if context_attn1 is None: context_attn1 = n - bank_styles = self.injection_holder.bank_styles[ref_controlnets[0].order] + bank_styles = self.injection_holder.bank_styles style_fidelity = bank_styles.get_avg_style_fidelity() + real_bank = bank_styles.bank.copy() + for idx, order in enumerate(bank_styles.cn_idx): + if ref_controlnets[idx].should_apply_effective_strength: + effective_strength = ref_controlnets[idx].get_effective_strength() + real_bank[idx] = real_bank[idx] * effective_strength + context_attn1 * (1-effective_strength) n_uc: Tensor = self.attn1( n, - context=torch.cat([context_attn1] + bank_styles.bank, dim=1), - value=torch.cat([value_attn1] + bank_styles.bank, dim=1) if value_attn1 is not None else value_attn1) + context=torch.cat([context_attn1] + real_bank, dim=1), + value=torch.cat([value_attn1] + real_bank, dim=1) if value_attn1 is not None else value_attn1) n_c = n_uc.clone() if len(uc_idx_mask) > 0 and not math.isclose(style_fidelity, 0.0): n_c[uc_idx_mask] = self.attn1( diff --git a/adv_control/nodes_reference.py b/adv_control/nodes_reference.py index 06236d2..f003f43 100644 --- a/adv_control/nodes_reference.py +++ b/adv_control/nodes_reference.py @@ -14,7 +14,8 @@ class ReferenceControlNetNode: return { "required": { "reference_type": (ReferenceType._LIST,), - "style_fidelity": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}) + "style_fidelity": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}), + "ref_weight": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), }, } @@ -23,8 +24,8 @@ class ReferenceControlNetNode: CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/Reference" - def load_controlnet(self, reference_type: str, style_fidelity: float): - ref_opts = ReferenceOptions(reference_type=reference_type, style_fidelity=style_fidelity) + def load_controlnet(self, reference_type: str, style_fidelity: float, ref_weight: float): + ref_opts = ReferenceOptions(reference_type=reference_type, style_fidelity=style_fidelity, ref_weight=ref_weight) ref_patch = ReferenceAttnPatch() controlnet = ReferenceAdvanced(patch_attn1=ref_patch, ref_opts=ref_opts, timestep_keyframes=None) return (controlnet,) From 37ae77b486f50c49ddd40bade1dff36b2769495c Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sat, 30 Mar 2024 05:29:22 -0500 Subject: [PATCH 11/17] Applied proper strength control to attn1_replace_patch as well --- adv_control/control_reference.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/adv_control/control_reference.py b/adv_control/control_reference.py index bc5d9fc..cadba8f 100644 --- a/adv_control/control_reference.py +++ b/adv_control/control_reference.py @@ -442,6 +442,10 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten bank_styles = self.injection_holder.bank_styles style_fidelity = bank_styles.get_avg_style_fidelity() real_bank = bank_styles.bank.copy() + for idx, order in enumerate(bank_styles.cn_idx): + if ref_controlnets[idx].should_apply_effective_strength: + effective_strength = ref_controlnets[idx].get_effective_strength() + real_bank[idx] = real_bank[idx] * effective_strength + context_attn1 * (1-effective_strength) n_uc = self.attn1.to_out(attn1_replace_patch[block_attn1]( n, self.attn1.to_k(torch.cat([context_attn1] + real_bank, dim=1)), From 7696628fdd4cfff660670a4414643b30109073fd Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sat, 30 Mar 2024 19:16:21 -0500 Subject: [PATCH 12/17] Simplified RefCN injection, model_opt no longer needed and should fix some intermittent bugs --- adv_control/control.py | 1 - adv_control/control_reference.py | 239 +++++++++++++------------------ adv_control/nodes_reference.py | 5 +- 3 files changed, 98 insertions(+), 147 deletions(-) diff --git a/adv_control/control.py b/adv_control/control.py index 8989efc..0226793 100644 --- a/adv_control/control.py +++ b/adv_control/control.py @@ -12,7 +12,6 @@ from comfy.model_patcher import ModelPatcher from .control_sparsectrl import SparseControlNet, SparseCtrlMotionWrapper, SparseMethod, SparseSettings, SparseSpreadMethod, PreprocSparseRGBWrapper from .control_lllite import LLLiteModule, LLLitePatch -from .control_reference import MachineState, ReferenceOptions, ReferenceAttnPatch 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) diff --git a/adv_control/control_reference.py b/adv_control/control_reference.py index cadba8f..c51d75f 100644 --- a/adv_control/control_reference.py +++ b/adv_control/control_reference.py @@ -4,6 +4,7 @@ import math import torch from torch import Tensor +import comfy.sample import comfy.model_patcher import comfy.utils from comfy.controlnet import ControlBase @@ -15,7 +16,89 @@ from .utils import (AdvancedControlBase, ControlWeights, TimestepKeyframeGroup, deepcopy_with_sharing, prepare_mask_batch, broadcast_image_to_full, ddpm_noise_latents, simple_noise_latents) +def refcn_sample_factory(orig_comfy_sample: Callable, is_custom=False) -> Callable: + def get_refcn(control: ControlBase, order: int=-1): + ref_set: set[ReferenceAdvanced] = set() + if control is None: + return ref_set + if type(control) == ReferenceAdvanced: + control.order = order + order -= 1 + ref_set.add(control) + ref_set.update(get_refcn(control.previous_controlnet, order=order)) + return ref_set + + def refcn_sample(model: ModelPatcher, *args, **kwargs): + # check if positive or negative conds contain ref cn + positive = args[-3] + negative = args[-2] + ref_set = set() + if positive is not None: + for cond in positive: + if "control" in cond[1]: + ref_set.update(get_refcn(cond[1]["control"])) + if negative is not None: + for cond in negative: + if "control" in cond[1]: + ref_set.update(get_refcn(cond[1]["control"])) + # if no ref cn found, do original function immediately + if len(ref_set) == 0: + return orig_comfy_sample(model, *args, **kwargs) + # otherwise, injection time + try: + # inject + # storage for all Reference-related injections + reference_injections = ReferenceInjections() + # first, handle attn module injection + all_modules = torch_dfs(model.model) + attn_modules: list[RefBasicTransformerBlock] = [] + for module in all_modules: + if isinstance(module, BasicTransformerBlock): + attn_modules.append(module) + attn_modules = [module for module in all_modules if isinstance(module, BasicTransformerBlock)] + attn_modules = sorted(attn_modules, key=lambda x: -x.norm1.normalized_shape[0]) + reference_injections.attn_modules = [] + for i, module in enumerate(attn_modules): + injection_holder = InjectionBasicTransformerBlockHolder(block=module, idx=i) + injection_holder.attn_weight = float(i) / float(len(attn_modules)) + module._forward = _forward_inject_BasicTransformerBlock.__get__(module, type(module)) + module.injection_holder = injection_holder + reference_injections.attn_modules.append(module) + # handle diffusion_model forward injection + reference_injections.diffusion_model_orig_forward = model.model.diffusion_model.forward + model.model.diffusion_model.forward = factory_forward_inject_UNetModel(reference_injections).__get__(model.model.diffusion_model, type(model.model.diffusion_model)) + # store ordered ref cns in model's transformer options + orig_model_options = model.model_options + new_model_options = model.model_options.copy() + new_model_options["transformer_options"] = model.model_options["transformer_options"].copy() + ref_list: list[ReferenceAdvanced] = list(ref_set) + new_model_options["transformer_options"][REF_CONTROL_LIST_ALL] = sorted(ref_list, key=lambda x: x.order) + model.model_options = new_model_options + # continue with original function + return orig_comfy_sample(model, *args, **kwargs) + finally: + # cleanup injections + # first, restore attn modules + attn_modules: list[RefBasicTransformerBlock] = reference_injections.attn_modules + for module in attn_modules: + module.injection_holder.restore(module) + module.injection_holder.clean() + del module.injection_holder + del attn_modules + # restore diffusion_model forward function + model.model.diffusion_model.forward = reference_injections.diffusion_model_orig_forward.__get__(model.model.diffusion_model, type(model.model.diffusion_model)) + # restore model_options + model.model_options = orig_model_options + # cleanup + reference_injections.cleanup() + return refcn_sample +# inject sample functions +comfy.sample.sample = refcn_sample_factory(comfy.sample.sample) +comfy.sample.sample_custom = refcn_sample_factory(comfy.sample.sample_custom, is_custom=True) + + REF_CONTROL_LIST = "ref_control_list" +REF_CONTROL_LIST_ALL = "ref_control_list_all" REF_CONTROL_INFO = "ref_control_info" REF_MACHINE_STATE = "ref_machine_state" REF_COND_IDXS = "ref_cond_idxs" @@ -55,49 +138,16 @@ class ReferencePreprocWrapper(AbstractPreprocWrapper): super().__init__(condhint) -class ReferenceAttnPatch: - def __init__(self, control: 'ReferenceAdvanced'=None): - self.control = control - - # def __call__(self, q: Tensor, k: Tensor, v: Tensor, extra_options: dict): - # # do nothing here - all ReferenceAttnPatch is trying to do is be a - # # ComfyUI-compliant way of tracking the corresponding ControlNet obj - # return q, k, v - - def __call__(self, x: Tensor, extra_options: dict): - # do nothing here - all ReferenceAttnPatch is trying to do is be a - # ComfyUI-compliant way of tracking the corresponding ControlNet obj - return x - - def set_control(self, control: 'ReferenceAdvanced') -> 'ReferenceAttnPatch': - self.control = control - return self - - def cleanup(self): - pass - - # make sure deepcopy does not copy control, and deepcopied patch should be assigned to control - def __deepcopy__(self, memo): - self.cleanup() - to_return: ReferenceAttnPatch = deepcopy_with_sharing(self, shared_attribute_names = ['control'], memo=memo) - #logger.warn(f"patch {id(self)} turned into {id(to_return)}") - try: - to_return.control.patch_attn1 = to_return - except Exception: - pass - return to_return - - class ReferenceAdvanced(ControlBase, AdvancedControlBase): - def __init__(self, patch_attn1: ReferenceAttnPatch, ref_opts: ReferenceOptions, timestep_keyframes: TimestepKeyframeGroup, device=None): + def __init__(self, ref_opts: ReferenceOptions, timestep_keyframes: TimestepKeyframeGroup, device=None): super().__init__(device) - AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controllllite(), require_model=True) - self.patch_attn1 = patch_attn1.set_control(self) + AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controllllite()) self.ref_opts = ref_opts self.order = 0 self.latent_format = None self.model_sampling_current = None self.should_apply_effective_strength = True + self.ref_latent_keyframe_mults = None def get_effective_strength(self): effective_strength = self.strength @@ -105,19 +155,6 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): effective_strength = effective_strength * self.current_timestep_keyframe.strength return effective_strength - #def should_apply_effective_strength(self): - # return not (math.isclose(self.strength, 1.0) and math.is_close(self.current_timestep_keyframe.strength, 1.0)) - - def patch_model(self, model: ModelPatcher): - # need to patch model so that control can be found later from it - model.set_model_attn1_output_patch(self.patch_attn1) - #model.set_model_attn1_patch(self.patch_attn1) - # need to add model_options to make patch/unpatch injection know it has to run - if not REF_CONTROL_INFO in model.model_options: - model.model_options[REF_CONTROL_INFO] = 0 - self.order = model.model_options[REF_CONTROL_INFO] - model.model_options[REF_CONTROL_INFO] - def pre_run_advanced(self, model, percent_to_timestep_function): AdvancedControlBase.pre_run_advanced(self, model, percent_to_timestep_function) if type(self.cond_hint_original) == ReferencePreprocWrapper: @@ -129,8 +166,6 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): self.ref_opts.style_fidelity = self.ref_opts.style_fidelity ** 3.0 else: self.ref_opts.style_fidelity = self.ref_opts.style_fidelity - # set control on patches - self.patch_attn1.set_control(self) def get_control_advanced(self, x_noisy: Tensor, t, cond, batched_number: int): # normal ControlNet stuff @@ -164,6 +199,7 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): self.cond_hint = ddpm_noise_latents(self.cond_hint, sigma=t, noise=None) timestep = self.model_sampling_current.timestep(t) self.should_apply_effective_strength = not (math.isclose(self.strength, 1.0) and math.isclose(self.current_timestep_keyframe.strength, 1.0)) + self.ref_latent_keyframe_mults = self.calc_latent_keyframe_mults(x=x_noisy, batched_number=batched_number) # prepare mask self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number) # done preparing; model patches will take care of everything now. @@ -172,19 +208,23 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): def cleanup_advanced(self): super().cleanup_advanced() - self.patch_attn1.cleanup() del self.latent_format self.latent_format = None del self.model_sampling_current self.model_sampling_current = None def copy(self): - c = ReferenceAdvanced(self.patch_attn1, self.ref_opts, self.timestep_keyframes) + c = ReferenceAdvanced(self.ref_opts, self.timestep_keyframes) c.order = self.order self.copy_to(c) self.copy_to_advanced(c) return c + # avoid deepcopy shenanigans by making deepcopy not do anything to the reference + # TODO: do the bookkeeping to do this in a proper way for all Adv-ControlNets + def __deepcopy__(self, memo): + return self + class BankStylesBasicTransformerBlock: def __init__(self): @@ -218,60 +258,6 @@ class InjectionBasicTransformerBlockHolder: self.bank_styles.clean() -# inject ModelPatcher.patch_model to apply -orig_modelpatcher_patch_model = comfy.model_patcher.ModelPatcher.patch_model -def patch_model_injection_ref(self: ModelPatcher, *args, **kwargs): - if REF_CONTROL_INFO in self.model_options: - # storage for all Reference-related injections - reference_injections = ReferenceInjections() - # first, handle attn module injection - all_modules = torch_dfs(self.model) - attn_modules: list[RefBasicTransformerBlock] = [] - for module in all_modules: - if isinstance(module, BasicTransformerBlock): - attn_modules.append(module) - attn_modules = [module for module in all_modules if isinstance(module, BasicTransformerBlock)] - attn_modules = sorted(attn_modules, key=lambda x: -x.norm1.normalized_shape[0]) - reference_injections.attn_modules = [] - for i, module in enumerate(attn_modules): - # if i != 11 and i != 12: - # continue - injection_holder = InjectionBasicTransformerBlockHolder(block=module, idx=i) - injection_holder.attn_weight = float(i) / float(len(attn_modules)) - module._forward = _forward_inject_BasicTransformerBlock.__get__(module, type(module)) - module.injection_holder = injection_holder - reference_injections.attn_modules.append(module) - - # handle diffusion_model forward injection - reference_injections.diffusion_model_orig_forward = self.model.diffusion_model.forward - self.model.diffusion_model.forward = factory_forward_inject_UNetModel(reference_injections).__get__(self.model.diffusion_model, type(self.model.diffusion_model)) - InjectMP.set_injected(self, reference_injections) - to_return = orig_modelpatcher_patch_model(self, *args, **kwargs) - return to_return -comfy.model_patcher.ModelPatcher.patch_model = patch_model_injection_ref - - -orig_modelpatcher_unpatch_model = comfy.model_patcher.ModelPatcher.unpatch_model -def unpatch_model_injection_ref(self: ModelPatcher, *args, **kwargs): - if REF_CONTROL_INFO in self.model_options: - reference_injections: ReferenceInjections = InjectMP.get_injected(self) - # first, restore attn modules - attn_modules: list[RefBasicTransformerBlock] = reference_injections.attn_modules - for module in attn_modules: - module.injection_holder.restore(module) - module.injection_holder.clean() - del module.injection_holder - del attn_modules - # restore diffusion_model forward function - self.model.diffusion_model.forward = reference_injections.diffusion_model_orig_forward.__get__(self.model.diffusion_model, type(self.model.diffusion_model)) - # cleanup - InjectMP.clean_injected(self) - reference_injections.cleanup() - to_return = orig_modelpatcher_unpatch_model(self, *args, **kwargs) - return to_return -comfy.model_patcher.ModelPatcher.unpatch_model = unpatch_model_injection_ref - - class ReferenceInjections: def __init__(self, attn_modules: list['RefBasicTransformerBlock']=None): self.attn_modules = attn_modules if attn_modules else [] @@ -291,27 +277,6 @@ class ReferenceInjections: self.diffusion_model_orig_forward = None -class InjectMP: - PARAM_INJECTED_REF = "___injected_ref" - - @staticmethod - def is_injected(model: ModelPatcher): - return getattr(model, InjectMP.PARAM_INJECTED_REF, False) - - def mark_injected(model: ModelPatcher): - setattr(model, InjectMP.PARAM_INJECTED_REF, True) - - def set_injected(model: ModelPatcher, value: ReferenceInjections): - setattr(model, InjectMP.PARAM_INJECTED_REF, value) - - def get_injected(model: ModelPatcher) -> ReferenceInjections: - return getattr(model, InjectMP.PARAM_INJECTED_REF) - - def clean_injected(model: ModelPatcher): - delattr(model, InjectMP.PARAM_INJECTED_REF) - InjectMP.mark_injected(model) - - def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections): def forward_inject_UNetModel(self, x: Tensor, *args, **kwargs): # get control and transformer_options from kwargs @@ -320,18 +285,9 @@ def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections): control = kwargs.get("control", None) transformer_options = kwargs.get("transformer_options", None) # look for ReferenceAttnPatch objects to get ReferenceAdvanced objects - patch_name = "attn1_output_patch" - ref_patches: list[ReferenceAttnPatch] = [] - if "patches" in transformer_options: - if patch_name in transformer_options["patches"]: - patches: list = transformer_options["patches"][patch_name] - for i in range(len(patches)): - if isinstance(patches[i], ReferenceAttnPatch): - ref_patches.append(patches[i]) - ref_controlnets: list[ReferenceAdvanced] = [x.control for x in ref_patches] + ref_controlnets: list[ReferenceAdvanced] = transformer_options[REF_CONTROL_LIST_ALL] # discard any controlnets that should not run ref_controlnets = [x for x in ref_controlnets if x.should_run()] - ref_controlnets = sorted(ref_controlnets, key=lambda x: x.order) # if nothing related to reference controlnets, do nothing special if len(ref_controlnets) == 0: return reference_injections.diffusion_model_orig_forward(x, *args, **kwargs) @@ -344,18 +300,16 @@ def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections): indiv_conds.extend([cond_type] * per_batch) transformer_options[REF_UNCOND_IDXS] = [i for i, x in enumerate(indiv_conds) if x == 1] transformer_options[REF_COND_IDXS] = [i for i, x in enumerate(indiv_conds) if x == 0] - # otherwise, need to handle ref controlnet stuff + # handle running diffusion with ref cond hints for control in ref_controlnets: transformer_options[REF_MACHINE_STATE] = MachineState.WRITE transformer_options[REF_CONTROL_LIST] = [control] - # handle masks - apply x to unmasked #strength_mask = torch.ones_like(x, dtype=x.dtype) * control.strength #control.apply_advanced_strengths_and_masks(x=strength_mask, batched_number=batched_number) #real_cond_hint = control.cond_hint * strength_mask + x * (1 - strength_mask) - reference_injections.diffusion_model_orig_forward(control.cond_hint.to(dtype=x.dtype).to(device=x.device), *args, **kwargs) - #reference_injections.diffusion_model_orig_forward(x, *args, **kwargs) + # run diffusion for real now transformer_options[REF_MACHINE_STATE] = MachineState.READ transformer_options[REF_CONTROL_LIST] = ref_controlnets return reference_injections.diffusion_model_orig_forward(x, *args, **kwargs) @@ -437,7 +391,6 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten value_attn1 = n n = self.attn1.to_q(n) # Reference CN READ - use attn1_replace_patch appropriately - # TODO: test this with a dummy attn1_replace_patch if ref_machine_state == MachineState.READ and len(self.injection_holder.bank_styles.bank) > 0: bank_styles = self.injection_holder.bank_styles style_fidelity = bank_styles.get_avg_style_fidelity() diff --git a/adv_control/nodes_reference.py b/adv_control/nodes_reference.py index f003f43..c4149c0 100644 --- a/adv_control/nodes_reference.py +++ b/adv_control/nodes_reference.py @@ -4,7 +4,7 @@ from nodes import VAEEncode import comfy.utils from comfy.sd import VAE -from .control_reference import ReferenceAdvanced, ReferenceAttnPatch, ReferenceOptions, ReferenceType, ReferencePreprocWrapper +from .control_reference import ReferenceAdvanced, ReferenceOptions, ReferenceType, ReferencePreprocWrapper # node for ReferenceCN @@ -26,8 +26,7 @@ class ReferenceControlNetNode: def load_controlnet(self, reference_type: str, style_fidelity: float, ref_weight: float): ref_opts = ReferenceOptions(reference_type=reference_type, style_fidelity=style_fidelity, ref_weight=ref_weight) - ref_patch = ReferenceAttnPatch() - controlnet = ReferenceAdvanced(patch_attn1=ref_patch, ref_opts=ref_opts, timestep_keyframes=None) + controlnet = ReferenceAdvanced(ref_opts=ref_opts, timestep_keyframes=None) return (controlnet,) From 2fdf8236f649b99a23e9a7f7e42603ca92491486 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sat, 30 Mar 2024 19:48:45 -0500 Subject: [PATCH 13/17] Don't apply normal CNs during ref cn WRITE runs --- adv_control/control_reference.py | 11 +++++++++-- adv_control/nodes_reference.py | 5 +++-- 2 files changed, 12 insertions(+), 4 deletions(-) diff --git a/adv_control/control_reference.py b/adv_control/control_reference.py index c51d75f..a042e19 100644 --- a/adv_control/control_reference.py +++ b/adv_control/control_reference.py @@ -122,14 +122,16 @@ class ReferenceType: class ReferenceOptions: - def __init__(self, reference_type: str, style_fidelity: float, ref_weight: float): + def __init__(self, reference_type: str, style_fidelity: float, ref_weight: float, ref_with_other_cns: bool=False): self.reference_type = reference_type self.original_style_fidelity = style_fidelity self.style_fidelity = style_fidelity self.ref_weight = ref_weight + self.ref_with_other_cns = ref_with_other_cns def clone(self): - return ReferenceOptions(reference_type=self.reference_type, style_fidelity=self.original_style_fidelity, ref_weight=self.ref_weight) + return ReferenceOptions(reference_type=self.reference_type, style_fidelity=self.original_style_fidelity, ref_weight=self.ref_weight, + ref_with_other_cns=self.ref_with_other_cns) class ReferencePreprocWrapper(AbstractPreprocWrapper): @@ -308,7 +310,12 @@ def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections): #strength_mask = torch.ones_like(x, dtype=x.dtype) * control.strength #control.apply_advanced_strengths_and_masks(x=strength_mask, batched_number=batched_number) #real_cond_hint = control.cond_hint * strength_mask + x * (1 - strength_mask) + orig_kwargs = kwargs + if not control.ref_opts.ref_with_other_cns: + kwargs = kwargs.copy() + kwargs["control"] = None reference_injections.diffusion_model_orig_forward(control.cond_hint.to(dtype=x.dtype).to(device=x.device), *args, **kwargs) + kwargs = orig_kwargs # run diffusion for real now transformer_options[REF_MACHINE_STATE] = MachineState.READ transformer_options[REF_CONTROL_LIST] = ref_controlnets diff --git a/adv_control/nodes_reference.py b/adv_control/nodes_reference.py index c4149c0..e51694b 100644 --- a/adv_control/nodes_reference.py +++ b/adv_control/nodes_reference.py @@ -24,8 +24,9 @@ class ReferenceControlNetNode: CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/Reference" - def load_controlnet(self, reference_type: str, style_fidelity: float, ref_weight: float): - ref_opts = ReferenceOptions(reference_type=reference_type, style_fidelity=style_fidelity, ref_weight=ref_weight) + def load_controlnet(self, reference_type: str, style_fidelity: float, ref_weight: float, ref_with_other_cns: bool=False): + ref_opts = ReferenceOptions(reference_type=reference_type, style_fidelity=style_fidelity, ref_weight=ref_weight, + ref_with_other_cns=ref_with_other_cns) controlnet = ReferenceAdvanced(ref_opts=ref_opts, timestep_keyframes=None) return (controlnet,) From e5841d1759f74791a44bd52c49030ef4c914ba26 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sun, 31 Mar 2024 00:34:56 -0500 Subject: [PATCH 14/17] Added latent_kf and masking support (masks don't work per-area currently, but do work per-latent) --- adv_control/control_reference.py | 52 +++++++++++++++++++++++++------- adv_control/utils.py | 14 ++++----- 2 files changed, 48 insertions(+), 18 deletions(-) diff --git a/adv_control/control_reference.py b/adv_control/control_reference.py index a042e19..895d3cb 100644 --- a/adv_control/control_reference.py +++ b/adv_control/control_reference.py @@ -64,6 +64,12 @@ def refcn_sample_factory(orig_comfy_sample: Callable, is_custom=False) -> Callab module._forward = _forward_inject_BasicTransformerBlock.__get__(module, type(module)) module.injection_holder = injection_holder reference_injections.attn_modules.append(module) + # figure out which module is middle block + if hasattr(model.model.diffusion_model, "middle_block"): + mid_modules = torch_dfs(model.model.diffusion_model.middle_block) + mid_attn_modules: list[RefBasicTransformerBlock] = [module for module in mid_modules if isinstance(module, BasicTransformerBlock)] + for module in mid_attn_modules: + module.injection_holder.is_middle = True # handle diffusion_model forward injection reference_injections.diffusion_model_orig_forward = model.model.diffusion_model.forward model.model.diffusion_model.forward = factory_forward_inject_UNetModel(reference_injections).__get__(model.model.diffusion_model, type(model.model.diffusion_model)) @@ -141,6 +147,8 @@ class ReferencePreprocWrapper(AbstractPreprocWrapper): class ReferenceAdvanced(ControlBase, AdvancedControlBase): + CHANNEL_TO_MULT = {320: 1, 640: 2, 1280: 4} + def __init__(self, ref_opts: ReferenceOptions, timestep_keyframes: TimestepKeyframeGroup, device=None): super().__init__(device) AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controllllite()) @@ -148,14 +156,32 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): self.order = 0 self.latent_format = None self.model_sampling_current = None - self.should_apply_effective_strength = True - self.ref_latent_keyframe_mults = None + self.should_apply_effective_strength = False + self.should_apply_effective_masks = False + self.latent_shape = None + def any_strength_to_apply(self): + return self.should_apply_effective_strength or self.should_apply_effective_masks + def get_effective_strength(self): effective_strength = self.strength if self.current_timestep_keyframe is not None: effective_strength = effective_strength * self.current_timestep_keyframe.strength return effective_strength + + def get_effective_mask_or_float(self, x: Tensor, channels: int, is_mid: bool): + if not self.should_apply_effective_masks: + return self.get_effective_strength() + if is_mid: + div = 8 + else: + div = self.CHANNEL_TO_MULT[channels] + real_mask = torch.ones([self.latent_shape[0], 1, self.latent_shape[2]//div, self.latent_shape[3]//div]).to(dtype=x.dtype, device=x.device) * self.strength + self.apply_advanced_strengths_and_masks(x=real_mask, batched_number=self.batched_number) + # mask is now shape [b, 1, h ,w]; need to turn into [b, h*w, 1] + b, c, h, w = real_mask.shape + real_mask = real_mask.permute(0, 2, 3, 1).reshape(b, h*w, c) + return real_mask def pre_run_advanced(self, model, percent_to_timestep_function): AdvancedControlBase.pre_run_advanced(self, model, percent_to_timestep_function) @@ -165,9 +191,9 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): self.model_sampling_current = model.model_sampling # SDXL is more sensitive to style_fidelity according to sd-webui-controlnet comments if type(model).__name__ == "SDXL": - self.ref_opts.style_fidelity = self.ref_opts.style_fidelity ** 3.0 + self.ref_opts.style_fidelity = self.ref_opts.original_style_fidelity ** 3.0 else: - self.ref_opts.style_fidelity = self.ref_opts.style_fidelity + self.ref_opts.style_fidelity = self.ref_opts.original_style_fidelity def get_control_advanced(self, x_noisy: Tensor, t, cond, batched_number: int): # normal ControlNet stuff @@ -201,9 +227,10 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): self.cond_hint = ddpm_noise_latents(self.cond_hint, sigma=t, noise=None) timestep = self.model_sampling_current.timestep(t) self.should_apply_effective_strength = not (math.isclose(self.strength, 1.0) and math.isclose(self.current_timestep_keyframe.strength, 1.0)) - self.ref_latent_keyframe_mults = self.calc_latent_keyframe_mults(x=x_noisy, batched_number=batched_number) - # prepare mask - self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number) + # prepare mask - use direct_attn, so the mask dims will match source latents (and be smaller) + self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number, direct_attn=True) + self.should_apply_effective_masks = self.latent_keyframes is not None or self.mask_cond_hint is not None or self.tk_mask_cond_hint is not None + self.latent_shape = list(x_noisy.shape) # done preparing; model patches will take care of everything now. # return normal controlnet stuff return control_prev @@ -214,6 +241,8 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): self.latent_format = None del self.model_sampling_current self.model_sampling_current = None + self.should_apply_effective_strength = False + self.should_apply_effective_masks = False def copy(self): c = ReferenceAdvanced(self.ref_opts, self.timestep_keyframes) @@ -251,6 +280,7 @@ class InjectionBasicTransformerBlockHolder: self.original_forward = block._forward self.idx = idx self.attn_weight = 1.0 + self.is_middle = False self.bank_styles = BankStylesBasicTransformerBlock() def restore(self, block: BasicTransformerBlock): @@ -403,8 +433,8 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten style_fidelity = bank_styles.get_avg_style_fidelity() real_bank = bank_styles.bank.copy() for idx, order in enumerate(bank_styles.cn_idx): - if ref_controlnets[idx].should_apply_effective_strength: - effective_strength = ref_controlnets[idx].get_effective_strength() + if ref_controlnets[idx].any_strength_to_apply(): + effective_strength = ref_controlnets[idx].get_effective_mask_or_float(x=n, channels=n.shape[2], is_mid=self.injection_holder.is_middle) real_bank[idx] = real_bank[idx] * effective_strength + context_attn1 * (1-effective_strength) n_uc = self.attn1.to_out(attn1_replace_patch[block_attn1]( n, @@ -434,8 +464,8 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten style_fidelity = bank_styles.get_avg_style_fidelity() real_bank = bank_styles.bank.copy() for idx, order in enumerate(bank_styles.cn_idx): - if ref_controlnets[idx].should_apply_effective_strength: - effective_strength = ref_controlnets[idx].get_effective_strength() + if ref_controlnets[idx].any_strength_to_apply(): + effective_strength = ref_controlnets[idx].get_effective_mask_or_float(x=n, channels=n.shape[2], is_mid=self.injection_holder.is_middle) real_bank[idx] = real_bank[idx] * effective_strength + context_attn1 * (1-effective_strength) n_uc: Tensor = self.attn1( n, diff --git a/adv_control/utils.py b/adv_control/utils.py index d153c43..f568a05 100644 --- a/adv_control/utils.py +++ b/adv_control/utils.py @@ -682,12 +682,12 @@ class AdvancedControlBase: o[i] += prev_val return out - def prepare_mask_cond_hint(self, x_noisy: Tensor, t, cond, batched_number, dtype=None): - self._prepare_mask("mask_cond_hint", self.mask_cond_hint_original, x_noisy, t, cond, batched_number, dtype) - self.prepare_tk_mask_cond_hint(x_noisy, t, cond, batched_number, dtype) + def prepare_mask_cond_hint(self, x_noisy: Tensor, t, cond, batched_number, dtype=None, direct_attn=False): + self._prepare_mask("mask_cond_hint", self.mask_cond_hint_original, x_noisy, t, cond, batched_number, dtype, direct_attn=direct_attn) + self.prepare_tk_mask_cond_hint(x_noisy, t, cond, batched_number, dtype, direct_attn=direct_attn) - def prepare_tk_mask_cond_hint(self, x_noisy: Tensor, t, cond, batched_number, dtype=None): - return self._prepare_mask("tk_mask_cond_hint", self.current_timestep_keyframe.mask_hint_orig, x_noisy, t, cond, batched_number, dtype) + def prepare_tk_mask_cond_hint(self, x_noisy: Tensor, t, cond, batched_number, dtype=None, direct_attn=False): + return self._prepare_mask("tk_mask_cond_hint", self.current_timestep_keyframe.mask_hint_orig, x_noisy, t, cond, batched_number, dtype, direct_attn=direct_attn) def prepare_weight_mask_cond_hint(self, x_noisy: Tensor, batched_number, dtype=None): return self._prepare_mask("weight_mask_cond_hint", self.weights.weight_mask, x_noisy, t=None, cond=None, batched_number=batched_number, dtype=dtype, direct_attn=True) @@ -696,12 +696,12 @@ class AdvancedControlBase: # make mask appropriate dimensions, if present if orig_mask is not None: out_mask = getattr(self, attr_name) - if self.sub_idxs is not None or out_mask is None or x_noisy.shape[2] * 8 != out_mask.shape[1] or x_noisy.shape[3] * 8 != out_mask.shape[2]: + multiplier = 1 if direct_attn else 8 + if self.sub_idxs is not None or out_mask is None or x_noisy.shape[2] * multiplier != out_mask.shape[1] or x_noisy.shape[3] * multiplier != out_mask.shape[2]: self._reset_attr(attr_name) del out_mask # TODO: perform upscale on only the sub_idxs masks at a time instead of all to conserve RAM # resize mask and match batch count - multiplier = 1 if direct_attn else 8 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) From fb4cd18a7fb04fe8bdba73864fb9c1341421bc0d Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sun, 31 Mar 2024 01:13:49 -0500 Subject: [PATCH 15/17] Only expose reference_attn for now since it's the only one with code support (adain coming soon) --- adv_control/control_reference.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/adv_control/control_reference.py b/adv_control/control_reference.py index 895d3cb..c24932f 100644 --- a/adv_control/control_reference.py +++ b/adv_control/control_reference.py @@ -124,7 +124,8 @@ class ReferenceType: ATTN_ADAIN = "reference_attn+adain" STYLE_ALIGN = "StyleAlign" - _LIST = [ATTN, ADAIN, ATTN_ADAIN] + _LIST = [ATTN] + _LIST_FULL = [ATTN, ADAIN, ATTN_ADAIN] class ReferenceOptions: From ab0c75999bd266294b2f16ecfaddfab16c78aaea Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sun, 31 Mar 2024 01:23:28 -0500 Subject: [PATCH 16/17] Update README.md - Reference support (reference_attn) --- README.md | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/README.md b/README.md index 498fb16..b4f86f2 100644 --- a/README.md +++ b/README.md @@ -1,5 +1,5 @@ # ComfyUI-Advanced-ControlNet -Nodes for scheduling ControlNet strength across timesteps and batched latents, as well as applying custom weights and attention masks. The ControlNet nodes here fully support sliding context sampling, like the one used in the [ComfyUI-AnimateDiff-Evolved](https://github.com/Kosinkadink/ComfyUI-AnimateDiff-Evolved) nodes. Currently supports ControlNets, T2IAdapters, ControlLoRAs, ControlLLLite, SparseCtrls, and SVD-ControlNets. +Nodes for scheduling ControlNet strength across timesteps and batched latents, as well as applying custom weights and attention masks. The ControlNet nodes here fully support sliding context sampling, like the one used in the [ComfyUI-AnimateDiff-Evolved](https://github.com/Kosinkadink/ComfyUI-AnimateDiff-Evolved) nodes. Currently supports ControlNets, T2IAdapters, ControlLoRAs, ControlLLLite, SparseCtrls, SVD-ControlNets, and Reference. Custom weights allow replication of the "My prompt is more important" feature of Auto1111's sd-webui ControlNet extension. @@ -8,12 +8,14 @@ ControlNet preprocessors are available through [comfyui_controlnet_aux](https:// ## Features - Timestep and latent strength scheduling - Attention masks -- Soft weights to replicate "My prompt is more important" feature from sd-webui ControlNet extension, and also change the scaling. -- ControlNet, T2IAdapter, and ControlLoRA support for sliding context windows. +- Soft weights to replicate "My prompt is more important" feature from sd-webui ControlNet extension, and also change the scaling +- ControlNet, T2IAdapter, and ControlLoRA support for sliding context windows - ControlLLLite support (requires model_optional to be passed into and out of Apply Advanced ControlNet node) - SparseCtrl support - SVD-ControlNet support - Stable Video Diffusion ControlNets trained by **CiaraRowles**: [Depth](https://huggingface.co/CiaraRowles/temporal-controlnet-depth-svd-v1/tree/main/controlnet), [Lineart](https://huggingface.co/CiaraRowles/temporal-controlnet-lineart-svd-v1/tree/main/controlnet) +- Reference support + - Currently, only ```reference_attn``` is exposed (equivalent of reference_only in Auto1111), reference_adain (and +attn) under construction. ```style_fidelity``` and ```ref_weight``` are equivalent to style_fidelity and control_weight in Auto1111, respectively, and strength of the Apply ControlNet is the balance between ref-influenced result and no-ref result. ## Table of Contents: - [Scheduling Explanation](#scheduling-explanation) From 829d10ddc0340483caab7225d11cf3d62422b3c6 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sun, 31 Mar 2024 01:30:07 -0500 Subject: [PATCH 17/17] Fixed BIGMIN and BIGMAX to conform with javascript limits, pretty-fied deprecated node --- adv_control/nodes.py | 2 +- adv_control/utils.py | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/adv_control/nodes.py b/adv_control/nodes.py index ce4c0e7..cea6ada 100644 --- a/adv_control/nodes.py +++ b/adv_control/nodes.py @@ -275,5 +275,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { # LOOSEControl #"ACN_ControlNetLoaderWithLoraAdvanced": "Load Adv. ControlNet Model w/ LoRA 🛂🅐🅒🅝", # Deprecated - "LoadImagesFromDirectory": "Load Images [DEPRECATED] 🛂🅐🅒🅝", + "LoadImagesFromDirectory": "🚫Load Images [DEPRECATED] 🛂🅐🅒🅝", } diff --git a/adv_control/utils.py b/adv_control/utils.py index f568a05..5caefbd 100644 --- a/adv_control/utils.py +++ b/adv_control/utils.py @@ -12,8 +12,8 @@ from comfy.model_patcher import ModelPatcher from .logger import logger -BIGMIN = -(2**63-1) -BIGMAX = (2**63-1) +BIGMIN = -(2**53-1) +BIGMAX = (2**53-1) def load_torch_file_with_dict_factory(controlnet_data: dict[str, Tensor], orig_load_torch_file: Callable): def load_torch_file_with_dict(*args, **kwargs):