sdpa update

This commit is contained in:
kijai
2024-09-30 13:35:02 +03:00
parent 2a5fe38004
commit e39ecb6f88
2 changed files with 23 additions and 28 deletions
+1 -4
View File
@@ -37,12 +37,11 @@ def get_sdpa_settings():
return old_gpu, use_flash_attn, math_kernel_on
OLD_GPU, USE_FLASH_ATTN, MATH_KERNEL_ON = get_sdpa_settings()
backends.append(SDPBackend.EFFICIENT_ATTENTION)
if USE_FLASH_ATTN:
backends.append(SDPBackend.FLASH_ATTENTION)
if MATH_KERNEL_ON:
backends.append(SDPBackend.MATH)
if OLD_GPU or not USE_FLASH_ATTN:
backends.append(SDPBackend.EFFICIENT_ATTENTION)
def remove_all_hooks(model: torch.nn.Module) -> None:
for child in model.children():
@@ -206,8 +205,6 @@ class Reference(Operator):
q, k, v = map(lambda t: rearrange(t, "b n (h d) -> b h n d", h=h), (q, k, v))
# Attention
N = q.shape[-2]
if not MATH_KERNEL_ON and OLD_GPU:
backends.append(SDPBackend.MATH)
with sdpa_kernel(backends):
if layer_ind > 12 or self.mode == 'normal':
attn_bias = None
+22 -24
View File
@@ -17,31 +17,19 @@ ops = comfy.ops.manual_cast
if version.parse(torch.__version__) >= version.parse("2.0.0"):
SDP_IS_AVAILABLE = True
from torch.backends.cuda import SDPBackend, sdp_kernel
from torch.nn.attention import SDPBackend, sdpa_kernel
BACKEND_MAP = {
SDPBackend.MATH: {
"enable_math": True,
"enable_flash": False,
"enable_mem_efficient": False,
},
SDPBackend.FLASH_ATTENTION: {
"enable_math": False,
"enable_flash": True,
"enable_mem_efficient": False,
},
SDPBackend.EFFICIENT_ATTENTION: {
"enable_math": False,
"enable_flash": False,
"enable_mem_efficient": True,
},
None: {"enable_math": True, "enable_flash": True, "enable_mem_efficient": True},
SDPBackend.MATH: [SDPBackend.MATH],
SDPBackend.FLASH_ATTENTION: [SDPBackend.FLASH_ATTENTION],
SDPBackend.EFFICIENT_ATTENTION: [SDPBackend.EFFICIENT_ATTENTION],
None: [SDPBackend.MATH, SDPBackend.FLASH_ATTENTION, SDPBackend.EFFICIENT_ATTENTION]
}
else:
from contextlib import nullcontext
SDP_IS_AVAILABLE = False
sdp_kernel = nullcontext
sdpa_kernel = nullcontext
BACKEND_MAP = {}
logpy.warn(
f"No SDP backend available, likely because you are running in pytorch "
@@ -60,7 +48,7 @@ except:
# from .diffusionmodules.util import mixed_checkpoint as checkpoint
backends = []
def exists(val):
return val is not None
@@ -346,7 +334,9 @@ class CrossAttention(nn.Module):
out = einsum('b i j, b j d -> b i d', sim, v)
"""
## new
with sdp_kernel(**BACKEND_MAP[self.backend]):
backends.extend(BACKEND_MAP[self.backend])
with sdpa_kernel(backends):
# print("dispatching into backend", self.backend, "q/k/v shape: ", q.shape, k.shape, v.shape)
out = F.scaled_dot_product_attention(
q, k, v, attn_mask=mask
@@ -440,7 +430,9 @@ class CrossAttention_Masked(nn.Module):
q, k, v = map(lambda t: rearrange(t, "b n (h d) -> b h n d", h=h), (q, k, v))
with sdp_kernel(**BACKEND_MAP[self.backend]):
backends.extend(BACKEND_MAP[self.backend])
with sdpa_kernel(backends):
# print("dispatching into backend", self.backend, "q/k/v shape: ", q.shape, k.shape, v.shape)
out = F.scaled_dot_product_attention(
q, k, v, attn_mask=mask
@@ -530,7 +522,9 @@ class CrossAttention_Cond(nn.Module):
q, k, v = map(lambda t: rearrange(t, "b n (h d) -> b h n d", h=h), (q, k, v))
with sdp_kernel(**BACKEND_MAP[self.backend]):
backends.extend(BACKEND_MAP[self.backend])
with sdpa_kernel(backends):
# print("dispatching into backend", self.backend, "q/k/v shape: ", q.shape, k.shape, v.shape)
out = F.scaled_dot_product_attention(
q, k, v, attn_mask=mask
@@ -737,7 +731,9 @@ class SelfAttentionRefconcatFirst(nn.Module):
q, k, v = map(lambda t: rearrange(t, "b n (h d) -> b h n d", h=h), (q, k, v))
with sdp_kernel(**BACKEND_MAP[self.backend]):
backends.extend(BACKEND_MAP[self.backend])
with sdpa_kernel(backends):
# print("dispatching into backend", self.backend, "q/k/v shape: ", q.shape, k.shape, v.shape)
out = F.scaled_dot_product_attention(
q, k, v, attn_mask=mask
@@ -833,7 +829,9 @@ class SelfAttentionRefonlyFirst(nn.Module):
q, k, v = map(lambda t: rearrange(t, "b n (h d) -> b h n d", h=h), (q, k, v))
with sdp_kernel(**BACKEND_MAP[self.backend]):
backends.extend(BACKEND_MAP[self.backend])
with sdpa_kernel(backends):
# print("dispatching into backend", self.backend, "q/k/v shape: ", q.shape, k.shape, v.shape)
out = F.scaled_dot_product_attention(
q, k, v, attn_mask=mask