Simplified ControlLLLite implementation, since transformer_options are passed into get_control with ComfyUI rework
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user