Converted sampler_sample_wrapper to a outer_sample_wrapper so that controlnet prerun can be ran properly
This commit is contained in:
+32
-34
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 = {
|
||||
|
||||
+18
-171
@@ -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
|
||||
|
||||
@@ -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]):
|
||||
|
||||
Reference in New Issue
Block a user