diff --git a/modules/attention.py b/modules/attention.py index 96296a687..7776af719 100644 --- a/modules/attention.py +++ b/modules/attention.py @@ -16,13 +16,14 @@ def set_dynamic_attention(): log.error(f'Torch attention: type="dynamic attention" {err}') return None + def set_triton_flash_attention(backend: str): try: if backend in {"rocm", "zluda"}: # flash_attn_triton_amd only works with AMD from modules.flash_attn_triton_amd import interface_fa sdpa_pre_triton_flash_atten = torch.nn.functional.scaled_dot_product_attention @wraps(sdpa_pre_triton_flash_atten) - def sdpa_triton_flash_atten(query: torch.FloatTensor, key: torch.FloatTensor, value: torch.FloatTensor, attn_mask: Optional[torch.FloatTensor] = None, dropout_p: float = 0.0, is_causal: bool = False, scale: Optional[float] = None, enable_gqa: bool = False, **kwargs) -> torch.FloatTensor: + def sdpa_triton_flash_atten(query: torch.FloatTensor, key: torch.FloatTensor, value: torch.FloatTensor, attn_mask: Optional[torch.Tensor] = None, dropout_p: float = 0.0, is_causal: bool = False, scale: Optional[float] = None, enable_gqa: bool = False, **kwargs) -> torch.FloatTensor: if query.shape[-1] <= 128 and attn_mask is None and query.dtype != torch.float32: if scale is None: scale = query.shape[-1] ** (-0.5) @@ -46,6 +47,42 @@ def set_triton_flash_attention(backend: str): except Exception as err: log.error(f'Torch attention: type="Triton Flash attention" {err}') + +def set_flex_attention(): + try: + from torch.nn.attention.flex_attention import flex_attention, create_block_mask + def flex_attention_causal_mask(b, h, q_idx, kv_idx): + return q_idx >= kv_idx + + sdpa_pre_flex_atten = torch.nn.functional.scaled_dot_product_attention + @wraps(sdpa_pre_flex_atten) + def sdpa_flex_atten(query: torch.FloatTensor, key: torch.FloatTensor, value: torch.FloatTensor, attn_mask: Optional[torch.Tensor] = None, dropout_p: float = 0.0, is_causal: bool = False, scale: Optional[float] = None, enable_gqa: bool = False, **kwargs) -> torch.FloatTensor: + score_mod = None + block_mask = None + if attn_mask is not None: + batch_size, num_heads = query.shape[:2] + seq_len_q = query.shape[-2] + seq_len_kv = key.shape[-2] + if attn_mask.ndim == 2: + attn_mask = attn_mask.view(attn_mask.shape[0], 1, attn_mask.size[1], 1) + attn_mask = attn_mask.expand(batch_size, num_heads, seq_len_q, seq_len_kv) + if attn_mask.dtype == torch.bool: + def mask_mod(batch_idx, head_idx, q_idx, kv_idx): + return attn_mask[batch_idx, head_idx, q_idx, kv_idx] + block_mask = create_block_mask(mask_mod, batch_size, None, seq_len_q, seq_len_kv, device=query.device) + else: + def score_mod(score, batch_idx, head_idx, q_idx, kv_idx): + return score + attn_mask[batch_idx, head_idx, q_idx, kv_idx] + elif is_causal: + block_mask = create_block_mask(flex_attention_causal_mask, query.shape[0], query.shape[1], query.shape[-2], key.shape[-2], device=query.device) + return flex_attention(query, key, value, score_mod=score_mod, block_mask=block_mask, scale=scale, enable_gqa=enable_gqa) + + torch.nn.functional.scaled_dot_product_attention = sdpa_flex_atten + log.debug('Torch attention: type="Flex attention"') + except Exception as err: + log.error(f'Torch attention: type="Flex attention" {err}') + + def set_ck_flash_attention(backend: str, device: torch.device): try: if backend == "rocm": @@ -58,7 +95,7 @@ def set_ck_flash_attention(backend: str, device: torch.device): from flash_attn import flash_attn_func sdpa_pre_flash_atten = torch.nn.functional.scaled_dot_product_attention @wraps(sdpa_pre_flash_atten) - def sdpa_flash_atten(query: torch.FloatTensor, key: torch.FloatTensor, value: torch.FloatTensor, attn_mask: Optional[torch.FloatTensor] = None, dropout_p: float = 0.0, is_causal: bool = False, scale: Optional[float] = None, enable_gqa: bool = False, **kwargs) -> torch.FloatTensor: + def sdpa_flash_atten(query: torch.FloatTensor, key: torch.FloatTensor, value: torch.FloatTensor, attn_mask: Optional[torch.Tensor] = None, dropout_p: float = 0.0, is_causal: bool = False, scale: Optional[float] = None, enable_gqa: bool = False, **kwargs) -> torch.FloatTensor: if query.shape[-1] <= 128 and attn_mask is None and query.dtype != torch.float32: is_unsqueezed = False if query.dim() == 3: @@ -87,6 +124,7 @@ def set_ck_flash_attention(backend: str, device: torch.device): except Exception as err: log.error(f'Torch attention: type="Flash attention" {err}') + def set_sage_attention(backend: str, device: torch.device): try: install('sageattention') @@ -96,7 +134,7 @@ def set_sage_attention(backend: str, device: torch.device): use_cuda_backend = True # Detect GPU architecture - sm86 confirmed to need CUDA backend workaround as Sage Attention + Triton causes NaNs try: from sageattention import sageattn_qk_int8_pv_fp16_cuda - except: + except Exception: use_cuda_backend = False if use_cuda_backend: @@ -123,7 +161,7 @@ def set_sage_attention(backend: str, device: torch.device): sdpa_pre_sage_atten = torch.nn.functional.scaled_dot_product_attention @wraps(sdpa_pre_sage_atten) - def sdpa_sage_atten(query: torch.FloatTensor, key: torch.FloatTensor, value: torch.FloatTensor, attn_mask: Optional[torch.FloatTensor] = None, dropout_p: float = 0.0, is_causal: bool = False, scale: Optional[float] = None, enable_gqa: bool = False, **kwargs) -> torch.FloatTensor: + def sdpa_sage_atten(query: torch.FloatTensor, key: torch.FloatTensor, value: torch.FloatTensor, attn_mask: Optional[torch.Tensor] = None, dropout_p: float = 0.0, is_causal: bool = False, scale: Optional[float] = None, enable_gqa: bool = False, **kwargs) -> torch.FloatTensor: if (query.shape[-1] in {128, 96, 64}) and (attn_mask is None) and (query.dtype != torch.float32): if enable_gqa: key = key.repeat_interleave(query.size(-3)//key.size(-3), -3) diff --git a/modules/devices.py b/modules/devices.py index ed29d5421..e6f21194c 100644 --- a/modules/devices.py +++ b/modules/devices.py @@ -474,6 +474,9 @@ def set_sdpa_params(): global sdpa_pre_dyanmic_atten # pylint: disable=global-statement sdpa_pre_dyanmic_atten = attention.set_dynamic_attention() + if 'Flex attention' in opts.sdp_overrides: + attention.set_flex_attention() + if 'Triton Flash attention' in opts.sdp_overrides: attention.set_triton_flash_attention(backend) diff --git a/modules/shared_defaults.py b/modules/shared_defaults.py index fef7462fd..9eac6dc5e 100644 --- a/modules/shared_defaults.py +++ b/modules/shared_defaults.py @@ -43,7 +43,7 @@ def get_default_modes(cmd_opts, mem_stat): default_sdp_choices = ['Flash', 'Memory', 'Math'] default_sdp_options = ['Flash', 'Memory', 'Math'] - default_sdp_override_choices = ['Dynamic attention', 'Flash attention', 'Sage attention'] + default_sdp_override_choices = ['Dynamic attention', 'Flex attention', 'Flash attention', 'Sage attention'] default_sdp_override_options = [] if devices.backend == "zluda":