mirror of
https://github.com/vladmandic/automatic
synced 2026-09-20 01:31:13 +02:00
IPEX remove the use of lambdas
This commit is contained in:
@@ -16,6 +16,28 @@ torch_version[0], torch_version[1] = int(torch_version[0]), int(torch_version[1]
|
||||
|
||||
# pylint: disable=protected-access, missing-function-docstring, line-too-long
|
||||
|
||||
def return_true(*args, **kwargs):
|
||||
return True
|
||||
|
||||
def return_false(*args, **kwargs):
|
||||
return False
|
||||
|
||||
def return_none(*args, **kwargs):
|
||||
return None
|
||||
|
||||
def return_zero(*args, **kwargs):
|
||||
return 0
|
||||
|
||||
def return_cuda_version(*args, **kwargs):
|
||||
return (12,1)
|
||||
|
||||
def return_xpu_string(*args, **kwargs):
|
||||
return "xpu"
|
||||
|
||||
def return_arch_list(*args, **kwargs):
|
||||
return ["pvc", "dg2", "ats-m150"]
|
||||
|
||||
|
||||
def ipex_init(): # pylint: disable=too-many-statements
|
||||
try:
|
||||
if hasattr(torch, "cuda") and hasattr(torch.cuda, "is_xpu_hijacked") and torch.cuda.is_xpu_hijacked:
|
||||
@@ -26,9 +48,9 @@ def ipex_init(): # pylint: disable=too-many-statements
|
||||
# import inductor utils to get around lazy import
|
||||
from torch._inductor import utils as torch_inductor_utils # pylint: disable=import-error, unused-import # noqa: F401,RUF100
|
||||
torch._inductor.utils.GPU_TYPES = ["xpu"]
|
||||
torch._inductor.utils.get_gpu_type = lambda *args, **kwargs: "xpu"
|
||||
torch._inductor.utils.get_gpu_type = return_xpu_string
|
||||
from triton import backends as triton_backends # pylint: disable=import-error
|
||||
triton_backends.backends["nvidia"].driver.is_active = lambda *args, **kwargs: False
|
||||
triton_backends.backends["nvidia"].driver.is_active = return_false
|
||||
except Exception:
|
||||
pass
|
||||
# Replace cuda with xpu:
|
||||
@@ -56,7 +78,7 @@ def ipex_init(): # pylint: disable=too-many-statements
|
||||
torch.cuda._get_device_index = torch.xpu._get_device_index
|
||||
torch.cuda._lazy_init = torch.xpu._lazy_init
|
||||
torch.cuda._lazy_call = torch.xpu._lazy_call
|
||||
torch.cuda.is_current_stream_capturing = lambda: False
|
||||
torch.cuda.is_current_stream_capturing = return_false
|
||||
|
||||
torch.cuda.__annotations__ = torch.xpu.__annotations__
|
||||
torch.cuda.__builtins__ = torch.xpu.__builtins__
|
||||
@@ -145,7 +167,7 @@ def ipex_init(): # pylint: disable=too-many-statements
|
||||
|
||||
# Memory:
|
||||
if "linux" in sys.platform and "WSL2" in os.popen("uname -a").read():
|
||||
torch.xpu.empty_cache = lambda: None
|
||||
torch.xpu.empty_cache = return_none
|
||||
torch.cuda.empty_cache = torch.xpu.empty_cache
|
||||
|
||||
if torch_version[0] >= 2 and torch_version[1] >= 8:
|
||||
@@ -180,21 +202,24 @@ def ipex_init(): # pylint: disable=too-many-statements
|
||||
torch.cuda.initial_seed = torch.xpu.initial_seed
|
||||
|
||||
# Fix functions with ipex:
|
||||
# torch.xpu.mem_get_info always returns the total memory as free memory
|
||||
torch.has_cuda = True
|
||||
torch.version.cuda = "12.1"
|
||||
torch.backends.cuda.is_built = lambda *args, **kwargs: True
|
||||
torch._utils._get_available_device_type = lambda: "xpu"
|
||||
torch.backends.cuda.is_built = return_true
|
||||
torch._utils._get_available_device_type = return_xpu_string
|
||||
|
||||
torch.xpu.mem_get_info = lambda device=None: [(torch.xpu.get_device_properties(device).total_memory - torch.xpu.memory_reserved(device)), torch.xpu.get_device_properties(device).total_memory]
|
||||
# torch.xpu.mem_get_info always returns the total memory as free memory
|
||||
def mem_get_info(device=None):
|
||||
return [(torch.xpu.get_device_properties(device).total_memory - torch.xpu.memory_reserved(device)), torch.xpu.get_device_properties(device).total_memory]
|
||||
torch.xpu.mem_get_info = mem_get_info
|
||||
torch.cuda.mem_get_info = torch.xpu.mem_get_info
|
||||
|
||||
torch.cuda.has_half = True
|
||||
torch.cuda.is_bf16_supported = getattr(torch.xpu, "is_bf16_supported", lambda *args, **kwargs: True)
|
||||
torch.cuda.is_fp16_supported = lambda *args, **kwargs: True
|
||||
torch.cuda.get_arch_list = getattr(torch.xpu, "get_arch_list", lambda: ["pvc", "dg2", "ats-m150"])
|
||||
torch.cuda.get_device_capability = lambda *args, **kwargs: (12,1)
|
||||
torch.cuda.ipc_collect = lambda *args, **kwargs: None
|
||||
torch.cuda.utilization = lambda *args, **kwargs: 0
|
||||
torch.cuda.is_bf16_supported = getattr(torch.xpu, "is_bf16_supported", return_true)
|
||||
torch.cuda.is_fp16_supported = getattr(torch.xpu, "is_fp16_supported", return_true)
|
||||
torch.cuda.get_arch_list = getattr(torch.xpu, "get_arch_list", return_arch_list)
|
||||
torch.cuda.get_device_capability = return_cuda_version
|
||||
torch.cuda.ipc_collect = return_none
|
||||
torch.cuda.utilization = return_zero
|
||||
|
||||
device_supports_fp64 = ipex_hijacks()
|
||||
try:
|
||||
|
||||
@@ -15,8 +15,10 @@ torch_version[0], torch_version[1] = int(torch_version[0]), int(torch_version[1]
|
||||
|
||||
device_supports_fp64 = torch.xpu.has_fp64_dtype() if hasattr(torch.xpu, "has_fp64_dtype") else torch.xpu.get_device_properties(devices.device).has_fp64
|
||||
|
||||
# pylint: disable=protected-access, missing-function-docstring, line-too-long, unnecessary-lambda, no-else-return
|
||||
# pylint: disable=protected-access, missing-function-docstring, line-too-long, no-else-return
|
||||
|
||||
def return_false(*args, **kwargs):
|
||||
return False
|
||||
|
||||
@property
|
||||
def is_cuda(self):
|
||||
@@ -437,6 +439,6 @@ def ipex_hijacks():
|
||||
|
||||
if not hasattr(torch.cuda.amp, "common"):
|
||||
torch.cuda.amp.common = nullcontext()
|
||||
torch.cuda.amp.common.amp_definitely_not_available = lambda: False
|
||||
torch.cuda.amp.common.amp_definitely_not_available = return_false
|
||||
|
||||
return device_supports_fp64
|
||||
|
||||
Reference in New Issue
Block a user