Merge PR #61 from Kosinkadink/develop SparseCtrl xformers fix

Fixed xformers issue with SparseCtrl after latest xformers update (will not use xformers ever for motion module attn)
This commit is contained in:
Jedrzej Kosinski
2024-02-02 02:45:55 -06:00
committed by GitHub
3 changed files with 26 additions and 13 deletions
+6 -3
View File
@@ -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)
-6
View File
@@ -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:
+20 -4
View File
@@ -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)