From 0ac2dfbcaaf3a226e51c87860f57ae6118a9cd4d Mon Sep 17 00:00:00 2001 From: Disty0 Date: Sat, 10 Feb 2024 20:02:29 +0300 Subject: [PATCH] Dyn Atten don't synchronize if not slicing --- modules/intel/ipex/attention.py | 2 +- modules/sd_hijack_dynamic_atten.py | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/modules/intel/ipex/attention.py b/modules/intel/ipex/attention.py index b46b3ab61..1ca5dd20d 100644 --- a/modules/intel/ipex/attention.py +++ b/modules/intel/ipex/attention.py @@ -171,7 +171,7 @@ def scaled_dot_product_attention_32_bit(query, key, value, attn_mask=None, dropo attn_mask=attn_mask[start_idx:end_idx] if attn_mask is not None else attn_mask, dropout_p=dropout_p, is_causal=is_causal ) + torch.xpu.synchronize(query.device) 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(query.device) return hidden_states diff --git a/modules/sd_hijack_dynamic_atten.py b/modules/sd_hijack_dynamic_atten.py index 378528e89..d9881beb7 100644 --- a/modules/sd_hijack_dynamic_atten.py +++ b/modules/sd_hijack_dynamic_atten.py @@ -90,10 +90,10 @@ def sliced_scaled_dot_product_attention(query, key, value, attn_mask=None, dropo attn_mask=attn_mask[start_idx:end_idx] if attn_mask is not None else attn_mask, dropout_p=dropout_p, is_causal=is_causal ) + if devices.backend != "directml": + getattr(torch, query.device.type).synchronize() else: return F.scaled_dot_product_attention(query, key, value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal) - if devices.backend != "directml": - getattr(torch, query.device.type).synchronize() return hidden_states