Make forward_inject_UNetModel a built-in diffusion_model wrapper instead to fix weird issue with ContextRef + ImageInjection

This commit is contained in:
Jedrzej Kosinski
2024-11-17 08:35:48 -06:00
parent a2430fb880
commit df899e8b7e
2 changed files with 35 additions and 24 deletions
+30 -17
View File
@@ -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
+5 -7
View File
@@ -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: