diff --git a/modules/devices.py b/modules/devices.py index 94f264532..533e4d4ad 100644 --- a/modules/devices.py +++ b/modules/devices.py @@ -40,7 +40,7 @@ def get_cuda_device_string(): def get_optimal_device_name(): - if cuda_ok or backend == 'ipex' or backend == 'directml': + if cuda_ok or backend == 'directml': return get_cuda_device_string() if has_mps(): return "mps" @@ -76,7 +76,7 @@ def torch_gc(force=False): if shared.opts.disable_gc and not force: return collected = gc.collect() - if cuda_ok or backend == 'ipex': + if cuda_ok: try: with torch.cuda.device(get_cuda_device_string()): torch.cuda.empty_cache() @@ -182,7 +182,7 @@ elif sys.platform == 'darwin': else: backend = 'cpu' -cuda_ok = torch.cuda.is_available() and not backend == 'ipex' +cuda_ok = torch.cuda.is_available() cpu = torch.device("cpu") device = device_interrogate = device_gfpgan = device_esrgan = device_codeformer = None dtype = torch.float16 @@ -221,8 +221,6 @@ def autocast(disable=False): return contextlib.nullcontext() if shared.cmd_opts.use_directml: return torch.dml.amp.autocast(dtype) - if backend == 'ipex': - return torch.xpu.amp.autocast(enabled=True, dtype=dtype) if cuda_ok: return torch.autocast("cuda") else: @@ -234,8 +232,6 @@ def without_autocast(disable=False): return contextlib.nullcontext() if shared.cmd_opts.use_directml: return torch.dml.amp.autocast(enabled=False) if torch.is_autocast_enabled() else contextlib.nullcontext() # pylint: disable=unexpected-keyword-arg - if backend == 'ipex': - return torch.xpu.amp.autocast(enabled=False) if torch.is_autocast_enabled() else contextlib.nullcontext() if cuda_ok: return torch.autocast("cuda", enabled=False) if torch.is_autocast_enabled() else contextlib.nullcontext() else: diff --git a/modules/ipex_specific/__init__.py b/modules/ipex_specific/__init__.py index 7de07d78b..b90328668 100644 --- a/modules/ipex_specific/__init__.py +++ b/modules/ipex_specific/__init__.py @@ -38,8 +38,71 @@ def ipex_init(): torch.cuda.FloatTensor = torch.xpu.FloatTensor torch.Tensor.cuda = torch.Tensor.xpu torch.Tensor.is_cuda = torch.Tensor.is_xpu + torch.cuda._initialization_lock = torch.xpu.lazy_init._initialization_lock + torch.cuda._initialized = torch.xpu.lazy_init._initialized + torch.cuda._lazy_seed_tracker = torch.xpu.lazy_init._lazy_seed_tracker + torch.cuda._queued_calls = torch.xpu.lazy_init._queued_calls + torch.cuda._tls = torch.xpu.lazy_init._tls + torch.cuda.threading = torch.xpu.lazy_init.threading + torch.cuda.traceback = torch.xpu.lazy_init.traceback + torch.cuda.Optional = torch.xpu.Optional + torch.cuda.__cached__ = torch.xpu.__cached__ + torch.cuda.__loader__ = torch.xpu.__loader__ + torch.cuda.ComplexFloatStorage = torch.xpu.ComplexFloatStorage + torch.cuda.Tuple = torch.xpu.Tuple + torch.cuda.streams = torch.xpu.streams + torch.cuda._lazy_new = torch.xpu._lazy_new + torch.cuda.FloatStorage = torch.xpu.FloatStorage + torch.cuda.Any = torch.xpu.Any + torch.cuda.__doc__ = torch.xpu.__doc__ + torch.cuda.default_generators = torch.xpu.default_generators + torch.cuda.HalfTensor = torch.xpu.HalfTensor + torch.cuda._get_device_index = torch.xpu._get_device_index + torch.cuda.__path__ = torch.xpu.__path__ + torch.cuda.Device = torch.xpu.Device + torch.cuda.IntTensor = torch.xpu.IntTensor + torch.cuda.ByteStorage = torch.xpu.ByteStorage + torch.cuda.set_stream = torch.xpu.set_stream + torch.cuda.BoolStorage = torch.xpu.BoolStorage + torch.cuda.get_device_capability = torch.xpu.get_device_capability + torch.cuda.os = torch.xpu.os + torch.cuda.torch = torch.xpu.torch + torch.cuda.BFloat16Storage = torch.xpu.BFloat16Storage + torch.cuda.Union = torch.xpu.Union + torch.cuda.DoubleTensor = torch.xpu.DoubleTensor + torch.cuda.ShortTensor = torch.xpu.ShortTensor + torch.cuda.LongTensor = torch.xpu.LongTensor + torch.cuda.IntStorage = torch.xpu.IntStorage + torch.cuda.LongStorage = torch.xpu.LongStorage + torch.cuda.__annotations__ = torch.xpu.__annotations__ + torch.cuda.__package__ = torch.xpu.__package__ + torch.cuda.__builtins__ = torch.xpu.__builtins__ + torch.cuda.CharTensor = torch.xpu.CharTensor + torch.cuda.List = torch.xpu.List + torch.cuda._lazy_init = torch.xpu._lazy_init + torch.cuda.BFloat16Tensor = torch.xpu.BFloat16Tensor + torch.cuda.DoubleStorage = torch.xpu.DoubleStorage + torch.cuda.ByteTensor = torch.xpu.ByteTensor + torch.cuda.StreamContext = torch.xpu.StreamContext + torch.cuda.ComplexDoubleStorage = torch.xpu.ComplexDoubleStorage + torch.cuda.ShortStorage = torch.xpu.ShortStorage + torch.cuda._lazy_call = torch.xpu._lazy_call + torch.cuda.HalfStorage = torch.xpu.HalfStorage + torch.cuda.random = torch.xpu.random + torch.cuda._device = torch.xpu._device + torch.cuda.classproperty = torch.xpu.classproperty + torch.cuda.__name__ = torch.xpu.__name__ + torch.cuda._device_t = torch.xpu._device_t + torch.cuda.warnings = torch.xpu.warnings + torch.cuda.__spec__ = torch.xpu.__spec__ + torch.cuda.BoolTensor = torch.xpu.BoolTensor + torch.cuda.CharStorage = torch.xpu.CharStorage + torch.cuda.__file__ = torch.xpu.__file__ + torch.cuda._is_in_bad_fork = torch.xpu.lazy_init._is_in_bad_fork + #torch.cuda.is_current_stream_capturing = torch.xpu.is_current_stream_capturing #Memory: + torch.cuda.memory = torch.xpu.memory if 'linux' in sys.platform and "WSL2" in os.popen("uname -a").read(): torch.xpu.empty_cache = lambda: None torch.cuda.empty_cache = torch.xpu.empty_cache @@ -49,8 +112,12 @@ def ipex_init(): torch.cuda.memory_allocated = torch.xpu.memory_allocated torch.cuda.max_memory_allocated = torch.xpu.max_memory_allocated torch.cuda.memory_reserved = torch.xpu.memory_reserved + torch.cuda.memory_cached = torch.xpu.memory_reserved torch.cuda.max_memory_reserved = torch.xpu.max_memory_reserved + torch.cuda.max_memory_cached = torch.xpu.max_memory_reserved torch.cuda.reset_peak_memory_stats = torch.xpu.reset_peak_memory_stats + torch.cuda.reset_max_memory_cached = torch.xpu.reset_peak_memory_stats + torch.cuda.reset_max_memory_allocated = torch.xpu.reset_peak_memory_stats torch.cuda.memory_stats_as_nested_dict = torch.xpu.memory_stats_as_nested_dict torch.cuda.reset_accumulated_memory_stats = torch.xpu.reset_accumulated_memory_stats @@ -65,7 +132,11 @@ def ipex_init(): torch.cuda.seed_all = torch.xpu.seed_all torch.cuda.initial_seed = torch.xpu.initial_seed - #Training: + #AMP: + torch.cuda.amp = torch.xpu.amp + if not hasattr(torch.cuda.amp, "common"): + torch.cuda.amp.common = contextlib.nullcontext() + torch.cuda.amp.common.amp_definitely_not_available = lambda: False try: torch.cuda.amp.GradScaler = torch.xpu.amp.GradScaler except Exception: @@ -79,8 +150,12 @@ 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.cuda.get_device_properties.major = 2023 - torch.cuda.get_device_properties.minor = 2 + torch.has_cuda = True + torch.cuda.has_half = True + torch.cuda.is_bf16_supported = True + torch.version.cuda = "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 diff --git a/modules/ipex_specific/hijacks.py b/modules/ipex_specific/hijacks.py index 05419f426..400afba77 100644 --- a/modules/ipex_specific/hijacks.py +++ b/modules/ipex_specific/hijacks.py @@ -1,6 +1,6 @@ import torch import intel_extension_for_pytorch as ipex -from modules import shared +from modules import devices from modules.sd_hijack_utils import CondFunc def ipex_no_cuda(orig_func, *args, **kwargs): # pylint: disable=redefined-outer-name @@ -8,7 +8,18 @@ def ipex_no_cuda(orig_func, *args, **kwargs): # pylint: disable=redefined-outer- orig_func(*args, **kwargs) torch.cuda.is_available = torch.xpu.is_available -#FP32: +#Autocast +original_autocast = torch.autocast +def ipex_autocast(*args, **kwargs): + if args[0] == "cuda": + if "dtype" in kwargs: + return original_autocast("xpu", *args[1:], **kwargs) + else: + return original_autocast("xpu", *args[1:], dtype=devices.dtype, **kwargs) + else: + return original_autocast(*args, **kwargs) + +#Diffusers BF16: original_linear_forward = torch.nn.modules.Linear.forward def linear_forward(self, input): if input.dtype != self.weight.data.dtype: @@ -37,7 +48,7 @@ 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"), size=size, scale_factor=scale_factor, mode=mode, - align_corners=align_corners, recompute_scale_factor=recompute_scale_factor, antialias=antialias).to(shared.device) + align_corners=align_corners, recompute_scale_factor=recompute_scale_factor, antialias=antialias).to(devices.device) else: return original_interpolate(input, size=size, scale_factor=scale_factor, mode=mode, align_corners=align_corners, recompute_scale_factor=recompute_scale_factor, antialias=antialias) @@ -46,18 +57,25 @@ def ipex_hijacks(): #Libraries that blindly uses cuda: #Adetailer: CondFunc('torch.Tensor.to', - lambda orig_func, self, device=None, *args, **kwargs: orig_func(self, shared.device, *args, **kwargs), + lambda orig_func, self, device=None, *args, **kwargs: orig_func(self, devices.device, *args, **kwargs), lambda orig_func, self, device=None, *args, **kwargs: (type(device) is torch.device and device.type == "cuda") or (type(device) is str and "cuda" in device)) CondFunc('torch.empty', - lambda orig_func, *args, device=None, **kwargs: orig_func(*args, device=shared.device, **kwargs), + lambda orig_func, *args, device=None, **kwargs: orig_func(*args, device=devices.device, **kwargs), lambda orig_func, *args, device=None, **kwargs: (type(device) is torch.device and device.type == "cuda") or (type(device) is str and "cuda" in device)) #ControlNet depth_leres CondFunc('torch.load', - lambda orig_func, *args, map_location=None, **kwargs: orig_func(*args, shared.device, **kwargs), + lambda orig_func, *args, map_location=None, **kwargs: orig_func(*args, devices.device, **kwargs), lambda orig_func, *args, map_location=None, **kwargs: (map_location is None) or (type(map_location) is torch.device and map_location.type == "cuda") or (type(map_location) is str and "cuda" in map_location)) #Diffusers Model CPU Offload: CondFunc('torch.randn', - lambda orig_func, *args, device=None, **kwargs: orig_func(*args, device=shared.device, **kwargs), + lambda orig_func, *args, device=None, **kwargs: orig_func(*args, device=devices.device, **kwargs), + lambda orig_func, *args, device=None, **kwargs: (type(device) is torch.device and device.type == "cuda") or (type(device) is str and "cuda" in device)) + #Other: + CondFunc('torch.ones', + lambda orig_func, *args, device=None, **kwargs: orig_func(*args, device=devices.device, **kwargs), + lambda orig_func, *args, device=None, **kwargs: (type(device) is torch.device and device.type == "cuda") or (type(device) is str and "cuda" in 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: (type(device) is torch.device and device.type == "cuda") or (type(device) is str and "cuda" in device)) #Broken functions when torch.cuda.is_available is True: @@ -67,6 +85,7 @@ def ipex_hijacks(): lambda orig_func, *args, **kwargs: True) #Functions with dtype errors: + #Original backend: CondFunc('torch.nn.modules.GroupNorm.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) @@ -84,7 +103,7 @@ def ipex_hijacks(): #Functions that does not work with the XPU: #UniPC: CondFunc('torch.linalg.solve', - lambda orig_func, A, B, *args, **kwargs: orig_func(A.to("cpu"), B.to("cpu"), *args, **kwargs).to(shared.device), + 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")) #SDE Samplers: CondFunc('torch.Generator', @@ -98,17 +117,18 @@ def ipex_hijacks(): #ControlNet and TiledVAE: 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=shared.device), - bias if bias is not None else torch.zeros(input.size()[1], device=shared.device), *args, **kwargs), + 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), lambda orig_func, input, *args, **kwargs: input.device != torch.device("cpu")) #ControlNet 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=shared.device), - bias if bias is not None else torch.zeros(input.size()[1], device=shared.device), *args, **kwargs), + 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), lambda orig_func, input, *args, **kwargs: input.device != torch.device("cpu")) #Functions that make compile mad with CondFunc: + torch.autocast = ipex_autocast torch.nn.modules.Linear.forward = linear_forward torch.cat = torch_cat torch.nn.functional.conv2d = conv2d diff --git a/modules/ipex_specific/openvino.py b/modules/ipex_specific/openvino.py index 917d3231f..9674ae5a4 100644 --- a/modules/ipex_specific/openvino.py +++ b/modules/ipex_specific/openvino.py @@ -7,13 +7,6 @@ from torch._dynamo.backends.common import fake_tensor_unsupported from torch._dynamo.backends.registry import register_backend from torch.fx.experimental.proxy_tensor import make_fx -class ModelState: - def __init__(self): - self.recompile = 1 - self.partition_id = 0 - -model_state = ModelState() - @register_backend @fake_tensor_unsupported def openvino_fx(subgraph, example_inputs):