aboutsummaryrefslogtreecommitdiffstats
path: root/modules/devices.py
diff options
context:
space:
mode:
authorAUTOMATIC <16777216c@gmail.com>2022-09-16 06:48:46 +0000
committerAUTOMATIC <16777216c@gmail.com>2022-09-16 06:48:46 +0000
commit83bce1a604f40356583d31c8c0b2f8b590dda071 (patch)
treec094f7241b9e001367fce750f6344dfa8a66ebcb /modules/devices.py
parentb44ddcb44398fbe922fd7515f66d8b0c2344bc54 (diff)
parent87e8b9a2ab3f033e7fdadbb2fe258857915980ac (diff)
downloadstable-diffusion-webui-gfx803-83bce1a604f40356583d31c8c0b2f8b590dda071.tar.gz
stable-diffusion-webui-gfx803-83bce1a604f40356583d31c8c0b2f8b590dda071.tar.bz2
stable-diffusion-webui-gfx803-83bce1a604f40356583d31c8c0b2f8b590dda071.zip
Merge branch 'batch-seed-attempt'
Diffstat (limited to 'modules/devices.py')
-rw-r--r--modules/devices.py10
1 files changed, 10 insertions, 0 deletions
diff --git a/modules/devices.py b/modules/devices.py
index e4430e1a..07bb2339 100644
--- a/modules/devices.py
+++ b/modules/devices.py
@@ -48,3 +48,13 @@ def randn(seed, shape):
torch.manual_seed(seed)
return torch.randn(shape, device=device)
+
+def randn_without_seed(shape):
+ # Pytorch currently doesn't handle setting randomness correctly when the metal backend is used.
+ if device.type == 'mps':
+ generator = torch.Generator(device=cpu)
+ noise = torch.randn(shape, generator=generator, device=cpu).to(device)
+ return noise
+
+ return torch.randn(shape, device=device)
+