Added fallback mechanism for attention to work around pytorch's CUDA compiled limit for scaled_dot_product_attention

This commit is contained in:
Jedrzej Kosinski
2024-04-23 02:22:10 -05:00
parent 1fde1be80e
commit 6ca21f2bb6
4 changed files with 31 additions and 4 deletions
+9
View File
@@ -781,6 +781,8 @@ class TemporalTransformer3DModel(nn.Module):
del self.temp_cameractrl_effect
self.temp_cameractrl_effect = None
self.prev_cameractrl_hidden_states_batch = 0
for block in self.transformer_blocks:
block.reset_temp_vars()
def get_scale_mask(self, hidden_states: Tensor) -> Union[Tensor, None]:
# if no raw mask, return None
@@ -972,6 +974,10 @@ class TemporalTransformerBlock(nn.Module):
for block in self.attention_blocks:
block.set_sub_idxs(sub_idxs)
def reset_temp_vars(self):
for block in self.attention_blocks:
block.reset_temp_vars()
def forward(
self,
hidden_states: Tensor,
@@ -1154,6 +1160,9 @@ class VersatileAttention(CrossAttentionMM):
def init_qkv_merge(self, ops=comfy.ops.disable_weight_init):
self.qkv_merge = zero_module(ops.Linear(in_features=self.query_dim, out_features=self.query_dim))
def reset_temp_vars(self):
self.reset_attention_type()
def forward(
self,
hidden_states: Tensor,
+3 -3
View File
@@ -21,8 +21,8 @@ from comfy.controlnet import ControlBase
import comfy.ops
from .context import ContextFuseMethod, ContextSchedules, get_context_weights, get_context_windows
from .sample_settings import IterationOptions, SampleSettings, SeedNoiseGeneration, prepare_mask_ad
from .utils_model import ModelTypeSD, wrap_function_to_inject_xformers_bug_info
from .sample_settings import IterationOptions, SampleSettings, SeedNoiseGeneration
from .utils_model import ModelTypeSD
from .model_injection import InjectionParams, ModelPatcherAndInjector, MotionModelGroup, MotionModelPatcher
from .motion_module_ad import AnimateDiffFormat, AnimateDiffInfo, AnimateDiffVersion, VanillaTemporalModule
from .logger import logger
@@ -380,7 +380,7 @@ def motion_sample_factory(orig_comfy_sample: Callable, is_custom: bool=False) ->
model.motion_models.pre_run(model)
if model.sample_settings is not None:
model.sample_settings.pre_run(model)
latents = wrap_function_to_inject_xformers_bug_info(orig_comfy_sample)(model, noise, *args, **kwargs)
latents = orig_comfy_sample(model, noise, *args, **kwargs)
return latents
finally:
del latents
+1
View File
@@ -370,6 +370,7 @@ def raise_if_not_checkpoint_sd1_5(model: ModelPatcher):
# TODO: remove this filth when xformers bug gets fixed in future xformers version
# NOTE: avoid using this for now to avoid false positives with pytorch or non-AD stuff like SVD
def wrap_function_to_inject_xformers_bug_info(function_to_wrap: Callable) -> Callable:
if not xformers_enabled:
return function_to_wrap
+18 -1
View File
@@ -14,12 +14,18 @@ from .logger import logger
# 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
# a fallback_attention_mm is selected to avoid CUDA configuration limitation with pytorch's scaled_dot_product
optimized_attention_mm = attention_basic
fallback_attention_mm = attention_basic
if model_management.xformers_enabled():
pass
#optimized_attention_mm = attention_xformers
if model_management.pytorch_attention_enabled():
optimized_attention_mm = attention_pytorch
if args.use_split_cross_attention:
fallback_attention_mm = attention_split
else:
fallback_attention_mm = attention_sub_quad
else:
if args.use_split_cross_attention:
optimized_attention_mm = attention_split
@@ -34,6 +40,7 @@ class CrossAttentionMM(nn.Module):
inner_dim = dim_head * heads
context_dim = default(context_dim, query_dim)
self.actual_attention = optimized_attention_mm
self.heads = heads
self.dim_head = dim_head
self.scale = None
@@ -45,6 +52,9 @@ class CrossAttentionMM(nn.Module):
self.to_out = nn.Sequential(operations.Linear(inner_dim, query_dim, dtype=dtype, device=device), nn.Dropout(dropout))
def reset_attention_type(self):
self.actual_attention = optimized_attention_mm
def forward(self, x, context=None, value=None, mask=None, scale_mask=None):
q = self.to_q(x)
context = default(context, x)
@@ -63,7 +73,14 @@ class CrossAttentionMM(nn.Module):
if scale_mask is not None:
k *= scale_mask
out = optimized_attention_mm(q, k, v, self.heads, mask)
try:
out = self.actual_attention(q, k, v, self.heads, mask)
except RuntimeError as e:
if str(e).startswith("CUDA error: invalid configuration argument"):
self.actual_attention = fallback_attention_mm
out = self.actual_attention(q, k, v, self.heads, mask)
else:
raise
return self.to_out(out)
# TODO: set up comfy.ops style classes for groupnorm and other functions