Add an opts override for DirectML.

This commit is contained in:
Seunghoon Lee
2023-07-06 01:35:00 +09:00
parent d30a55e523
commit 993de932ab
3 changed files with 10 additions and 2 deletions
+4 -1
View File
@@ -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()
+5
View File
@@ -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.