sdpa update
This commit is contained in:
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user