Unstable & need more test.
This commit is contained in:
Seunghoon Lee
2023-04-26 12:45:44 +09:00
parent 8b75033a11
commit df0e89be48
6 changed files with 8 additions and 10 deletions
+5 -7
View File
@@ -1,11 +1,9 @@
import torch
import torch_directml
import modules.dml.kdiffusion
import modules.dml.stablediffusion
import modules.dml.torch
import modules.dml.hijack
from optimizer.unknown import UnknownOptimizer
from modules.dml.optimizer.unknown import UnknownOptimizer
class DirectML():
def get_optimizer(self, device: torch.device):
@@ -13,11 +11,11 @@ class DirectML():
try:
device_name = torch_directml.device_name(device.index)
if 'NVIDIA' in device_name or 'GeForce' in device_name:
from optimizer.nvidia import nVidiaOptimizer as optimizer
from modules.dml.optimizer.nvidia import nVidiaOptimizer as optimizer
elif 'AMD' in device_name or 'Radeon' in device_name:
from optimizer.amd import AMDOptimizer as optimizer
from modules.dml.optimizer.amd import AMDOptimizer as optimizer
elif 'Intel' in device_name:
from optimizer.intel import IntelOptimizer as optimizer
from modules.dml.optimizer.intel import IntelOptimizer as optimizer
else:
return UnknownOptimizer
return optimizer
+3
View File
@@ -0,0 +1,3 @@
import modules.dml.hijack.kdiffusion
import modules.dml.hijack.stablediffusion
import modules.dml.hijack.torch
-3
View File
@@ -20,9 +20,6 @@ if shared.opts.cross_attention_optimization == "xFormers":
except Exception:
pass
if shared.device.type == 'privateuseone':
import dml
def get_available_vram():
if shared.device.type == 'cuda':