From 7f84e973d6eaab58a7dd14675112277fa45b5b79 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Thu, 14 Nov 2024 10:53:59 -0600 Subject: [PATCH] Converted sampler_sample_wrapper to a outer_sample_wrapper so that controlnet prerun can be ran properly --- adv_control/control.py | 66 ++++++----- adv_control/control_reference.py | 10 +- adv_control/dinklink.py | 8 +- adv_control/nodes.py | 5 - adv_control/sampling.py | 189 +++---------------------------- adv_control/utils.py | 4 +- 6 files changed, 61 insertions(+), 221 deletions(-) diff --git a/adv_control/control.py b/adv_control/control.py index 79b5e60..c589ed1 100644 --- a/adv_control/control.py +++ b/adv_control/control.py @@ -529,17 +529,16 @@ def convert_to_advanced(control, timestep_keyframe: TimestepKeyframeGroup=None): return control -def convert_all_to_advanced(conds: list[list[dict[str]]]) -> tuple[bool, list]: +def convert_all_to_advanced(conds: dict[str, list[dict[str]]]) -> tuple[bool, list]: cache = {} modified = False - new_conds = [] - for cond in conds: - converted_cond = None + new_conds = {} + for cond_type in conds: + converted_cond: list[dict[str]] = None + cond = conds[cond_type] if cond is not None: - need_to_convert = False - # first, check if there is even a need to convert - for sub_cond in cond: - actual_cond = sub_cond[1] + for actual_cond in cond: + need_to_convert = False if "control" in actual_cond: if not are_all_advanced_controlnet(actual_cond["control"]): need_to_convert = True @@ -548,23 +547,20 @@ def convert_all_to_advanced(conds: list[list[dict[str]]]) -> tuple[bool, list]: converted_cond = cond else: converted_cond = [] - for sub_cond in cond: - new_sub_cond: list = [] - for actual_cond in sub_cond: - if not type(actual_cond) == dict: - new_sub_cond.append(actual_cond) - continue - if "control" not in actual_cond: - new_sub_cond.append(actual_cond) - elif are_all_advanced_controlnet(actual_cond["control"]): - new_sub_cond.append(actual_cond) - else: - actual_cond = actual_cond.copy() - actual_cond["control"] = _convert_all_control_to_advanced(actual_cond["control"], cache) - new_sub_cond.append(actual_cond) - modified = True - converted_cond.append(new_sub_cond) - new_conds.append(converted_cond) + for actual_cond in cond: + if not isinstance(actual_cond, dict): + converted_cond.append(actual_cond) + continue + if "control" not in actual_cond: + converted_cond.append(actual_cond) + elif are_all_advanced_controlnet(actual_cond["control"]): + converted_cond.append(actual_cond) + else: + actual_cond = actual_cond.copy() + actual_cond["control"] = _convert_all_control_to_advanced(actual_cond["control"], cache) + converted_cond.append(actual_cond) + modified = True + new_conds[cond_type] = converted_cond return modified, new_conds @@ -606,19 +602,21 @@ def _convert_all_control_to_advanced(input_object: ControlBase, cache: dict): return output_object -def restore_all_controlnet_conns(conds: list[list[dict[str]]]): +def restore_all_controlnet_conns(conds: dict[str, list[dict[str]]]): # if a cn has an _orig_previous_controlnet property, restore it and delete - for main_cond in conds: - if main_cond is not None: - for cond in main_cond: - if "control" in cond[1]: + for cond_type in conds: + cond = conds[cond_type] + if cond is not None: + for actual_cond in cond: + if "control" in actual_cond: # if ACN is the one to have initialized it, delete it # TODO: maybe check if someone else did a similar hack, and carefully pluck out our stuff? - if CONTROL_INIT_BY_ACN in cond[1]: - cond[1].pop("control") - cond[1].pop(CONTROL_INIT_BY_ACN) + if CONTROL_INIT_BY_ACN in actual_cond: + actual_cond.pop("control") + actual_cond.pop(CONTROL_INIT_BY_ACN) else: - _restore_all_controlnet_conns(cond[1]["control"]) + _restore_all_controlnet_conns(actual_cond["control"]) + def _restore_all_controlnet_conns(input_object: ControlBase): diff --git a/adv_control/control_reference.py b/adv_control/control_reference.py index c1311a8..5455efb 100644 --- a/adv_control/control_reference.py +++ b/adv_control/control_reference.py @@ -316,7 +316,7 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): return self -def handle_context_ref_setup(contextref_obj, transformer_options: dict, conds: list[dict]): +def handle_context_ref_setup(contextref_obj, transformer_options: dict, conds: dict[str, list[dict, str]]): transformer_options[CONTEXTREF_MACHINE_STATE] = MachineState.OFF # verify version is compatible if contextref_obj.version > HIGHEST_VERSION_SUPPORT: @@ -371,7 +371,7 @@ def _create_tks_from_dict_list(dlist: list[dict[str]]) -> TimestepKeyframeGroup: return tks -def _add_context_ref_to_conds(conds: list[list[dict[str]]], context_ref: ReferenceAdvanced): +def _add_context_ref_to_conds(conds: dict[list[dict[str]]], context_ref: ReferenceAdvanced): def _add_context_ref_to_existing_control(control: ControlBase, context_ref: ReferenceAdvanced): curr_cn = control while curr_cn is not None: @@ -395,10 +395,10 @@ def _add_context_ref_to_conds(conds: list[list[dict[str]]], context_ref: Referen actual_cond[CONTROL_INIT_BY_ACN] = True # either add context_ref to end of existing cnet chain, or init 'control' key on actual cond - for cond in conds: + for cond_type in conds: + cond = conds[cond_type] if cond is not None: - for sub_cond in cond: - actual_cond = sub_cond[1] + for actual_cond in cond: _add_context_ref(actual_cond, context_ref) diff --git a/adv_control/dinklink.py b/adv_control/dinklink.py index 1c0fcc0..a0e8b39 100644 --- a/adv_control/dinklink.py +++ b/adv_control/dinklink.py @@ -14,7 +14,7 @@ from __future__ import annotations import comfy.hooks from comfy.patcher_extension import WrappersMP -from .sampling import acn_sampler_sample_wrapper +from .sampling import acn_outer_sample_wrapper from .utils import WrapperConsts DINKLINK = "__DINKLINK" @@ -33,6 +33,6 @@ def prepare_dinklink(): d = get_dinklink() 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) + link_acn[WrapperConsts.ACN_CREATE_SAMPLER_SAMPLE_WRAPPER] = (WrappersMP.OUTER_SAMPLE, + WrapperConsts.ACN_OUTER_SAMPLE_WRAPPER_KEY, + acn_outer_sample_wrapper) diff --git a/adv_control/nodes.py b/adv_control/nodes.py index e1c33f9..32593a5 100644 --- a/adv_control/nodes.py +++ b/adv_control/nodes.py @@ -18,11 +18,6 @@ from .nodes_deprecated import (LoadImagesFromDirectory, ScaledSoftUniversalWeigh ControlNetLoaderAdvancedDEPR, DiffControlNetLoaderAdvancedDEPR) 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) - # NODE MAPPING NODE_CLASS_MAPPINGS = { diff --git a/adv_control/sampling.py b/adv_control/sampling.py index d16aae6..59f2ce9 100644 --- a/adv_control/sampling.py +++ b/adv_control/sampling.py @@ -72,32 +72,31 @@ 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, +def should_register_outer_sample_wrapper(hook, model, model_options: dict, target, registered: list): + wrappers = comfy.patcher_extension.get_wrappers_with_key(comfy.patcher_extension.WrappersMP.OUTER_SAMPLE, + WrapperConsts.ACN_OUTER_SAMPLE_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, + comfy.patcher_extension.add_wrapper_with_key(comfy.patcher_extension.WrappersMP.OUTER_SAMPLE, + WrapperConsts.ACN_OUTER_SAMPLE_WRAPPER_KEY, + acn_outer_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 + hook.hook_id = WrapperConsts.ACN_OUTER_SAMPLE_WRAPPER_KEY + hook.custom_should_register = should_register_outer_sample_wrapper hooks.add(hook) return hooks -def acn_sampler_sample_wrapper(executor, *args, **kwargs): +def acn_outer_sample_wrapper(executor, *args, **kwargs): controlnets_modified = False - guider: comfy.samplers.CFGGuider = args[0] + guider: comfy.samplers.CFGGuider = executor.class_obj model = guider.model_patcher - extra_args: dict = args[2] orig_conds = guider.conds - orig_model_options = extra_args["model_options"] + orig_model_options = guider.model_options try: new_model_options = orig_model_options # if context options present, perform some special actions that may be required @@ -105,13 +104,13 @@ def acn_sampler_sample_wrapper(executor, *args, **kwargs): if has_sliding_context_windows(guider.model_patcher): 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()) + controlnets_modified, conds = support_sliding_context_windows(orig_conds) if controlnets_modified: guider.conds = conds # 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, new_model_options["transformer_options"], guider.conds.values()) + context_refs = handle_context_ref_setup(existing_contextref_obj, new_model_options["transformer_options"], guider.conds) controlnets_modified = True # look for Advanced ControlNets that will require intervention to work ref_set = set() @@ -199,7 +198,7 @@ def acn_sampler_sample_wrapper(executor, *args, **kwargs): 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 - extra_args["model_options"] = new_model_options + guider.model_options = new_model_options # continue with original function return executor(*args, **kwargs) finally: @@ -224,165 +223,13 @@ def acn_sampler_sample_wrapper(executor, *args, **kwargs): reference_injections.cleanup() finally: # restore model_options - extra_args["model_options"] = orig_model_options + guider.model_options = orig_model_options # restore guider.conds guider.conds = orig_conds # restore controlnets in conds, if needed if controlnets_modified: - restore_all_controlnet_conns(orig_conds.values()) + restore_all_controlnet_conns(guider.conds) + del orig_conds + del orig_model_options del model del guider - - -def acn_sample_factory(orig_comfy_sample: Callable, is_custom=False) -> Callable: - def acn_sample(model: ModelPatcher, *args, **kwargs): - controlnets_modified = False - orig_positive = args[-3] - orig_negative = args[-2] - try: - orig_model_options = model.model_options - # check if positive or negative conds contain ref cn - positive = args[-3] - negative = args[-2] - # if context options present, perform some special actions that may be required - context_refs = [] - if has_sliding_context_windows(model): - model.model_options = model.model_options.copy() - model.model_options["transformer_options"] = model.model_options["transformer_options"].copy() - # convert all CNs to Advanced if needed - controlnets_modified, conds = support_sliding_context_windows([positive, negative]) - positive, negative = conds - if controlnets_modified: - args = list(args) - args[-3] = positive - args[-2] = negative - args = tuple(args) - # enable ContextRef, if requested - existing_contextref_obj = get_contextref_obj(model) - if existing_contextref_obj is not None: - context_refs = handle_context_ref_setup(existing_contextref_obj, model.model_options["transformer_options"], [positive, negative]) - 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 - if positive is not None: - for cond in positive: - if "control" in cond[1]: - ref_set.update(get_refcn(cond[1]["control"])) - lllite_dict.update(get_lllitecn(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"])) - lllite_dict.update(get_lllitecn(cond[1]["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() - 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) - # if no ref cn found, do original function immediately - if len(ref_set) == 0 and len(context_refs) == 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]) - for i, module in enumerate(attn_modules): - injection_holder = InjectionBasicTransformerBlockHolder(block=module, idx=i) - injection_holder.attn_weight = float(i) / float(len(attn_modules)) - if hasattr(module, "_forward"): # backward compatibility - module._forward = _forward_inject_BasicTransformerBlock.__get__(module, type(module)) - else: - 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 - - # next, handle gn module injection (TimestepEmbedSequential) - # TODO: figure out the logic behind these hardcoded indexes - if type(model.model).__name__ == "SDXL": - input_block_indices = [4, 5, 7, 8] - output_block_indices = [0, 1, 2, 3, 4, 5] - else: - input_block_indices = [4, 5, 7, 8, 10, 11] - output_block_indices = [0, 1, 2, 3, 4, 5, 6, 7] - if hasattr(model.model.diffusion_model, "middle_block"): - module = model.model.diffusion_model.middle_block - injection_holder = InjectionTimestepEmbedSequentialHolder(block=module, idx=0, is_middle=True) - injection_holder.gn_weight = 0.0 - module.injection_holder = injection_holder - reference_injections.gn_modules.append(module) - for w, i in enumerate(input_block_indices): - module = model.model.diffusion_model.input_blocks[i] - injection_holder = InjectionTimestepEmbedSequentialHolder(block=module, idx=i, is_input=True) - injection_holder.gn_weight = 1.0 - float(w) / float(len(input_block_indices)) - module.injection_holder = injection_holder - reference_injections.gn_modules.append(module) - for w, i in enumerate(output_block_indices): - module = model.model.diffusion_model.output_blocks[i] - injection_holder = InjectionTimestepEmbedSequentialHolder(block=module, idx=i, is_output=True) - injection_holder.gn_weight = float(w) / float(len(output_block_indices)) - module.injection_holder = injection_holder - reference_injections.gn_modules.append(module) - # hack gn_module forwards and update weights - for i, module in enumerate(reference_injections.gn_modules): - module.injection_holder.gn_weight *= 2 - - # 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 - 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) - new_model_options["transformer_options"][CONTEXTREF_CLEAN_FUNC] = reference_injections.clean_contextref_module_mem - model.model_options = new_model_options - # continue with original function - return orig_comfy_sample(model, *args, **kwargs) - finally: - # cleanup injections - # 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_all() - del module.injection_holder - del attn_modules - # restore gn modules - gn_modules: list[RefTimestepEmbedSequential] = reference_injections.gn_modules - for module in gn_modules: - module.injection_holder.restore(module) - module.injection_holder.clean_all() - del module.injection_holder - del gn_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)) - # cleanup - reference_injections.cleanup() - finally: - # restore model_options - model.model_options = orig_model_options - # restore controlnets in conds, if needed - if controlnets_modified: - restore_all_controlnet_conns([orig_positive, orig_negative]) - - return acn_sample diff --git a/adv_control/utils.py b/adv_control/utils.py index e1d2ed6..acc26d6 100644 --- a/adv_control/utils.py +++ b/adv_control/utils.py @@ -34,8 +34,8 @@ def load_torch_file_with_dict_factory(controlnet_data: dict[str, Tensor], orig_l class WrapperConsts: ACN = "ACN" VERSION = "version" - ACN_SAMPLER_SAMPLER_WRAPPER_KEY = "ACN_sampler_sample_wrapper" - CREATE_SAMPLER_SAMPLE_WRAPPER = "create_sampler_sample_wrapper" + ACN_OUTER_SAMPLE_WRAPPER_KEY = "ACN_outer_sample_wrapper" + ACN_CREATE_SAMPLER_SAMPLE_WRAPPER = "create_outer_sample_wrapper" def get_properly_arranged_t2i_weights(initial_weights: list[float]):