diff --git a/adv_control/control_reference.py b/adv_control/control_reference.py index 5455efb..f5fe8f1 100644 --- a/adv_control/control_reference.py +++ b/adv_control/control_reference.py @@ -6,6 +6,7 @@ import torch from torch import Tensor import comfy.model_management +import comfy.patcher_extension import comfy.sample import comfy.hooks import comfy.model_patcher @@ -687,7 +688,6 @@ class ReferenceInjections: def __init__(self, attn_modules: list['RefBasicTransformerBlock']=None, gn_modules: list['RefTimestepEmbedSequential']=None): self.attn_modules = attn_modules if attn_modules else [] self.gn_modules = gn_modules if gn_modules else [] - self.diffusion_model_orig_forward: Callable = None def clean_ref_module_mem(self): for attn_module in self.attn_modules: @@ -731,16 +731,30 @@ class ReferenceInjections: self.attn_modules = [] del self.gn_modules self.gn_modules = [] - self.diffusion_model_orig_forward = None -def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections): - def forward_inject_UNetModel(self, x: Tensor, *args, **kwargs): - # get control and transformer_options from kwargs +def handle_reference_injection(model_options: dict, reference_injections: ReferenceInjections): + # register wrapper functions on transformer_options + comfy.patcher_extension.add_wrapper_with_key(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, + "ACN_refcn_diffusion_model", + refcn_diffusion_model_wrapper_factory(reference_injections), + model_options, is_model_options=True) + + +def refcn_diffusion_model_wrapper_factory(reference_injections: ReferenceInjections): + def refcn_diffusion_model_wrapper(executor, x, *args, **kwargs): + # get control and transformer_options from args real_args = list(args) real_kwargs = list(kwargs.keys()) - control = kwargs.get("control", None) - transformer_options: dict[str] = kwargs.get("transformer_options", {}) + # args values (x is treated separately, so all args are actually shifted by -1): + # -1: x + # 0: timesteps + # 1: context + # 2: y + # 3: control + # 4: transformer_options + control = args[3] + transformer_options = args[4] # NOTE: adds support for both ReferenceCN and ContextRef, so need to track them separately # get ReferenceAdvanced objects ref_controlnets: list[ReferenceAdvanced] = transformer_options.get(REF_CONTROL_LIST_ALL, []) @@ -758,7 +772,7 @@ def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections): 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: - return reference_injections.diffusion_model_orig_forward(x, *args, **kwargs) + return executor(x, *args, **kwargs) try: # assign cond and uncond idxs batched_number = len(transformer_options["cond_or_uncond"]) @@ -813,13 +827,14 @@ def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections): transformer_options[REF_READ_ADAIN_CONTROL_LIST] = read_adain_list transformer_options[REF_WRITE_ADAIN_CONTROL_LIST] = write_adain_list - orig_kwargs = kwargs + orig_args = args # disable other controlnets for this run, if specified if not control.ref_opts.ref_with_other_cns: - kwargs = kwargs.copy() - kwargs["control"] = None - reference_injections.diffusion_model_orig_forward(control.cond_hint.to(dtype=x.dtype).to(device=x.device), *args, **kwargs) - kwargs = orig_kwargs + args = list(args) + args[3] = None + args = tuple(args) + executor(control.cond_hint.to(dtype=x.dtype).to(device=x.device), *args, **kwargs) + args = orig_args # prepare running diffusion for real now read_attn_list = [] write_attn_list = [] @@ -853,7 +868,7 @@ def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections): transformer_options[REF_WRITE_ADAIN_CONTROL_LIST] = write_adain_list # run diffusion for real try: - return reference_injections.diffusion_model_orig_forward(x, *args, **kwargs) + return executor(x, *args, **kwargs) finally: # increment current cond idx if len(context_controlnets) > 0: @@ -864,9 +879,7 @@ def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections): reference_injections.clean_ref_module_mem() if len(adain_controlnets) > 0 or len(context_adain_controlnets) > 0: openaimodel.forward_timestep_embed = orig_forward_timestep_embed - - - return forward_inject_UNetModel + return refcn_diffusion_model_wrapper # dummy class just to help IDE keep track of injected variables diff --git a/adv_control/sampling.py b/adv_control/sampling.py index e6eb3ea..d97cd98 100644 --- a/adv_control/sampling.py +++ b/adv_control/sampling.py @@ -14,8 +14,8 @@ from .control import convert_all_to_advanced, restore_all_controlnet_conns from .control_reference import (ReferenceAdvanced, ReferenceInjections, RefBasicTransformerBlock, RefTimestepEmbedSequential, InjectionBasicTransformerBlockHolder, InjectionTimestepEmbedSequentialHolder, - _forward_inject_BasicTransformerBlock, factory_forward_inject_UNetModel, - handle_context_ref_setup, + _forward_inject_BasicTransformerBlock, + handle_context_ref_setup, handle_reference_injection, REF_CONTROL_LIST_ALL, CONTEXTREF_CLEAN_FUNC) from .utils import torch_dfs, WrapperConsts @@ -172,11 +172,11 @@ def acn_outer_sample_wrapper(executor, *args, **kwargs): 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 = comfy.model_patcher.create_model_options_clone(new_model_options) + # handle diffusion_model forward injection + handle_reference_injection(new_model_options, reference_injections) + 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 @@ -199,8 +199,6 @@ def acn_outer_sample_wrapper(executor, *args, **kwargs): 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: