diff --git a/modules/intel/ipex/attention.py b/modules/intel/ipex/attention.py index eacb6a7d0..4f74a2c2e 100644 --- a/modules/intel/ipex/attention.py +++ b/modules/intel/ipex/attention.py @@ -120,10 +120,10 @@ def dynamic_scaled_dot_product_attention(query, key, value, attn_mask=None, drop value[start_idx:end_idx, :, :, :], 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, **kwargs - ) - if is_unsqueezed: - hidden_states.squeeze(0) + ) 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, **kwargs) + hidden_states = original_scaled_dot_product_attention(query, key, value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, **kwargs) + if is_unsqueezed: + hidden_states.squeeze(0) return hidden_states diff --git a/modules/sd_hijack_dynamic_atten.py b/modules/sd_hijack_dynamic_atten.py index 304f3f355..49f89b414 100644 --- a/modules/sd_hijack_dynamic_atten.py +++ b/modules/sd_hijack_dynamic_atten.py @@ -115,12 +115,12 @@ def dynamic_scaled_dot_product_attention(query, key, value, attn_mask=None, drop 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, **kwargs ) - if is_unsqueezed: - hidden_states.squeeze(0) if devices.backend != "directml": getattr(torch, query.device.type).synchronize() else: return devices.sdpa_pre_dyanmic_atten(query, key, value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, **kwargs) + if is_unsqueezed: + hidden_states.squeeze(0) return hidden_states