Add dml_specific.

This commit is contained in:
Seunghoon Lee
2023-04-25 23:19:03 +09:00
parent 32634298d7
commit eb072db23c
2 changed files with 3 additions and 8 deletions
@@ -2,6 +2,7 @@ import torch
from tqdm.auto import tqdm
from modules.shared import device
from modules.sd_hijack_utils import CondFunc
# k-diffusion
from k_diffusion import sampling
@@ -172,10 +173,4 @@ DDIMSampler.p_sample_ddim = p_sample_ddim
# torch
Generator_init = torch.Generator.__init__
def Generator_init_fix(self, device = None, *args, **kwargs):
if device is not None and device.type == 'privateuseone':
return Generator_init(self, 'cpu', *args, **kwargs) # DML Solution: torch.Generator fallback to cpu.
else:
return Generator_init(self, device, *args, **kwargs)
torch.Generator.__init__ = Generator_init_fix
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')
+1 -1
View File
@@ -64,7 +64,7 @@ clip_model = None
if device.type == 'privateuseone':
import modules.sd_hijack_directml
import modules.dml_specific
is_device_dml = True