Files
automatic/modules/dml/__init__.py
T
2023-07-08 13:35:25 -04:00

40 lines
1.3 KiB
Python

# pylint: disable=no-member,no-self-argument,no-method-argument
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
class DirectML():
is_autocast_enabled = False
autocast_gpu_dtype = torch.float16
def get_optimizer(device: torch.device):
assert device.type == 'privateuseone'
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
elif 'AMD' in device_name or 'Radeon' in device_name:
from .optimizer.amd import AMDOptimizer as optimizer
elif 'Intel' in device_name:
from .optimizer.intel import IntelOptimizer as optimizer
else:
return UnknownOptimizer
return optimizer
except Exception:
return UnknownOptimizer
def memory_stats(device: torch.device):
optimizer = DirectML.get_optimizer(device)
return optimizer.memory_stats(device.index)
DirectML.amp = amp
# Alternative of torch.cuda for DirectML.
torch.dml = DirectML
override_opts()