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