Implement fallback_attention_mm for SparseCtrl to mirror AnimateDiff-Evolved

This commit is contained in:
Jedrzej Kosinski
2024-05-23 15:59:50 -05:00
parent 67884487aa
commit c808e875c4
+30 -3
View File
@@ -3,6 +3,7 @@
#and then taken from comfy/cldm/cldm.py and modified again
from abc import ABC, abstractmethod
import copy
import math
import numpy as np
from typing import Iterable, Union
@@ -34,12 +35,18 @@ from .utils import TimestepKeyframeGroup, disable_weight_init_clean_groupnorm, p
# 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 comfy.model_management.xformers_enabled():
pass
#optimized_attention_mm = attention_xformers
if comfy.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
@@ -641,6 +648,8 @@ class TemporalTransformer3DModel(nn.Module):
del self.temp_scale_mask
self.temp_scale_mask = None
self.prev_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
@@ -779,6 +788,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,
@@ -833,13 +846,14 @@ class PositionalEncoding(nn.Module):
return self.dropout(x)
class CrossAttentionMM(nn.Module):
class CrossAttentionMMSparse(nn.Module):
def __init__(self, query_dim, context_dim=None, heads=8, dim_head=64, dropout=0., dtype=None, device=None,
operations=comfy.ops.disable_weight_init):
super().__init__()
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
@@ -850,6 +864,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)
@@ -868,11 +885,18 @@ 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)
class VersatileAttention(CrossAttentionMM):
class VersatileAttention(CrossAttentionMMSparse):
def __init__(
self,
attention_mode=None,
@@ -911,6 +935,9 @@ class VersatileAttention(CrossAttentionMM):
if self.pos_encoder != None:
self.pos_encoder.set_sub_idxs(sub_idxs)
def reset_temp_vars(self):
self.reset_attention_type()
def forward(
self,
hidden_states: Tensor,