IPEX add enable_gqa support and add result dtype check to sdpa

This commit is contained in:
Disty0
2025-09-10 05:20:49 +03:00
parent de499bcdb6
commit 1b0c06bc71
+9 -2
View File
@@ -1,3 +1,5 @@
from typing import Tuple, Optional
import os
from functools import wraps
from contextlib import nullcontext
@@ -134,14 +136,19 @@ else:
original_scaled_dot_product_attention = torch.nn.functional.scaled_dot_product_attention
@wraps(torch.nn.functional.scaled_dot_product_attention)
def scaled_dot_product_attention(query, key, value, attn_mask=None, dropout_p=0.0, is_causal=False, **kwargs):
def scaled_dot_product_attention(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:
if query.dtype != key.dtype:
key = key.to(dtype=query.dtype)
if query.dtype != value.dtype:
value = value.to(dtype=query.dtype)
if attn_mask is not None and query.dtype != attn_mask.dtype:
attn_mask = attn_mask.to(dtype=query.dtype)
return original_scaled_dot_product_attention(query, key, value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, **kwargs)
if enable_gqa:
kwargs["enable_gqa"] = enable_gqa
result = original_scaled_dot_product_attention(query, key, value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, scale=scale, **kwargs)
if result.dtype != query.dtype:
result = result.to(dtype=query.dtype)
return result
# Data Type Errors:
original_torch_bmm = torch.bmm