diff --git a/modules/dml/Generator.py b/modules/dml/Generator.py new file mode 100644 index 000000000..c35f87361 --- /dev/null +++ b/modules/dml/Generator.py @@ -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") diff --git a/modules/dml/__init__.py b/modules/dml/__init__.py index 015ba7d63..3b9c8cf63 100644 --- a/modules/dml/__init__.py +++ b/modules/dml/__init__.py @@ -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 @@ -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) diff --git a/modules/dml/backend.py b/modules/dml/backend.py index 1c7ce5f28..887991e86 100644 --- a/modules/dml/backend.py +++ b/modules/dml/backend.py @@ -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 diff --git a/modules/dml/hijack/torch.py b/modules/dml/hijack/torch.py index edada0c9a..084671425 100644 --- a/modules/dml/hijack/torch.py +++ b/modules/dml/hijack/torch.py @@ -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