From 326b26ed959a155c05666beff0f2bb134bb93949 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Thu, 14 Nov 2024 12:10:28 -0600 Subject: [PATCH] Simplified ControlLLLite implementation, since transformer_options are passed into get_control with ComfyUI rework --- adv_control/control_lllite.py | 22 ++++++++++------------ adv_control/sampling.py | 18 ------------------ 2 files changed, 10 insertions(+), 30 deletions(-) diff --git a/adv_control/control_lllite.py b/adv_control/control_lllite.py index 60e0a5f..465b571 100644 --- a/adv_control/control_lllite.py +++ b/adv_control/control_lllite.py @@ -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 diff --git a/adv_control/sampling.py b/adv_control/sampling.py index 59f2ce9..e6eb3ea 100644 --- a/adv_control/sampling.py +++ b/adv_control/sampling.py @@ -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)