Started to replace anc_sample_factory with acn_sampler_sample_wrapper WrapperHook, ReferenceCN now can respect individual conds

This commit is contained in:
Jedrzej Kosinski
2024-11-14 08:12:54 -06:00
parent 57dfa7d1ce
commit 52f8220260
5 changed files with 63 additions and 25 deletions
+16 -3
View File
@@ -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:
+6 -8
View File
@@ -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)
+2 -2
View File
@@ -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
+32 -12
View File
@@ -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:
+7
View File
@@ -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)