From bdc4bd846f48c14148937bd22af8983a1faeb787 Mon Sep 17 00:00:00 2001 From: Disty0 Date: Sat, 18 Nov 2023 15:42:53 +0300 Subject: [PATCH] IPEX fix torch.load --- modules/intel/ipex/hijacks.py | 7 ++++--- modules/upscaler.py | 1 - 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/modules/intel/ipex/hijacks.py b/modules/intel/ipex/hijacks.py index 34342f84c..432e805de 100644 --- a/modules/intel/ipex/hijacks.py +++ b/modules/intel/ipex/hijacks.py @@ -103,9 +103,6 @@ def ipex_hijacks(): CondFunc('torch.empty', 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, 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=return_xpu(device), **kwargs), lambda orig_func, *args, device=None, **kwargs: check_device(device)) @@ -121,6 +118,10 @@ 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.load', + lambda orig_func, f, map_location=None, pickle_module=None, *, weights_only=False, mmap=None, **kwargs: + orig_func(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)) CondFunc('torch.Generator', lambda orig_func, device=None: torch.xpu.Generator(device), diff --git a/modules/upscaler.py b/modules/upscaler.py index d15f44ad3..fa23ddaa9 100644 --- a/modules/upscaler.py +++ b/modules/upscaler.py @@ -228,7 +228,6 @@ def compile_upscaler(model, name=""): if modules.shared.opts.cuda_compile_backend == "openvino_fx": from modules.intel.openvino import openvino_fx # pylint: disable=unused-import - from modules.sd_models_compile import CompiledModelState # pylint: disable=unused-import torch._dynamo.eval_frame.check_if_dynamo_supported = lambda: True # pylint: disable=protected-access log_level = logging.WARNING if modules.shared.opts.cuda_compile_verbose else logging.CRITICAL # pylint: disable=protected-access