Started to replace anc_sample_factory with acn_sampler_sample_wrapper WrapperHook, ReferenceCN now can respect individual conds
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user