Simplified ControlLLLite implementation, since transformer_options are passed into get_control with ComfyUI rework

This commit is contained in:
Jedrzej Kosinski
2024-11-14 12:10:28 -06:00
parent b9c53b79c9
commit 326b26ed95
2 changed files with 10 additions and 30 deletions
+10 -12
View File
@@ -19,8 +19,8 @@ from .utils import (AdvancedControlBase, TimestepKeyframeGroup, ControlWeights,
# based on set_model_patch code in comfy/model_patcher.py
def set_model_patch(model_options, patch, name):
to = model_options["transformer_options"]
def set_model_patch(transformer_options, patch, name):
to = transformer_options
# check if patch was already added
if "patches" in to:
current_patches = to["patches"].get(name, [])
@@ -30,11 +30,11 @@ def set_model_patch(model_options, patch, name):
to["patches"] = {}
to["patches"][name] = to["patches"].get(name, []) + [patch]
def set_model_attn1_patch(model_options, patch):
set_model_patch(model_options, patch, "attn1_patch")
def set_model_attn1_patch(transformer_options, patch):
set_model_patch(transformer_options, patch, "attn1_patch")
def set_model_attn2_patch(model_options, patch):
set_model_patch(model_options, patch, "attn2_patch")
def set_model_attn2_patch(transformer_options, patch):
set_model_patch(transformer_options, patch, "attn2_patch")
def extra_options_to_module_prefix(extra_options):
@@ -298,10 +298,6 @@ class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase):
self.latent_dims_div2 = None
self.latent_dims_div4 = None
def live_model_patches(self, model_options):
set_model_attn1_patch(model_options, self.patch_attn1.set_control(self))
set_model_attn2_patch(model_options, self.patch_attn2.set_control(self))
def set_cond_hint_inject(self, *args, **kwargs):
to_return = super().set_cond_hint_inject(*args, **kwargs)
# cond hint for LLLite needs to be scaled between (-1, 1) instead of (0, 1)
@@ -315,7 +311,7 @@ class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase):
self.patch_attn2.set_control(self)
#logger.warn(f"in pre_run_advanced: {id(self)}")
def get_control_advanced(self, x_noisy: Tensor, t, cond, batched_number: int, transformer_options):
def get_control_advanced(self, x_noisy: Tensor, t, cond, batched_number: int, transformer_options: dict):
# normal ControlNet stuff
control_prev = None
if self.previous_controlnet is not None:
@@ -368,7 +364,9 @@ class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase):
self.latent_dims_div4 = (new_h, new_w)
# prepare mask
self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number)
# done preparing; model patches will take care of everything now.
# done preparing; model patches will take care of everything now
set_model_attn1_patch(transformer_options, self.patch_attn1.set_control(self))
set_model_attn2_patch(transformer_options, self.patch_attn2.set_control(self))
# return normal controlnet stuff
return control_prev
-18
View File
@@ -17,7 +17,6 @@ from .control_reference import (ReferenceAdvanced, ReferenceInjections,
_forward_inject_BasicTransformerBlock, factory_forward_inject_UNetModel,
handle_context_ref_setup,
REF_CONTROL_LIST_ALL, CONTEXTREF_CLEAN_FUNC)
from .control_lllite import (ControlLLLiteAdvanced)
from .utils import torch_dfs, WrapperConsts
@@ -63,14 +62,6 @@ def get_refcn(control: ControlBase, order: int=-1):
ref_set.update(get_refcn(control.previous_controlnet, order=order))
return ref_set
def get_lllitecn(control: ControlBase):
cn_dict: dict[ControlLLLiteAdvanced,None] = {}
if control is None:
return cn_dict
if type(control) == ControlLLLiteAdvanced:
cn_dict[control] = None
cn_dict.update(get_lllitecn(control.previous_controlnet))
return cn_dict
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,
@@ -114,19 +105,10 @@ def acn_outer_sample_wrapper(executor, *args, **kwargs):
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:
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())
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(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)