From 6fbb9ad7f39bdb3c23622568ca77af2277ef1794 Mon Sep 17 00:00:00 2001 From: Disty0 Date: Mon, 8 Jan 2024 20:31:54 +0300 Subject: [PATCH] IPEX fix lock-ups at very high resolutions --- modules/intel/ipex/attention.py | 2 ++ modules/intel/ipex/diffusers.py | 2 ++ 2 files changed, 4 insertions(+) diff --git a/modules/intel/ipex/attention.py b/modules/intel/ipex/attention.py index e9f927a9c..ce955fc0b 100644 --- a/modules/intel/ipex/attention.py +++ b/modules/intel/ipex/attention.py @@ -124,6 +124,7 @@ def torch_bmm_32_bit(input, mat2, *, out=None): ) else: return original_torch_bmm(input, mat2, out=out) + torch.xpu.synchronize() return hidden_states original_scaled_dot_product_attention = torch.nn.functional.scaled_dot_product_attention @@ -172,4 +173,5 @@ def scaled_dot_product_attention_32_bit(query, key, value, attn_mask=None, dropo ) else: return original_scaled_dot_product_attention(query, key, value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal) + torch.xpu.synchronize() return hidden_states diff --git a/modules/intel/ipex/diffusers.py b/modules/intel/ipex/diffusers.py index 47b0375ae..48cf11b73 100644 --- a/modules/intel/ipex/diffusers.py +++ b/modules/intel/ipex/diffusers.py @@ -149,6 +149,7 @@ class SlicedAttnProcessor: # pylint: disable=too-few-public-methods hidden_states[start_idx:end_idx, start_idx_2:end_idx_2] = attn_slice del attn_slice + torch.xpu.synchronize() else: query_slice = query[start_idx:end_idx] key_slice = key[start_idx:end_idx] @@ -283,6 +284,7 @@ class AttnProcessor: hidden_states[start_idx:end_idx] = attn_slice del attn_slice + torch.xpu.synchronize() else: attention_probs = attn.get_attention_scores(query, key, attention_mask) hidden_states = torch.bmm(attention_probs, value)