From 4de0480d9fa07d00be9d88386eb7bdb5daa7fa63 Mon Sep 17 00:00:00 2001 From: Disty0 Date: Mon, 23 Sep 2024 22:53:40 +0300 Subject: [PATCH] Upcast bf16 fftn to fp32 --- modules/intel/ipex/diffusers.py | 4 ++-- modules/sd_hijack.py | 40 +++++++++++++++++++++++++++++++++ 2 files changed, 42 insertions(+), 2 deletions(-) diff --git a/modules/intel/ipex/diffusers.py b/modules/intel/ipex/diffusers.py index 55908c7cf..50332121b 100644 --- a/modules/intel/ipex/diffusers.py +++ b/modules/intel/ipex/diffusers.py @@ -13,8 +13,8 @@ device_supports_fp64 = torch.xpu.has_fp64_dtype() if hasattr(torch.xpu, "has_fp6 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 +# Diffusers FreeU +# Diffusers is imported before ipex hijacks so fourier_filter needs hijacking too original_fourier_filter = diffusers.utils.torch_utils.fourier_filter @wraps(diffusers.utils.torch_utils.fourier_filter) def fourier_filter(x_in, threshold, scale): diff --git a/modules/sd_hijack.py b/modules/sd_hijack.py index 2f707d098..15e6dba7c 100644 --- a/modules/sd_hijack.py +++ b/modules/sd_hijack.py @@ -1,8 +1,10 @@ from types import MethodType, SimpleNamespace import io import contextlib +from functools import wraps import torch from torch.nn.functional import silu +import diffusers from modules import shared shared.log.debug('Importing LDM') @@ -331,3 +333,41 @@ ldm.models.diffusion.plms.PLMSSampler.register_buffer = register_buffer # Ensure samping from Guassian for DDPM follows types ldm.modules.distributions.distributions.DiagonalGaussianDistribution.sample = lambda self: self.mean.to(self.parameters.dtype) + self.std.to(self.parameters.dtype) * torch.randn(self.mean.shape, dtype=self.parameters.dtype).to(device=self.parameters.device) + + +# Upcast BF16 to FP32 +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 + if input.dtype == torch.bfloat16: + input = input.to(dtype=torch.float32) + return original_fft_fftn(input, s=s, dim=dim, norm=norm, out=out).to(dtype=return_dtype) + + +# Upcast BF16 to FP32 +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 + if input.dtype == torch.bfloat16: + input = input.to(dtype=torch.float32) + return original_fft_ifftn(input, s=s, dim=dim, norm=norm, out=out).to(dtype=return_dtype) + + +# Diffusers FreeU +# Diffusers is imported before sd_hijacks so fourier_filter needs hijacking too +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 + if x_in.dtype == torch.bfloat16: + x_in = x_in.to(dtype=torch.float32) + return original_fourier_filter(x_in, threshold, scale).to(dtype=return_dtype) + + +# IPEX always upcasts +if devices.backend != "ipex": + torch.fft.fftn = fft_fftn + torch.fft.ifftn = fft_ifftn + diffusers.utils.torch_utils.fourier_filter = fourier_filter