From 31b66cfa4a9e2c56fa50ae52e3dc8b86439a3a8b Mon Sep 17 00:00:00 2001 From: Disty0 Date: Mon, 23 Sep 2024 21:43:01 +0300 Subject: [PATCH] IPEX fix FreeU --- modules/intel/ipex/__init__.py | 13 +++++-------- modules/intel/ipex/diffusers.py | 21 ++++++++++++++++++--- modules/intel/ipex/hijacks.py | 16 ++++++++++++++++ 3 files changed, 39 insertions(+), 11 deletions(-) diff --git a/modules/intel/ipex/__init__.py b/modules/intel/ipex/__init__.py index 189dd07d0..b84a853ed 100644 --- a/modules/intel/ipex/__init__.py +++ b/modules/intel/ipex/__init__.py @@ -16,8 +16,6 @@ def ipex_init(): # pylint: disable=too-many-statements if hasattr(torch, "cuda") and hasattr(torch.cuda, "is_xpu_hijacked") and torch.cuda.is_xpu_hijacked: return True, "Skipping IPEX hijack" else: - device_supports_fp64 = torch.xpu.has_fp64_dtype() if hasattr(torch.xpu, "has_fp64_dtype") else torch.xpu.get_device_properties("xpu").has_fp64 - # Replace cuda with xpu: torch.cuda.current_device = torch.xpu.current_device torch.cuda.current_stream = torch.xpu.current_stream @@ -201,12 +199,11 @@ def ipex_init(): # pylint: disable=too-many-statements torch.cuda.utilization = lambda *args, **kwargs: 0 ipex_hijacks() - if not device_supports_fp64 or os.environ.get('IPEX_FORCE_ATTENTION_SLICE', None) is not None: - try: - from .diffusers import ipex_diffusers - ipex_diffusers() - except Exception: # pylint: disable=broad-exception-caught - pass + try: + from .diffusers import ipex_diffusers + ipex_diffusers() + except Exception: # pylint: disable=broad-exception-caught + pass torch.cuda.is_xpu_hijacked = True except Exception as e: return False, e diff --git a/modules/intel/ipex/diffusers.py b/modules/intel/ipex/diffusers.py index 613f90cca..55908c7cf 100644 --- a/modules/intel/ipex/diffusers.py +++ b/modules/intel/ipex/diffusers.py @@ -1,4 +1,5 @@ import os +from functools import wraps import torch import diffusers #0.29.1 # pylint: disable=import-error from diffusers.models.attention_processor import Attention @@ -11,6 +12,16 @@ from functools import cache device_supports_fp64 = torch.xpu.has_fp64_dtype() if hasattr(torch.xpu, "has_fp64_dtype") else torch.xpu.get_device_properties("xpu").has_fp64 attention_slice_rate = float(os.environ.get('IPEX_ATTENTION_SLICE_RATE', 4)) + +# Force FP32 upcast +# diffusers is imported before ipex hijacks and doesn't apply so hijack this separately +original_fourier_filter = diffusers.utils.torch_utils.fourier_filter +@wraps(diffusers.utils.torch_utils.fourier_filter) +def fourier_filter(x_in, threshold, scale): + return_dtype = x_in.dtype + return original_fourier_filter(x_in.to(dtype=torch.float32), threshold, scale).to(dtype=return_dtype) + + # fp64 error def rope(pos: torch.Tensor, dim: int, theta: int) -> torch.Tensor: assert dim % 2 == 0, "The dimension must be even." @@ -27,6 +38,7 @@ def rope(pos: torch.Tensor, dim: int, theta: int) -> torch.Tensor: out = stacked_out.view(batch_size, -1, dim // 2, 2, 2) return out.float() + @cache def find_slice_size(slice_size, slice_block_size): while (slice_size * slice_block_size) > attention_slice_rate: @@ -74,6 +86,7 @@ def find_attention_slice_sizes(query_shape, query_element_size, query_device_typ return do_split, do_split_2, do_split_3, split_slice_size, split_2_slice_size, split_3_slice_size + class SlicedAttnProcessor: # pylint: disable=too-few-public-methods r""" Processor for implementing sliced attention. @@ -320,9 +333,11 @@ class AttnProcessor: return hidden_states + def ipex_diffusers(): + diffusers.utils.torch_utils.fourier_filter = fourier_filter #ARC GPUs can't allocate more than 4GB to a single block: - diffusers.models.attention_processor.SlicedAttnProcessor = SlicedAttnProcessor - diffusers.models.attention_processor.AttnProcessor = AttnProcessor - if not device_supports_fp64 and hasattr(transformers, "transformer_flux"): + if not device_supports_fp64 or os.environ.get('IPEX_FORCE_ATTENTION_SLICE', None) is not None: + diffusers.models.attention_processor.SlicedAttnProcessor = SlicedAttnProcessor + diffusers.models.attention_processor.AttnProcessor = AttnProcessor diffusers.models.transformers.transformer_flux.rope = rope diff --git a/modules/intel/ipex/hijacks.py b/modules/intel/ipex/hijacks.py index e5a71f5e2..7d2b2cb31 100644 --- a/modules/intel/ipex/hijacks.py +++ b/modules/intel/ipex/hijacks.py @@ -105,6 +105,20 @@ def scaled_dot_product_attention(query, key, value, attn_mask=None, dropout_p=0. 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) +# Diffusers FreeU +original_fft_fftn = torch.fft.fftn +@wraps(torch.fft.fftn) +def fft_fftn(input, s=None, dim=None, norm=None, *, out=None): + return_dtype = input.dtype + return original_fft_fftn(input.to(dtype=torch.float32), s=s, dim=dim, norm=norm, out=out).to(dtype=return_dtype) + +# Diffusers FreeU +original_fft_ifftn = torch.fft.ifftn +@wraps(torch.fft.ifftn) +def fft_ifftn(input, s=None, dim=None, norm=None, *, out=None): + return_dtype = input.dtype + return original_fft_ifftn(input.to(dtype=torch.float32), s=s, dim=dim, norm=norm, out=out).to(dtype=return_dtype) + # A1111 FP16 original_functional_group_norm = torch.nn.functional.group_norm @wraps(torch.nn.functional.group_norm) @@ -309,6 +323,8 @@ def ipex_hijacks(): torch.nn.functional.pad = functional_pad torch.bmm = torch_bmm + torch.fft.fftn = fft_fftn + torch.fft.ifftn = fft_ifftn if not device_supports_fp64: torch.from_numpy = from_numpy torch.as_tensor = as_tensor