From 993de932ab50f3b3a6e41c76e86a38b49bf761a8 Mon Sep 17 00:00:00 2001 From: Seunghoon Lee Date: Thu, 6 Jul 2023 01:35:00 +0900 Subject: [PATCH] Add an opts override for DirectML. --- extensions-builtin/sd-extension-system-info | 2 +- modules/dml/__init__.py | 5 ++++- modules/dml/opts.py | 5 +++++ 3 files changed, 10 insertions(+), 2 deletions(-) create mode 100644 modules/dml/opts.py diff --git a/extensions-builtin/sd-extension-system-info b/extensions-builtin/sd-extension-system-info index fcdd10c79..b30e32455 160000 --- a/extensions-builtin/sd-extension-system-info +++ b/extensions-builtin/sd-extension-system-info @@ -1 +1 @@ -Subproject commit fcdd10c7957f85504a4511f2040ec7e01f2054a8 +Subproject commit b30e324552517e68012f2487cf7b0be43616a4cc diff --git a/modules/dml/__init__.py b/modules/dml/__init__.py index 3dc89173a..9e721591c 100644 --- a/modules/dml/__init__.py +++ b/modules/dml/__init__.py @@ -3,6 +3,7 @@ import torch import torch_directml # pylint: disable=import-error import modules.dml.hijack import modules.dml.amp as amp +from modules.dml.opts import override_opts from .optimizer.unknown import UnknownOptimizer @@ -43,6 +44,8 @@ class DirectML(): DirectML._is_autocast_enabled = enabled -# Alternative of torch.cuda for DirectML. DirectML.amp = amp +# Alternative of torch.cuda for DirectML. torch.dml = DirectML + +override_opts() diff --git a/modules/dml/opts.py b/modules/dml/opts.py new file mode 100644 index 000000000..32e24a4d5 --- /dev/null +++ b/modules/dml/opts.py @@ -0,0 +1,5 @@ +from modules import shared + +def override_opts(): + if shared.cmd_opts.backend.lower() == "diffusers": + shared.opts.diffusers_generator_device = "cpu" # DirectML does not support torch.Generator API.