diff --git a/adv_control/control_reference.py b/adv_control/control_reference.py index 2c860cd..c1311a8 100644 --- a/adv_control/control_reference.py +++ b/adv_control/control_reference.py @@ -7,6 +7,7 @@ from torch import Tensor import comfy.model_management import comfy.sample +import comfy.hooks import comfy.model_patcher import comfy.utils from comfy.controlnet import ControlBase @@ -46,6 +47,7 @@ RETURNED_CONTEXTREF_VERSION = 1 class RefConst: OPTS = "refcn_opts" CREF_MODE = "contextref_mode" + REFCN_PRESENT_IN_CONDS = "refcn_present_in_conds" class MachineState: @@ -142,7 +144,7 @@ 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): + def __init__(self, ref_opts: ReferenceOptions, timestep_keyframes: TimestepKeyframeGroup, extra_hooks: comfy.hooks.HookGroup=None): super().__init__() AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controllllite(), allow_condhint_latents=True) # TODO: allow vae_optional to be used instead of preprocessor @@ -155,6 +157,8 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): self.should_apply_adain_effective_strength = False self.should_apply_effective_masks = False self.latent_shape = None + # wrapper hooks + self.extra_hooks = extra_hooks.clone() if extra_hooks else self.import_and_create_wrapper_hooks() # ContextRef stuff self.is_context_ref = False self.contextref_cond_idx = -1 @@ -166,6 +170,10 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): return self._current_timestep_keyframe.control_weights.extras.get(RefConst.OPTS, self._ref_opts) return self._ref_opts + def import_and_create_wrapper_hooks(self): + from .sampling import create_wrapper_hooks + return create_wrapper_hooks() + def any_attn_strength_to_apply(self): return self.should_apply_attn_effective_strength or self.should_apply_effective_masks @@ -280,6 +288,7 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): 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. + transformer_options[RefConst.REFCN_PRESENT_IN_CONDS] = True # return normal controlnet stuff return control_prev @@ -294,7 +303,7 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): self.should_apply_effective_masks = False def copy(self): - c = ReferenceAdvanced(self.ref_opts, self.timestep_keyframes) + c = ReferenceAdvanced(self.ref_opts, self.timestep_keyframes, self.extra_hooks) c.order = self.order c.is_context_ref = self.is_context_ref self.copy_to(c) @@ -741,7 +750,11 @@ def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections): reference_injections.clean_contextref_module_mem() context_controlnets = [] # discard any controlnets that should not run - ref_controlnets = [z for z in ref_controlnets if z.should_run()] + refcn_present_in_conds = transformer_options.get(RefConst.REFCN_PRESENT_IN_CONDS, False) + if refcn_present_in_conds: + ref_controlnets = [z for z in ref_controlnets if z.should_run()] + else: + ref_controlnets = [] context_controlnets = [z for z in context_controlnets if z.should_run()] # if nothing related to reference controlnets, do nothing special if len(ref_controlnets) == 0 and len(context_controlnets) == 0: diff --git a/adv_control/dinklink.py b/adv_control/dinklink.py index e30a3bd..1c0fcc0 100644 --- a/adv_control/dinklink.py +++ b/adv_control/dinklink.py @@ -15,6 +15,7 @@ import comfy.hooks from comfy.patcher_extension import WrappersMP from .sampling import acn_sampler_sample_wrapper +from .utils import WrapperConsts DINKLINK = "__DINKLINK" @@ -27,14 +28,11 @@ def init_dinklink(): def get_dinklink() -> dict[str, dict[str]]: return getattr(comfy.hooks, DINKLINK) -class Consts: - ACN = "ACN" - VERSION = "version" - CREATE_SAMPLER_SAMPLE_WRAPPER = "create_sampler_sample_wrapper" - def prepare_dinklink(): # expose acn_sampler_sample_wrapper d = get_dinklink() - link_acn = d.setdefault(Consts.ACN, {}) - link_acn[Consts.VERSION] = 1 - link_acn[Consts.CREATE_SAMPLER_SAMPLE_WRAPPER] = (WrappersMP.SAMPLER_SAMPLE, Consts.ACN, acn_sampler_sample_wrapper) + link_acn = d.setdefault(WrapperConsts.ACN, {}) + link_acn[WrapperConsts.VERSION] = 1 + link_acn[WrapperConsts.CREATE_SAMPLER_SAMPLE_WRAPPER] = (WrappersMP.SAMPLER_SAMPLE, + WrapperConsts.ACN_SAMPLER_SAMPLER_WRAPPER_KEY, + acn_sampler_sample_wrapper) diff --git a/adv_control/nodes.py b/adv_control/nodes.py index 1849fb0..e1c33f9 100644 --- a/adv_control/nodes.py +++ b/adv_control/nodes.py @@ -20,8 +20,8 @@ from .logger import logger from .sampling import acn_sample_factory # inject sample functions -comfy.sample.sample = acn_sample_factory(comfy.sample.sample) -comfy.sample.sample_custom = acn_sample_factory(comfy.sample.sample_custom, is_custom=True) +#comfy.sample.sample = acn_sample_factory(comfy.sample.sample) +#comfy.sample.sample_custom = acn_sample_factory(comfy.sample.sample_custom, is_custom=True) # NODE MAPPING diff --git a/adv_control/sampling.py b/adv_control/sampling.py index 2e8016a..d16aae6 100644 --- a/adv_control/sampling.py +++ b/adv_control/sampling.py @@ -1,6 +1,8 @@ from typing import Callable, Union +import comfy.hooks import comfy.model_patcher +import comfy.patcher_extension import comfy.sample import comfy.samplers from comfy.model_patcher import ModelPatcher @@ -16,7 +18,7 @@ from .control_reference import (ReferenceAdvanced, ReferenceInjections, handle_context_ref_setup, REF_CONTROL_LIST_ALL, CONTEXTREF_CLEAN_FUNC) from .control_lllite import (ControlLLLiteAdvanced) -from .utils import torch_dfs +from .utils import torch_dfs, WrapperConsts def support_sliding_context_windows(conds) -> tuple[bool, list[dict]]: @@ -70,6 +72,25 @@ def get_lllitecn(control: ControlBase): cn_dict.update(get_lllitecn(control.previous_controlnet)) return cn_dict +def should_register_sampler_sampler_wrapper(hook, model, model_options: dict, target, registered: list): + wrappers = comfy.patcher_extension.get_wrappers_with_key(comfy.patcher_extension.WrappersMP.SAMPLER_SAMPLE, + WrapperConsts.ACN_SAMPLER_SAMPLER_WRAPPER_KEY, + model_options, is_model_options=True) + return len(wrappers) == 0 + +def create_wrapper_hooks(): + wrappers = {} + comfy.patcher_extension.add_wrapper_with_key(comfy.patcher_extension.WrappersMP.SAMPLER_SAMPLE, + WrapperConsts.ACN_SAMPLER_SAMPLER_WRAPPER_KEY, + acn_sampler_sample_wrapper, + transformer_options=wrappers) + hooks = comfy.hooks.HookGroup() + hook = comfy.hooks.WrapperHook(wrappers) + hook.hook_id = WrapperConsts.ACN_SAMPLER_SAMPLER_WRAPPER_KEY + hook.custom_should_register = should_register_sampler_sampler_wrapper + hooks.add(hook) + return hooks + def acn_sampler_sample_wrapper(executor, *args, **kwargs): controlnets_modified = False guider: comfy.samplers.CFGGuider = args[0] @@ -78,10 +99,11 @@ def acn_sampler_sample_wrapper(executor, *args, **kwargs): orig_conds = guider.conds orig_model_options = extra_args["model_options"] try: + new_model_options = orig_model_options # if context options present, perform some special actions that may be required context_refs = [] if has_sliding_context_windows(guider.model_patcher): - extra_args["model_options"] = comfy.model_patcher.create_model_options_clone(orig_model_options) + new_model_options = comfy.model_patcher.create_model_options_clone(new_model_options) # convert all CNs to Advanced if needed controlnets_modified, conds = support_sliding_context_windows(orig_conds.values()) if controlnets_modified: @@ -89,24 +111,23 @@ def acn_sampler_sample_wrapper(executor, *args, **kwargs): # enable ContextRef, if requested existing_contextref_obj = get_contextref_obj(guider.model_patcher) if existing_contextref_obj is not None: - context_refs = handle_context_ref_setup(existing_contextref_obj, extra_args["model_options"]["transformer_options"], guider.conds.values()) + context_refs = handle_context_ref_setup(existing_contextref_obj, new_model_options["transformer_options"], guider.conds.values()) controlnets_modified = True # look for Advanced ControlNets that will require intervention to work ref_set = set() lllite_dict: dict[ControlLLLiteAdvanced, None] = {} # dicts preserve insertion order since py3.7 for outer_cond in guider.conds.values(): for cond in outer_cond: - if "control" in cond[1]: - ref_set.update(get_refcn(cond[1]["control"])) - lllite_dict.update(get_lllitecn(cond[1]["control"])) + if "control" in cond: + ref_set.update(get_refcn(cond["control"])) + lllite_dict.update(get_lllitecn(cond["control"])) # if lllite found, apply patches to a cloned model_options, and continue if len(lllite_dict) > 0: lllite_list = list(lllite_dict.keys()) - model.model_options = model.model_options.copy() - model.model_options["transformer_options"] = model.model_options["transformer_options"].copy() + new_model_options = comfy.model_patcher.create_model_options_clone(new_model_options) lllite_list.reverse() # reverse so that patches will be applied in expected order for lll in lllite_list: - lll.live_model_patches(model.model_options) + lll.live_model_patches(new_model_options) # if no ref cn found, do original function immediately if len(ref_set) == 0 and len(context_refs) == 0: return executor(*args, **kwargs) @@ -174,12 +195,11 @@ def acn_sampler_sample_wrapper(executor, *args, **kwargs): 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 - new_model_options = model.model_options.copy() - new_model_options["transformer_options"] = model.model_options["transformer_options"].copy() + new_model_options = comfy.model_patcher.create_model_options_clone(new_model_options) ref_list: list[ReferenceAdvanced] = list(ref_set) new_model_options["transformer_options"][REF_CONTROL_LIST_ALL] = sorted(ref_list, key=lambda x: x.order) new_model_options["transformer_options"][CONTEXTREF_CLEAN_FUNC] = reference_injections.clean_contextref_module_mem - model.model_options = new_model_options + extra_args["model_options"] = new_model_options # continue with original function return executor(*args, **kwargs) finally: diff --git a/adv_control/utils.py b/adv_control/utils.py index 8c98f89..e1d2ed6 100644 --- a/adv_control/utils.py +++ b/adv_control/utils.py @@ -31,6 +31,13 @@ def load_torch_file_with_dict_factory(controlnet_data: dict[str, Tensor], orig_l return load_torch_file_with_dict +class WrapperConsts: + ACN = "ACN" + VERSION = "version" + ACN_SAMPLER_SAMPLER_WRAPPER_KEY = "ACN_sampler_sample_wrapper" + CREATE_SAMPLER_SAMPLE_WRAPPER = "create_sampler_sample_wrapper" + + def get_properly_arranged_t2i_weights(initial_weights: list[float]): new_weights = [] new_weights.extend([initial_weights[0]]*3)