Torch 2.0 for IPEX

This commit is contained in:
Disty0
2023-08-03 16:29:56 +03:00
parent 246989129f
commit a293e3cdcb
2 changed files with 6 additions and 1 deletions
+1 -1
View File
@@ -323,7 +323,7 @@ def check_torch():
os.environ.setdefault('NEOReadDebugKeys', '1')
os.environ.setdefault('ClDeviceGlobalMemSizeAvailablePercent', '100')
os.environ.setdefault('TENSORFLOW_PACKAGE', 'tensorflow==2.12.0 intel-extension-for-tensorflow[gpu]')
torch_command = os.environ.get('TORCH_COMMAND', 'torch==1.13.0a0+git6c9b55e torchvision==0.14.1a0 intel_extension_for_pytorch==1.13.120+xpu -f https://developer.intel.com/ipex-whl-stable-xpu')
torch_command = os.environ.get('TORCH_COMMAND', 'torch==2.0.1a0 torchvision==0.15.2a0 intel_extension_for_pytorch==2.0.110+xpu -f https://developer.intel.com/ipex-whl-stable-xpu')
else:
machine = platform.machine()
if sys.platform == 'darwin':
+5
View File
@@ -1,4 +1,5 @@
import os
import contextlib
import torch
import intel_extension_for_pytorch as ipex
from modules import shared
@@ -16,6 +17,9 @@ def ipex_no_cuda(orig_func, *args, **kwargs): # pylint: disable=redefined-outer-
orig_func(*args, **kwargs)
torch.cuda.is_available = torch.xpu.is_available
def return_null_context(*args, **kwargs):
return contextlib.nullcontext()
def ipex_init():
#Fix functions with ipex
torch.cuda.is_available = torch.xpu.is_available
@@ -27,6 +31,7 @@ def ipex_init():
torch._utils._get_available_device_type = lambda: "xpu" # pylint: disable=protected-access
torch.cuda.set_device = torch.xpu.set_device
torch.cuda.synchronize = torch.xpu.synchronize
torch.backends.cuda.sdp_kernel = return_null_context
torch.Tensor.cuda = torch.Tensor.xpu
torch.nn.DataParallel = DummyDataParallel