From 9cadf4fc10e4e3918f283b2581cae95fc0d3478d Mon Sep 17 00:00:00 2001 From: Disty0 Date: Fri, 8 Sep 2023 03:51:44 +0300 Subject: [PATCH] IPEX fix and make SDP default --- modules/intel/ipex/__init__.py | 43 ++++----- modules/intel/ipex/attention.py | 131 ++++++++++++++++++++++++++ modules/intel/ipex/diffusers.py | 154 ++----------------------------- modules/intel/ipex/gradscaler.py | 18 ++-- modules/intel/ipex/hijacks.py | 126 ++++++++++++++++++------- modules/shared.py | 2 +- 6 files changed, 257 insertions(+), 217 deletions(-) create mode 100644 modules/intel/ipex/attention.py diff --git a/modules/intel/ipex/__init__.py b/modules/intel/ipex/__init__.py index 9794f98b7..bd2e8e142 100644 --- a/modules/intel/ipex/__init__.py +++ b/modules/intel/ipex/__init__.py @@ -2,22 +2,14 @@ import os import sys import contextlib import torch -import intel_extension_for_pytorch as ipex -from modules import shared -from .diffusers import ipex_diffusers +import intel_extension_for_pytorch as ipex # pylint: disable=import-error, unused-import from .hijacks import ipex_hijacks +from .attention import attention_init +from .diffusers import ipex_diffusers -#ControlNet depth_leres++ -class DummyDataParallel(torch.nn.Module): - def __new__(cls, module, device_ids=None, output_device=None, dim=0): - if type(device_ids) is list and len(device_ids) > 1: - shared.log.warning("IPEX backend doesn't support DataParallel on multiple XPU devices") - return module.to(shared.device) +# pylint: disable=protected-access, missing-function-docstring, line-too-long -def return_null_context(*args, **kwargs): - return contextlib.nullcontext() - -def ipex_init(): +def ipex_init(): # pylint: disable=too-many-statements try: #Replace cuda with xpu: torch.cuda.current_device = torch.xpu.current_device @@ -140,10 +132,13 @@ def ipex_init(): torch.cuda.amp.common.amp_definitely_not_available = lambda: False try: torch.cuda.amp.GradScaler = torch.xpu.amp.GradScaler - except Exception: - from .gradscaler import gradscaler_init - gradscaler_init() - torch.cuda.amp.GradScaler = torch.xpu.amp.GradScaler + except Exception: # pylint: disable=broad-exception-caught + try: + from .gradscaler import gradscaler_init # pylint: disable=import-outside-toplevel, import-error + gradscaler_init() + torch.cuda.amp.GradScaler = torch.xpu.amp.GradScaler + except Exception: # pylint: disable=broad-exception-caught + torch.cuda.amp.GradScaler = ipex.cpu.autocast._grad_scaler.GradScaler #C torch._C._cuda_getCurrentRawStream = ipex._C._getCurrentStream @@ -152,20 +147,20 @@ def ipex_init(): #Fix functions with ipex: torch.cuda.mem_get_info = lambda device=None: [(torch.xpu.get_device_properties(device).total_memory - torch.xpu.memory_allocated(device)), torch.xpu.get_device_properties(device).total_memory] - torch._utils._get_available_device_type = lambda: "xpu" # pylint: disable=protected-access + torch._utils._get_available_device_type = lambda: "xpu" torch.has_cuda = True torch.cuda.has_half = True - torch.cuda.is_bf16_supported = True + torch.cuda.is_bf16_supported = lambda *args, **kwargs: True + torch.cuda.is_fp16_supported = lambda *args, **kwargs: True #torch.version.cuda = "11.7" #Breaks System Info - torch.cuda.get_device_capability = lambda: [11,7] + torch.cuda.get_device_capability = lambda *args, **kwargs: [11,7] torch.cuda.get_device_properties.major = 11 torch.cuda.get_device_properties.minor = 7 - torch.backends.cuda.sdp_kernel = return_null_context - torch.nn.DataParallel = DummyDataParallel - torch.cuda.ipc_collect = lambda: None - torch.cuda.utilization = lambda: 0 + torch.cuda.ipc_collect = lambda *args, **kwargs: None + torch.cuda.utilization = lambda *args, **kwargs: 0 ipex_hijacks() + attention_init() ipex_diffusers() except Exception as e: return False, e diff --git a/modules/intel/ipex/attention.py b/modules/intel/ipex/attention.py new file mode 100644 index 000000000..87b14e978 --- /dev/null +++ b/modules/intel/ipex/attention.py @@ -0,0 +1,131 @@ +import torch +import intel_extension_for_pytorch as ipex # pylint: disable=import-error, unused-import +import diffusers #0.20.2 # pylint: disable=import-error + +# pylint: disable=protected-access, missing-function-docstring, line-too-long + +Attention = diffusers.models.attention_processor.Attention + +original_torch_bmm = torch.bmm +def torch_bmm(input, mat2, *, out=None): + if input.dtype != mat2.dtype: + mat2 = mat2.to(input.dtype) + + #ARC GPUs can't allocate more than 4GB to a single block, Slice it: + batch_size_attention, input_tokens, mat2_shape = input.shape[0], input.shape[1], mat2.shape[2] + block_multiply = 2.4 if input.dtype == torch.float32 else 1.2 + block_size = (batch_size_attention * input_tokens * mat2_shape) / 1024 * block_multiply #MB + split_slice_size = batch_size_attention + if block_size >= 4000: + do_split = True + #Find something divisible with the input_tokens + while ((split_slice_size * input_tokens * mat2_shape) / 1024 * block_multiply) > 4000: + split_slice_size = split_slice_size // 2 + if split_slice_size <= 1: + split_slice_size = 1 + break + else: + do_split = False + + split_block_size = (split_slice_size * input_tokens * mat2_shape) / 1024 * block_multiply #MB + split_2_slice_size = input_tokens + if split_block_size >= 4000: + do_split_2 = True + #Find something divisible with the input_tokens + while ((split_slice_size * split_2_slice_size * mat2_shape) / 1024 * block_multiply) > 4000: + split_2_slice_size = split_2_slice_size // 2 + if split_2_slice_size <= 1: + split_2_slice_size = 1 + break + else: + do_split_2 = False + + if do_split: + hidden_states = torch.zeros(input.shape[0], input.shape[1], mat2.shape[2], device=input.device, dtype=input.dtype) + for i in range(batch_size_attention // split_slice_size): + start_idx = i * split_slice_size + end_idx = (i + 1) * split_slice_size + if do_split_2: + for i2 in range(input_tokens // split_2_slice_size): # pylint: disable=invalid-name + start_idx_2 = i2 * split_2_slice_size + end_idx_2 = (i2 + 1) * split_2_slice_size + hidden_states[start_idx:end_idx, start_idx_2:end_idx_2] = original_torch_bmm( + input[start_idx:end_idx, start_idx_2:end_idx_2], + mat2[start_idx:end_idx, start_idx_2:end_idx_2], + out=out + ) + else: + hidden_states[start_idx:end_idx] = original_torch_bmm( + input[start_idx:end_idx], + mat2[start_idx:end_idx], + out=out + ) + else: + return original_torch_bmm(input, mat2, out=out) + return hidden_states + +original_scaled_dot_product_attention = torch.nn.functional.scaled_dot_product_attention +def scaled_dot_product_attention(query, key, value, attn_mask=None, dropout_p=0.0, is_causal=False): + #ARC GPUs can't allocate more than 4GB to a single block, Slice it: + shape_one, batch_size_attention, query_tokens, shape_four = query.shape + block_multiply = 2.4 if query.dtype == torch.float32 else 1.2 + block_size = (shape_one * batch_size_attention * query_tokens * shape_four) / 1024 * block_multiply #MB + split_slice_size = batch_size_attention + if block_size >= 4000: + do_split = True + #Find something divisible with the shape_one + while ((shape_one * split_slice_size * query_tokens * shape_four) / 1024 * block_multiply) > 4000: + split_slice_size = split_slice_size // 2 + if split_slice_size <= 1: + split_slice_size = 1 + break + else: + do_split = False + + split_block_size = (shape_one * split_slice_size * query_tokens * shape_four) / 1024 * block_multiply #MB + split_2_slice_size = query_tokens + if split_block_size >= 4000: + do_split_2 = True + #Find something divisible with the batch_size_attention + while ((shape_one * split_slice_size * split_2_slice_size * shape_four) / 1024 * block_multiply) > 4000: + split_2_slice_size = split_2_slice_size // 2 + if split_2_slice_size <= 1: + split_2_slice_size = 1 + break + else: + do_split_2 = False + + if do_split: + hidden_states = torch.zeros(query.shape, device=query.device, dtype=query.dtype) + for i in range(batch_size_attention // split_slice_size): + start_idx = i * split_slice_size + end_idx = (i + 1) * split_slice_size + if do_split_2: + for i2 in range(query_tokens // split_2_slice_size): # pylint: disable=invalid-name + start_idx_2 = i2 * split_2_slice_size + end_idx_2 = (i2 + 1) * split_2_slice_size + hidden_states[:, start_idx:end_idx, start_idx_2:end_idx_2] = original_scaled_dot_product_attention( + query[:, start_idx:end_idx, start_idx_2:end_idx_2], + key[:, start_idx:end_idx, start_idx_2:end_idx_2], + value[:, start_idx:end_idx, start_idx_2:end_idx_2], + attn_mask=attn_mask[:, start_idx:end_idx, start_idx_2:end_idx_2] if attn_mask is not None else attn_mask, + dropout_p=dropout_p, is_causal=is_causal + ) + else: + hidden_states[:, start_idx:end_idx] = original_scaled_dot_product_attention( + query[:, start_idx:end_idx], + key[:, start_idx:end_idx], + 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 + ) + else: + return original_scaled_dot_product_attention( + query, key, value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal + ) + return hidden_states + +def attention_init(): + #ARC GPUs can't allocate more than 4GB to a single block: + torch.bmm = torch_bmm + torch.nn.functional.scaled_dot_product_attention = scaled_dot_product_attention diff --git a/modules/intel/ipex/diffusers.py b/modules/intel/ipex/diffusers.py index e03592197..3435abe14 100644 --- a/modules/intel/ipex/diffusers.py +++ b/modules/intel/ipex/diffusers.py @@ -1,11 +1,10 @@ import torch -import intel_extension_for_pytorch as ipex -import torch.nn.functional as F -import diffusers #0.20.2 +import intel_extension_for_pytorch as ipex # pylint: disable=import-error, unused-import +import diffusers #0.20.2 # pylint: disable=import-error -Attention = diffusers.models.attention_processor.Attention +# pylint: disable=protected-access, missing-function-docstring, line-too-long -class SlicedAttnProcessor: +class SlicedAttnProcessor: # pylint: disable=too-few-public-methods r""" Processor for implementing sliced attention. @@ -18,7 +17,7 @@ class SlicedAttnProcessor: def __init__(self, slice_size): self.slice_size = slice_size - def __call__(self, attn: Attention, hidden_states, encoder_hidden_states=None, attention_mask=None): + def __call__(self, attn: diffusers.models.attention_processor.Attention, hidden_states, encoder_hidden_states=None, attention_mask=None): # pylint: disable=too-many-statements, too-many-locals, too-many-branches residual = hidden_states input_ndim = hidden_states.ndim @@ -74,7 +73,7 @@ class SlicedAttnProcessor: end_idx = (i + 1) * self.slice_size if do_split_2: - for i2 in range(query_tokens // split_2_slice_size): + for i2 in range(query_tokens // split_2_slice_size): # pylint: disable=invalid-name start_idx_2 = i2 * split_2_slice_size end_idx_2 = (i2 + 1) * split_2_slice_size @@ -114,147 +113,6 @@ class SlicedAttnProcessor: return hidden_states -class AttnProcessor2_0: - r""" - Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). - """ - - def __init__(self): - if not hasattr(F, "scaled_dot_product_attention"): - raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.") - - def __call__( - self, - attn: Attention, - hidden_states, - encoder_hidden_states=None, - attention_mask=None, - temb=None, - ): - residual = hidden_states - - if attn.spatial_norm is not None: - hidden_states = attn.spatial_norm(hidden_states, temb) - - input_ndim = hidden_states.ndim - - if input_ndim == 4: - batch_size, channel, height, width = hidden_states.shape - hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) - - batch_size, sequence_length, _ = ( - hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape - ) - - if attention_mask is not None: - attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) - # scaled_dot_product_attention expects attention_mask shape to be - # (batch, heads, source_length, target_length) - attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) - - if attn.group_norm is not None: - hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) - - query = attn.to_q(hidden_states) - - if encoder_hidden_states is None: - encoder_hidden_states = hidden_states - elif attn.norm_cross: - encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) - - key = attn.to_k(encoder_hidden_states) - value = attn.to_v(encoder_hidden_states) - - inner_dim = key.shape[-1] - head_dim = inner_dim // attn.heads - - query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) - - #ARC GPUs can't allocate more than 4GB to a single block, Slice it: - shape_one, batch_size_attention, query_tokens, shape_four = query.shape - block_multiply = 2.4 if query.dtype == torch.float32 else 1.2 - block_size = (shape_one * batch_size_attention * query_tokens * shape_four) / 1024 * block_multiply #MB - split_slice_size = batch_size_attention - if block_size >= 4000: - do_split = True - #Find something divisible with the shape_one - while ((shape_one * split_slice_size * query_tokens * shape_four) / 1024 * block_multiply) > 4000: - split_slice_size = split_slice_size // 2 - if split_slice_size <= 1: - split_slice_size = 1 - break - else: - do_split = False - - split_block_size = (shape_one * split_slice_size * query_tokens * shape_four) / 1024 * block_multiply #MB - split_2_slice_size = query_tokens - if split_block_size >= 4000: - do_split_2 = True - #Find something divisible with the batch_size_attention - while ((shape_one * split_slice_size * split_2_slice_size * shape_four) / 1024 * block_multiply) > 4000: - split_2_slice_size = split_2_slice_size // 2 - if split_2_slice_size <= 1: - split_2_slice_size = 1 - break - else: - do_split_2 = False - - if do_split: - hidden_states = torch.zeros(query.shape, device=query.device, dtype=query.dtype) - for i in range(batch_size_attention // split_slice_size): - start_idx = i * split_slice_size - end_idx = (i + 1) * split_slice_size - if do_split_2: - for i2 in range(query_tokens // split_2_slice_size): - start_idx_2 = i2 * split_2_slice_size - end_idx_2 = (i2 + 1) * split_2_slice_size - - query_slice = query[:, start_idx:end_idx, start_idx_2:end_idx_2] - key_slice = key[:, start_idx:end_idx, start_idx_2:end_idx_2] - attn_mask_slice = attention_mask[:, start_idx:end_idx, start_idx_2:end_idx_2] if attention_mask is not None else None - - attn_slice = F.scaled_dot_product_attention( - query_slice, key_slice, value[:, start_idx:end_idx, start_idx_2:end_idx_2], - attn_mask=attn_mask_slice, dropout_p=0.0, is_causal=False - ) - hidden_states[:, start_idx:end_idx, start_idx_2:end_idx_2] = attn_slice - else: - query_slice = query[:, start_idx:end_idx] - key_slice = key[:, start_idx:end_idx] - attn_mask_slice = attention_mask[:, start_idx:end_idx] if attention_mask is not None else None - - attn_slice = F.scaled_dot_product_attention( - query_slice, key_slice, value[:, start_idx:end_idx], - attn_mask=attn_mask_slice, dropout_p=0.0, is_causal=False - ) - hidden_states[:, start_idx:end_idx] = attn_slice - else: - hidden_states = F.scaled_dot_product_attention( - query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False - ) - - hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) - hidden_states = hidden_states.to(query.dtype) - - # linear proj - hidden_states = attn.to_out[0](hidden_states) - # dropout - hidden_states = attn.to_out[1](hidden_states) - - if input_ndim == 4: - hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) - - if attn.residual_connection: - hidden_states = hidden_states + residual - - hidden_states = hidden_states / attn.rescale_output_factor - - return hidden_states - def ipex_diffusers(): #ARC GPUs can't allocate more than 4GB to a single block: diffusers.models.attention_processor.SlicedAttnProcessor = SlicedAttnProcessor - diffusers.models.attention_processor.AttnProcessor2_0 = AttnProcessor2_0 diff --git a/modules/intel/ipex/gradscaler.py b/modules/intel/ipex/gradscaler.py index 217f6cc7d..530212101 100644 --- a/modules/intel/ipex/gradscaler.py +++ b/modules/intel/ipex/gradscaler.py @@ -1,14 +1,15 @@ -import torch from collections import defaultdict -import intel_extension_for_pytorch as ipex -import intel_extension_for_pytorch._C as core -from modules import shared +import torch +import intel_extension_for_pytorch as ipex # pylint: disable=import-error, unused-import +import intel_extension_for_pytorch._C as core # pylint: disable=import-error, unused-import + +# pylint: disable=protected-access, missing-function-docstring, line-too-long OptState = ipex.cpu.autocast._grad_scaler.OptState _MultiDeviceReplicator = ipex.cpu.autocast._grad_scaler._MultiDeviceReplicator _refresh_per_optimizer_state = ipex.cpu.autocast._grad_scaler._refresh_per_optimizer_state -def _unscale_grads_(self, optimizer, inv_scale, found_inf, allow_fp16): +def _unscale_grads_(self, optimizer, inv_scale, found_inf, allow_fp16): # pylint: disable=unused-argument per_device_inv_scale = _MultiDeviceReplicator(inv_scale) per_device_found_inf = _MultiDeviceReplicator(found_inf) @@ -40,7 +41,7 @@ def _unscale_grads_(self, optimizer, inv_scale, found_inf, allow_fp16): else: to_unscale = param.grad - # TODO: is there a way to split by device and dtype without appending in the inner loop? + # -: is there a way to split by device and dtype without appending in the inner loop? to_unscale = to_unscale.to("cpu") per_device_and_dtype_grads[to_unscale.device][ to_unscale.dtype @@ -86,7 +87,7 @@ def unscale_(self, optimizer): optimizer_state = self._per_optimizer_states[id(optimizer)] - if optimizer_state["stage"] is OptState.UNSCALED: + if optimizer_state["stage"] is OptState.UNSCALED: # pylint: disable=no-else-raise raise RuntimeError( "unscale_() has already been called on this optimizer since the last update()." ) @@ -175,5 +176,4 @@ def gradscaler_init(): torch.xpu.amp.GradScaler._unscale_grads_ = _unscale_grads_ torch.xpu.amp.GradScaler.unscale_ = unscale_ torch.xpu.amp.GradScaler.update = update - - + return torch.xpu.amp.GradScaler diff --git a/modules/intel/ipex/hijacks.py b/modules/intel/ipex/hijacks.py index 9e1d9713e..1ad90d72b 100644 --- a/modules/intel/ipex/hijacks.py +++ b/modules/intel/ipex/hijacks.py @@ -1,19 +1,63 @@ +import contextlib import torch -import intel_extension_for_pytorch as ipex -from modules import devices +import intel_extension_for_pytorch as ipex # pylint: disable=import-error, unused-import from modules.sd_hijack_utils import CondFunc +from modules import devices + +# pylint: disable=protected-access, missing-function-docstring, line-too-long, unnecessary-lambda, no-else-return + +def _shutdown_workers(self): + if torch.utils.data._utils is None or torch.utils.data._utils.python_exit_status is True or torch.utils.data._utils.python_exit_status is None: + return + if hasattr(self, "_shutdown") and not self._shutdown: + self._shutdown = True + try: + if hasattr(self, '_pin_memory_thread'): + self._pin_memory_thread_done_event.set() + self._worker_result_queue.put((None, None)) + self._pin_memory_thread.join() + self._worker_result_queue.cancel_join_thread() + self._worker_result_queue.close() + self._workers_done_event.set() + for worker_id in range(len(self._workers)): + if self._persistent_workers or self._workers_status[worker_id]: + self._mark_worker_as_unavailable(worker_id, shutdown=True) + for w in self._workers: # pylint: disable=invalid-name + w.join(timeout=torch.utils.data._utils.MP_STATUS_CHECK_INTERVAL) + for q in self._index_queues: # pylint: disable=invalid-name + q.cancel_join_thread() + q.close() + finally: + if self._worker_pids_set: + torch.utils.data._utils.signal_handling._remove_worker_pids(id(self)) + self._worker_pids_set = False + for w in self._workers: # pylint: disable=invalid-name + if w.is_alive(): + w.terminate() + +class DummyDataParallel(torch.nn.Module): # pylint: disable=missing-class-docstring, unused-argument, too-few-public-methods + def __new__(cls, module, device_ids=None, output_device=None, dim=0): # pylint: disable=unused-argument + if isinstance(device_ids, list) and len(device_ids) > 1: + print("IPEX backend doesn't support DataParallel on multiple XPU devices") + return module.to(devices.device) + +def return_null_context(*args, **kwargs): # pylint: disable=unused-argument + return contextlib.nullcontext() def check_device(device): return bool((isinstance(device, torch.device) and device.type == "cuda") or (isinstance(device, str) and "cuda" in device) or isinstance(device, int)) -def ipex_no_cuda(orig_func, *args, **kwargs): # pylint: disable=redefined-outer-name +def return_xpu(device): + return f"xpu:{device[-1]}" if isinstance(device, str) and ":" in device else f"xpu:{device}" if isinstance(device, int) else torch.device(devices.device) if isinstance(device, torch.device) else devices.device + +def ipex_no_cuda(orig_func, *args, **kwargs): torch.cuda.is_available = lambda: False orig_func(*args, **kwargs) torch.cuda.is_available = torch.xpu.is_available original_autocast = torch.autocast def ipex_autocast(*args, **kwargs): - if args[0] == "cuda" or args[0] == "xpu": + if len(args) > 0 and args[0] == "cuda" or args[0] == "xpu": if "dtype" in kwargs: return original_autocast("xpu", *args[1:], **kwargs) else: @@ -23,66 +67,75 @@ def ipex_autocast(*args, **kwargs): #Embedding BF16 original_torch_cat = torch.cat -def torch_cat(input, *args, **kwargs): - if len(input) == 3 and (input[0].dtype != input[1].dtype or input[2].dtype != input[1].dtype): - return original_torch_cat([input[0].to(input[1].dtype), input[1], input[2].to(input[1].dtype)], *args, **kwargs) +def torch_cat(tensor, *args, **kwargs): + if len(tensor) == 3 and (tensor[0].dtype != tensor[1].dtype or tensor[2].dtype != tensor[1].dtype): + return original_torch_cat([tensor[0].to(tensor[1].dtype), tensor[1], tensor[2].to(tensor[1].dtype)], *args, **kwargs) else: - return original_torch_cat(input, *args, **kwargs) + return original_torch_cat(tensor, *args, **kwargs) #Latent antialias: original_interpolate = torch.nn.functional.interpolate -def interpolate(input, size=None, scale_factor=None, mode='nearest', align_corners=None, recompute_scale_factor=None, antialias=False): - if antialias: - return original_interpolate(input.to("cpu", dtype=torch.float32), size=size, scale_factor=scale_factor, mode=mode, - align_corners=align_corners, recompute_scale_factor=recompute_scale_factor, antialias=antialias).to(devices.device, dtype=devices.dtype) +def interpolate(tensor, size=None, scale_factor=None, mode='nearest', align_corners=None, recompute_scale_factor=None, antialias=False): # pylint: disable=too-many-arguments + if antialias or align_corners is not None: + return_device = tensor.device + return_dtype = tensor.dtype + return original_interpolate(tensor.to("cpu", dtype=torch.float32), size=size, scale_factor=scale_factor, mode=mode, + align_corners=align_corners, recompute_scale_factor=recompute_scale_factor, antialias=antialias).to(return_device, dtype=return_dtype) else: - return original_interpolate(input, size=size, scale_factor=scale_factor, mode=mode, + return original_interpolate(tensor, size=size, scale_factor=scale_factor, mode=mode, align_corners=align_corners, recompute_scale_factor=recompute_scale_factor, antialias=antialias) +original_linalg_solve = torch.linalg.solve +def linalg_solve(A, B, *args, **kwargs): # pylint: disable=invalid-name + if A.device != torch.device("cpu") or B.device != torch.device("cpu"): + return_device = A.device + return original_linalg_solve(A.to("cpu"), B.to("cpu"), *args, **kwargs).to(return_device) + else: + return original_linalg_solve(A, B, *args, **kwargs) + def ipex_hijacks(): CondFunc('torch.Tensor.to', - lambda orig_func, self, device=None, *args, **kwargs: orig_func(self, devices.device, *args, **kwargs), + lambda orig_func, self, device=None, *args, **kwargs: orig_func(self, return_xpu(device), *args, **kwargs), lambda orig_func, self, device=None, *args, **kwargs: check_device(device)) CondFunc('torch.Tensor.cuda', - lambda orig_func, self, device=None, *args, **kwargs: orig_func(self, devices.device, *args, **kwargs), + lambda orig_func, self, device=None, *args, **kwargs: orig_func(self, return_xpu(device), *args, **kwargs), lambda orig_func, self, device=None, *args, **kwargs: check_device(device)) CondFunc('torch.empty', - lambda orig_func, *args, device=None, **kwargs: orig_func(*args, device=devices.device, **kwargs), + lambda orig_func, *args, device=None, **kwargs: orig_func(*args, device=return_xpu(device), **kwargs), lambda orig_func, *args, device=None, **kwargs: check_device(device)) CondFunc('torch.load', - lambda orig_func, *args, map_location=None, **kwargs: orig_func(*args, devices.device, **kwargs), + lambda orig_func, *args, map_location=None, **kwargs: orig_func(*args, return_xpu(map_location), **kwargs), lambda orig_func, *args, map_location=None, **kwargs: map_location is None or check_device(map_location)) CondFunc('torch.randn', - lambda orig_func, *args, device=None, **kwargs: orig_func(*args, device=devices.device, **kwargs), + lambda orig_func, *args, device=None, **kwargs: orig_func(*args, device=return_xpu(device), **kwargs), lambda orig_func, *args, device=None, **kwargs: check_device(device)) CondFunc('torch.ones', - lambda orig_func, *args, device=None, **kwargs: orig_func(*args, device=devices.device, **kwargs), + lambda orig_func, *args, device=None, **kwargs: orig_func(*args, device=return_xpu(device), **kwargs), lambda orig_func, *args, device=None, **kwargs: check_device(device)) CondFunc('torch.zeros', - lambda orig_func, *args, device=None, **kwargs: orig_func(*args, device=devices.device, **kwargs), + lambda orig_func, *args, device=None, **kwargs: orig_func(*args, device=return_xpu(device), **kwargs), lambda orig_func, *args, device=None, **kwargs: check_device(device)) CondFunc('torch.tensor', - lambda orig_func, *args, device=None, **kwargs: orig_func(*args, device=devices.device, **kwargs), + lambda orig_func, *args, device=None, **kwargs: orig_func(*args, device=return_xpu(device), **kwargs), + lambda orig_func, *args, device=None, **kwargs: check_device(device)) + CondFunc('torch.linspace', + lambda orig_func, *args, device=None, **kwargs: orig_func(*args, device=return_xpu(device), **kwargs), lambda orig_func, *args, device=None, **kwargs: check_device(device)) CondFunc('torch.Generator', - lambda orig_func, device: torch.xpu.Generator(device), - lambda orig_func, device: device != torch.device("cpu") and device != "cpu") - #Crashes the GPU: - CondFunc('torch.linalg.solve', - lambda orig_func, A, B, *args, **kwargs: orig_func(A.to("cpu"), B.to("cpu"), *args, **kwargs).to(devices.device), - lambda orig_func, A, B, *args, **kwargs: A.device != torch.device("cpu") or B.device != torch.device("cpu")) + lambda orig_func, device=None: torch.xpu.Generator(device), + lambda orig_func, device=None: device is not None and device != torch.device("cpu") and device != "cpu") #TiledVAE and ControlNet: CondFunc('torch.batch_norm', lambda orig_func, input, weight, bias, *args, **kwargs: orig_func(input, - weight if weight is not None else torch.ones(input.size()[1], device=devices.device), - bias if bias is not None else torch.zeros(input.size()[1], device=devices.device), *args, **kwargs), + weight if weight is not None else torch.ones(input.size()[1], device=input.device), + bias if bias is not None else torch.zeros(input.size()[1], device=input.device), *args, **kwargs), lambda orig_func, input, *args, **kwargs: input.device != torch.device("cpu")) CondFunc('torch.instance_norm', lambda orig_func, input, weight, bias, *args, **kwargs: orig_func(input, - weight if weight is not None else torch.ones(input.size()[1], device=devices.device), - bias if bias is not None else torch.zeros(input.size()[1], device=devices.device), *args, **kwargs), + weight if weight is not None else torch.ones(input.size()[1], device=input.device), + bias if bias is not None else torch.zeros(input.size()[1], device=input.device), *args, **kwargs), lambda orig_func, input, *args, **kwargs: input.device != torch.device("cpu")) #Functions with dtype errors: @@ -94,10 +147,9 @@ def ipex_hijacks(): CondFunc('torch.nn.modules.linear.Linear.forward', lambda orig_func, self, input: orig_func(self, input.to(self.weight.data.dtype)), lambda orig_func, self, input: input.dtype != self.weight.data.dtype) - #Embedding FP32: - CondFunc('torch.bmm', - lambda orig_func, input, mat2, *args, **kwargs: orig_func(input, mat2.to(input.dtype), *args, **kwargs), - lambda orig_func, input, mat2, *args, **kwargs: input.dtype != mat2.dtype) + CondFunc('torch.nn.modules.conv.Conv2d.forward', + lambda orig_func, self, input: orig_func(self, input.to(self.weight.data.dtype)), + lambda orig_func, self, input: input.dtype != self.weight.data.dtype) #BF16: CondFunc('torch.nn.functional.layer_norm', lambda orig_func, input, normalized_shape=None, weight=None, *args, **kwargs: @@ -118,6 +170,10 @@ def ipex_hijacks(): lambda orig_func, *args, **kwargs: True) #Functions that make compile mad with CondFunc: + torch.utils.data.dataloader._MultiProcessingDataLoaderIter._shutdown_workers = _shutdown_workers + torch.nn.DataParallel = DummyDataParallel torch.autocast = ipex_autocast torch.cat = torch_cat + torch.linalg.solve = linalg_solve torch.nn.functional.interpolate = interpolate + torch.backends.cuda.sdp_kernel = return_null_context diff --git a/modules/shared.py b/modules/shared.py index 3a8fb74b9..c9e3c49ea 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -350,7 +350,7 @@ if devices.backend == "cpu": elif devices.backend == "mps": cross_attention_optimization_default = "Doggettx's" elif devices.backend == "ipex": - cross_attention_optimization_default = "Sub-quadratic" + cross_attention_optimization_default = "Scaled-Dot-Product" elif devices.backend == "directml": cross_attention_optimization_default = "Sub-quadratic" elif devices.backend == "rocm":