From f63dd1c92e2e360969c4120d3aebfad57dfe3978 Mon Sep 17 00:00:00 2001 From: Disty0 Date: Sat, 10 Jun 2023 22:01:09 +0300 Subject: [PATCH] Fix torch.linalg.solve with IPEX & Diffusers UniPC --- modules/devices.py | 7 ++++--- modules/models/diffusion/uni_pc/uni_pc.py | 14 ++------------ 2 files changed, 6 insertions(+), 15 deletions(-) diff --git a/modules/devices.py b/modules/devices.py index 1b7cb428b..ef847862d 100644 --- a/modules/devices.py +++ b/modules/devices.py @@ -34,9 +34,7 @@ def get_cuda_device_string(): def get_optimal_device_name(): - if shared.cmd_opts.use_ipex: - return get_cuda_device_string() - elif cuda_ok and not shared.cmd_opts.use_directml: + if (cuda_ok or shared.cmd_opts.use_ipex) and not shared.cmd_opts.use_directml: return get_cuda_device_string() if has_mps(): return "mps" @@ -172,6 +170,9 @@ if args.use_ipex: CondFunc('torch.nn.modules.Linear.forward', lambda orig_func, *args, **kwargs: orig_func(args[0], args[1].to(args[0].weight.data.dtype)), lambda *args, **kwargs: args[2].dtype != args[1].weight.data.dtype) + CondFunc('torch.linalg.solve', + lambda orig_func, *args, **kwargs: orig_func(args[0].to("cpu"), args[1].to("cpu")).to(get_cuda_device_string()), + lambda *args, **kwargs: True) #Use XPU instead of CPU. %20 Perf improvement on weak CPUs. if args.device_id is not None: diff --git a/modules/models/diffusion/uni_pc/uni_pc.py b/modules/models/diffusion/uni_pc/uni_pc.py index 3e1fec68a..0ff426ae3 100644 --- a/modules/models/diffusion/uni_pc/uni_pc.py +++ b/modules/models/diffusion/uni_pc/uni_pc.py @@ -685,12 +685,7 @@ class UniPC: if order == 2: rhos_p = torch.tensor([0.5], device=b.device) else: - if shared.cmd_opts.use_ipex: - #Running torch.linalg.solve on XPU crashes the GPU. - rhos_p = torch.linalg.solve(R[:-1, :-1].to("cpu"), b[:-1].to("cpu")) - rhos_p = rhos_p.to(b.device) - else: - rhos_p = torch.linalg.solve(R[:-1, :-1], b[:-1]) + rhos_p = torch.linalg.solve(R[:-1, :-1], b[:-1]) else: D1s = None @@ -700,12 +695,7 @@ class UniPC: if order == 1: rhos_c = torch.tensor([0.5], device=b.device) else: - if shared.cmd_opts.use_ipex: - #Running torch.linalg.solve on XPU crashes the GPU. - rhos_c = torch.linalg.solve(R.to("cpu"), b.to("cpu")) - rhos_c = rhos_c.to(b.device) - else: - rhos_c = torch.linalg.solve(R, b) + rhos_c = torch.linalg.solve(R, b) model_t = None if self.predict_x0: