Improve MPS handling and add VAE slicing/tiling
This commit is contained in:
@@ -1907,7 +1907,8 @@ class App(TkinterDnD.Tk): # type: ignore
|
|||||||
raise RuntimeError(f"Model file not found: {model_path}")
|
raise RuntimeError(f"Model file not found: {model_path}")
|
||||||
has_mps = bool(getattr(torch.backends, "mps", None) and torch.backends.mps.is_available())
|
has_mps = bool(getattr(torch.backends, "mps", None) and torch.backends.mps.is_available())
|
||||||
device = "cuda" if torch.cuda.is_available() else "mps" if has_mps else "cpu"
|
device = "cuda" if torch.cuda.is_available() else "mps" if has_mps else "cpu"
|
||||||
dtype = torch.float16 if device in {"cuda", "mps"} else torch.float32
|
# MPS often produces black outputs in float16; keep float32 there for stability
|
||||||
|
dtype = torch.float16 if device == "cuda" else torch.float32
|
||||||
pipe = StableDiffusionXLPipeline.from_single_file(
|
pipe = StableDiffusionXLPipeline.from_single_file(
|
||||||
str(model_path),
|
str(model_path),
|
||||||
torch_dtype=dtype,
|
torch_dtype=dtype,
|
||||||
@@ -1924,6 +1925,11 @@ class App(TkinterDnD.Tk): # type: ignore
|
|||||||
pipe.enable_attention_slicing()
|
pipe.enable_attention_slicing()
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
try:
|
||||||
|
pipe.enable_vae_slicing()
|
||||||
|
pipe.enable_vae_tiling()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
try:
|
try:
|
||||||
pipe.set_progress_bar_config(disable=True)
|
pipe.set_progress_bar_config(disable=True)
|
||||||
except Exception:
|
except Exception:
|
||||||
|
|||||||
Reference in New Issue
Block a user