diff --git a/modules/intel/ipex/__init__.py b/modules/intel/ipex/__init__.py index c78547915..333504935 100644 --- a/modules/intel/ipex/__init__.py +++ b/modules/intel/ipex/__init__.py @@ -140,6 +140,7 @@ def ipex_init(): # pylint: disable=too-many-statements # C torch._C._cuda_getCurrentRawStream = ipex._C._getCurrentStream + ipex._C._DeviceProperties.multi_processor_count = ipex._C._DeviceProperties.gpu_eu_count ipex._C._DeviceProperties.major = 2023 ipex._C._DeviceProperties.minor = 2 diff --git a/modules/intel/ipex/attention.py b/modules/intel/ipex/attention.py index 5bffe159b..e9f927a9c 100644 --- a/modules/intel/ipex/attention.py +++ b/modules/intel/ipex/attention.py @@ -171,7 +171,5 @@ def scaled_dot_product_attention_32_bit(query, key, value, attn_mask=None, dropo 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 original_scaled_dot_product_attention(query, key, value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal) return hidden_states diff --git a/modules/intel/ipex/diffusers.py b/modules/intel/ipex/diffusers.py index 3275379b1..617c12369 100644 --- a/modules/intel/ipex/diffusers.py +++ b/modules/intel/ipex/diffusers.py @@ -67,7 +67,9 @@ class SlicedAttnProcessor: # pylint: disable=too-few-public-methods def __init__(self, slice_size): self.slice_size = slice_size - def __call__(self, attn: Attention, hidden_states, encoder_hidden_states=None, attention_mask=None): # pylint: disable=too-many-statements, too-many-locals, too-many-branches + def __call__(self, attn: Attention, hidden_states: torch.FloatTensor, + encoder_hidden_states=None, attention_mask=None) -> torch.FloatTensor: # pylint: disable=too-many-statements, too-many-locals, too-many-branches + residual = hidden_states input_ndim = hidden_states.ndim @@ -182,15 +184,10 @@ class AttnProcessor: Default processor for performing attention-related computations. """ - def __call__( - self, - attn: Attention, - hidden_states: torch.FloatTensor, - encoder_hidden_states=None, - attention_mask=None, - temb=None, - scale: float = 1.0, - ) -> torch.Tensor: + def __call__(self, attn: Attention, hidden_states: torch.FloatTensor, + encoder_hidden_states=None, attention_mask=None, + temb=None, scale: float = 1.0) -> torch.Tensor: # pylint: disable=too-many-statements, too-many-locals, too-many-branches + residual = hidden_states args = () if USE_PEFT_BACKEND else (scale,) diff --git a/modules/intel/ipex/hijacks.py b/modules/intel/ipex/hijacks.py index 6cdcded62..0fc5b15f7 100644 --- a/modules/intel/ipex/hijacks.py +++ b/modules/intel/ipex/hijacks.py @@ -6,35 +6,6 @@ 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: @@ -44,22 +15,18 @@ class DummyDataParallel(torch.nn.Module): # pylint: disable=missing-class-docstr def return_null_context(*args, **kwargs): # pylint: disable=unused-argument return contextlib.nullcontext() +@property +def is_cuda(self): + return self.device.type == 'xpu' + 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 return_xpu(device): return f"xpu:{device.split(':')[-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 - -@property -def is_cuda(self): - return self.device.type == 'xpu' - +# Autocast original_autocast = torch.autocast def ipex_autocast(*args, **kwargs): if len(args) > 0 and args[0] == "cuda" or args[0] == "xpu": @@ -70,7 +37,6 @@ def ipex_autocast(*args, **kwargs): else: return original_autocast(*args, **kwargs) - # Latent Antialias CPU Offload: original_interpolate = torch.nn.functional.interpolate def interpolate(tensor, size=None, scale_factor=None, mode='nearest', align_corners=None, recompute_scale_factor=None, antialias=False): # pylint: disable=too-many-arguments @@ -83,15 +49,13 @@ def interpolate(tensor, size=None, scale_factor=None, mode='nearest', align_corn return original_interpolate(tensor, size=size, scale_factor=scale_factor, mode=mode, align_corners=align_corners, recompute_scale_factor=recompute_scale_factor, antialias=antialias) -# Torch Linalg Solve CPU Offload for IPEX 2.0 and older: -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) +# Diffusers Float64 (Alchemist GPUs doesn't support 64 bit): +original_from_numpy = torch.from_numpy +def from_numpy(ndarray): + if ndarray.dtype == float: + return original_from_numpy(ndarray.astype('float32')) else: - return original_linalg_solve(A, B, *args, **kwargs) - + return original_from_numpy(ndarray) if torch.xpu.has_fp64_dtype(): original_torch_bmm = torch.bmm @@ -162,6 +126,14 @@ def torch_cat(tensor, *args, **kwargs): else: return original_torch_cat(tensor, *args, **kwargs) +# SwinIR BF16: +original_funtional_pad = torch.nn.functional.pad +def funtional_pad(input, pad, mode='constant', value=None): + if mode == 'reflect' and input.dtype == torch.bfloat16: + return original_funtional_pad(input.to(torch.float32), pad, mode=mode, value=value).to(dtype=torch.bfloat16) + else: + return original_funtional_pad(input, pad, mode=mode, value=value) + def ipex_hijacks(): CondFunc('torch.tensor', @@ -194,65 +166,29 @@ def ipex_hijacks(): 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=None: orig_func(return_xpu(device)), + lambda orig_func, device=None: check_device(device)) CondFunc('torch.load', lambda orig_func, f, map_location=None, pickle_module=None, *, weights_only=False, mmap=None, **kwargs: orig_func(f, map_location=return_xpu(map_location), pickle_module=pickle_module, weights_only=weights_only, mmap=mmap, **kwargs), lambda orig_func, f, map_location=None, pickle_module=None, *, weights_only=False, mmap=None, **kwargs: check_device(map_location)) - if hasattr(torch.xpu, "Generator"): - CondFunc('torch.Generator', - lambda orig_func, device=None: torch.xpu.Generator(return_xpu(device)), - lambda orig_func, device=None: device is not None and device != torch.device("cpu") and device != "cpu") - else: - CondFunc('torch.Generator', - lambda orig_func, device=None: orig_func(return_xpu(device)), - lambda orig_func, device=None: check_device(device)) - - # A1111 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=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=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")) - - # SwinIR BF16: - CondFunc('torch.nn.functional.pad', - lambda orig_func, input, pad, mode='constant', value=None: orig_func(input.to(torch.float32), pad, mode=mode, value=value).to(dtype=torch.bfloat16), - lambda orig_func, input, pad, mode='constant', value=None: mode == 'reflect' and input.dtype == torch.bfloat16) - - # Diffusers Float64 (Alchemist GPUs doesn't support 64 bit): - if not torch.xpu.has_fp64_dtype(): - CondFunc('torch.from_numpy', - lambda orig_func, ndarray: orig_func(ndarray.astype('float32')), - lambda orig_func, ndarray: ndarray.dtype == float) - - # Broken functions when torch.cuda.is_available is True: - # Pin Memory: - CondFunc('torch.utils.data.dataloader._BaseDataLoaderIter.__init__', - lambda orig_func, *args, **kwargs: ipex_no_cuda(orig_func, *args, **kwargs), - lambda orig_func, *args, **kwargs: True) # Hijack Functions: - torch.nn.DataParallel = DummyDataParallel - torch.utils.data.dataloader._MultiProcessingDataLoaderIter._shutdown_workers = _shutdown_workers - torch.UntypedStorage.is_cuda = is_cuda - - torch.autocast = ipex_autocast torch.backends.cuda.sdp_kernel = return_null_context + torch.nn.DataParallel = DummyDataParallel + torch.UntypedStorage.is_cuda = is_cuda + torch.autocast = ipex_autocast + torch.nn.functional.scaled_dot_product_attention = scaled_dot_product_attention torch.nn.functional.group_norm = functional_group_norm torch.nn.functional.layer_norm = functional_layer_norm torch.nn.functional.linear = functional_linear torch.nn.functional.conv2d = functional_conv2d - torch.nn.functional.interpolate = interpolate - if hasattr(torch.xpu, "Generator"): - torch.linalg.solve = linalg_solve + torch.nn.functional.pad = funtional_pad torch.bmm = torch_bmm torch.cat = torch_cat - torch.nn.functional.scaled_dot_product_attention = scaled_dot_product_attention + if not torch.xpu.has_fp64_dtype(): + torch.from_numpy = from_numpy