diff --git a/inference/model_hack.py b/inference/model_hack.py index e3cd2f4..df32238 100644 --- a/inference/model_hack.py +++ b/inference/model_hack.py @@ -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 diff --git a/sgm/modules/attention.py b/sgm/modules/attention.py index 5858d75..a93e647 100644 --- a/sgm/modules/attention.py +++ b/sgm/modules/attention.py @@ -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