diff --git a/adv_control/control.py b/adv_control/control.py index f84e011..ee19d9c 100644 --- a/adv_control/control.py +++ b/adv_control/control.py @@ -301,7 +301,7 @@ class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase): def pre_run_advanced(self, *args, **kwargs): AdvancedControlBase.pre_run_advanced(self, *args, **kwargs) - self.patch.control = self + self.patch.set_control(self) def get_control_advanced(self, x_noisy: Tensor, t, cond, batched_number: int): # normal ControlNet stuff @@ -341,6 +341,10 @@ class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase): self.copy_to(c) self.copy_to_advanced(c) return c + + # deepcopy needs to properly keep track of objects to work between model.clone calls! + def __deepcopy__(self, *args, **kwargs): + return self # def get_models(self): # # get_models is called once at the start of every KSampler run - use to reset already_patched status @@ -358,7 +362,6 @@ def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroup=None, mo for key in controlnet_data: # LLLLite check if "lllite" in key: - logger.info("ControlLLLite controlnet!") controlnet_type = ControlWeightType.CONTROLLLLITE break # SparseCtrl check @@ -596,7 +599,7 @@ def load_controllllite(ckpt_path: str, controlnet_data: dict[str, Tensor]=None, if len(modules) == 1: module.is_first = True - logger.info(f"loaded {ckpt_path} successfully, {len(modules)} modules") + #logger.info(f"loaded {ckpt_path} successfully, {len(modules)} modules") patch = LLLitePatch(modules=modules) control = ControlLLLiteAdvanced(patch=patch, timestep_keyframes=timestep_keyframe) diff --git a/adv_control/control_lllite.py b/adv_control/control_lllite.py index e92a416..c5b7928 100644 --- a/adv_control/control_lllite.py +++ b/adv_control/control_lllite.py @@ -48,7 +48,6 @@ class LLLitePatch: # it turns out comparing single-value tensors to floats is extremely slow # a: Tensor = extra_options["sigmas"][0] if self.control.t > self.control.timestep_range[0] or self.control.t < self.control.timestep_range[1]: - logger.info("Stopping short!!!") return q, k, v module_pfx = extra_options_to_module_prefix(extra_options) @@ -63,11 +62,6 @@ class LLLitePatch: module_pfx_to_k = module_pfx + "_to_k" module_pfx_to_v = module_pfx + "_to_v" - # if masks present, get masks with same dims as attention - # if q.shape != k.shape or q.shape != v.shape: - # logger.warn(f"mismatch!!! q:{q.shape}, k:{k.shape}, v:{v.shape}") - #logger.warn(f"{q.shape}") - if module_pfx_to_q in self.modules: q = q + self.modules[module_pfx_to_q](q, self.control) if module_pfx_to_k in self.modules: diff --git a/adv_control/control_sparsectrl.py b/adv_control/control_sparsectrl.py index 9d51e6e..5885ed5 100644 --- a/adv_control/control_sparsectrl.py +++ b/adv_control/control_sparsectrl.py @@ -17,19 +17,35 @@ from comfy.ldm.modules.diffusionmodules.util import ( timestep_embedding, ) +from comfy.cli_args import args from comfy.cldm.cldm import ControlNet as ControlNetCLDM from comfy.ldm.modules.attention import SpatialTransformer -from comfy.ldm.modules.diffusionmodules.openaimodel import TimestepEmbedSequential, ResBlock, Downsample -from comfy.ldm.util import exists -from comfy.ldm.modules.attention import default, optimized_attention +from comfy.ldm.modules.attention import attention_basic, attention_pytorch, attention_split, attention_sub_quad, default from comfy.ldm.modules.attention import FeedForward, SpatialTransformer +from comfy.ldm.modules.diffusionmodules.openaimodel import TimestepEmbedSequential, ResBlock, Downsample from comfy.controlnet import broadcast_image_to from comfy.utils import repeat_to_batch_size import comfy.ops +import comfy.model_management from .utils import TimestepKeyframeGroup, disable_weight_init_clean_groupnorm, prepare_mask_batch +# until xformers bug is fixed, do not use xformers for VersatileAttention! TODO: change this when fix is out +# logic for choosing optimized_attention method taken from comfy/ldm/modules/attention.py +optimized_attention_mm = attention_basic +if comfy.model_management.xformers_enabled(): + pass + #optimized_attention_mm = attention_xformers +if comfy.model_management.pytorch_attention_enabled(): + optimized_attention_mm = attention_pytorch +else: + if args.use_split_cross_attention: + optimized_attention_mm = attention_split + else: + optimized_attention_mm = attention_sub_quad + + class SparseControlNet(ControlNetCLDM): def __init__(self, *args,**kwargs): super().__init__(*args, **kwargs) @@ -810,7 +826,7 @@ class CrossAttentionMM(nn.Module): if scale_mask is not None: k *= scale_mask - out = optimized_attention(q, k, v, self.heads, mask) + out = optimized_attention_mm(q, k, v, self.heads, mask) return self.to_out(out)