From 1b0c06bc71e0ca0030c8235e8997f341cbf5d51f Mon Sep 17 00:00:00 2001 From: Disty0 Date: Wed, 10 Sep 2025 05:20:49 +0300 Subject: [PATCH] IPEX add enable_gqa support and add result dtype check to sdpa --- modules/intel/ipex/hijacks.py | 11 +++++++++-- 1 file changed, 9 insertions(+), 2 deletions(-) diff --git a/modules/intel/ipex/hijacks.py b/modules/intel/ipex/hijacks.py index eb8936c3a..58aeae9e0 100644 --- a/modules/intel/ipex/hijacks.py +++ b/modules/intel/ipex/hijacks.py @@ -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