update dml

This commit is contained in:
Seunghoon Lee
2023-10-08 02:19:56 +09:00
parent c6226605dc
commit e5f8b7f0a4
4 changed files with 49 additions and 1 deletions
+6
View File
@@ -0,0 +1,6 @@
import torch
from typing import Optional
class Generator(torch.Generator):
def __init__(self, device: Optional[torch.device] = None):
super().__init__("cpu")
+6
View File
@@ -10,6 +10,7 @@ if platform.system() == "Windows":
memory_providers.append("Performance Counter")
default_memory_provider = "Performance Counter"
do_nothing = lambda: None # pylint: disable=unnecessary-lambda-assignment
do_nothing_with_self = lambda self: None # pylint: disable=unnecessary-lambda-assignment
def _set_memory_provider():
from modules.shared import opts, cmd_opts, log
@@ -66,7 +67,12 @@ def directml_do_hijack():
import modules.dml.hijack # pylint: disable=unused-import
from modules.devices import device
CondFunc('torch.Generator',
lambda orig_func, device: orig_func("cpu"),
lambda orig_func, device: True)
if not torch.dml.has_float64_support(device):
torch.Tensor.__str__ = do_nothing_with_self
CondFunc('torch.from_numpy',
lambda orig_func, *args, **kwargs: orig_func(args[0].astype('float32')),
lambda *args, **kwargs: args[1].dtype == float)
+2
View File
@@ -6,6 +6,7 @@ import modules.dml.amp as amp
from .utils import rDevice, get_device
from .device import device
from .Generator import Generator
from .device_properties import DeviceProperties
def amd_mem_get_info(device: Optional[rDevice]=None) -> tuple[int, int]:
@@ -22,6 +23,7 @@ def mem_get_info(device: Optional[rDevice]=None) -> tuple[int, int]:
class DirectML:
amp = amp
device = device
Generator = Generator
context_device: Optional[torch.device] = None
+35 -1
View File
@@ -5,4 +5,38 @@ from modules.sd_hijack_utils import CondFunc
CondFunc('torchsde._brownian.brownian_interval._randn', lambda _, size, dtype, device, seed: torch.randn(size, dtype=dtype, device=torch.device("cpu"), generator=torch.Generator(torch.device("cpu")).manual_seed(int(seed))).to(device), lambda _, size, dtype, device, seed: device.type == 'privateuseone')
# https://github.com/microsoft/DirectML/issues/400
CondFunc('torch.Tensor.new', lambda orig, self, *args, **kwargs: orig(self.cpu(), *args, **kwargs), lambda orig, self, *args, **kwargs: torch.dml.is_directml_device(self.device))
CondFunc('torch.Tensor.new', lambda orig, self, *args, **kwargs: orig(self.cpu(), *args, **kwargs).to(self.device), lambda orig, self, *args, **kwargs: torch.dml.is_directml_device(self.device))
_lerp = torch.lerp
def lerp(*args, **kwargs) -> torch.Tensor:
rep = None
for i in range(0, len(args)):
if torch.is_tensor(args[i]):
rep = args[i]
break
if rep is None:
for key in kwargs:
if torch.is_tensor(kwargs[key]):
rep = kwargs[key]
break
if torch.dml.is_directml_device(rep.device):
args = list(args)
if rep.dtype == torch.float16:
for i in range(len(args)):
if torch.is_tensor(args[i]):
args[i] = args[i].float()
for i in range(len(args)):
if torch.is_tensor(args[i]):
args[i] = args[i].cpu()
if rep.dtype == torch.float16:
for kwarg in kwargs:
if torch.is_tensor(kwargs[kwarg]):
kwargs[kwarg] = kwargs[kwarg].float()
for kwarg in kwargs:
if torch.is_tensor(kwargs[kwarg]):
kwargs[kwarg] = kwargs[kwarg].cpu()
return _lerp(*args, **kwargs).to(rep.device).type(rep.dtype)
return _lerp(*args, **kwargs)
torch.lerp = lerp