diff options
author | yfszzx <yfszzx@gmail.com> | 2022-10-23 08:17:37 +0000 |
---|---|---|
committer | yfszzx <yfszzx@gmail.com> | 2022-10-23 08:17:37 +0000 |
commit | 6a9ea40d7f64139f23d634efd7c2fb2c743a546f (patch) | |
tree | 8d420a5590c09e1d4c005a709166479856e70580 /modules/devices.py | |
parent | 67b78f0ea6f196bfdca49932da062631bb40d0b1 (diff) | |
parent | 1ef32c8b8fa3e16a1e7b287eb19d4fc943d1f2a5 (diff) | |
download | stable-diffusion-webui-gfx803-6a9ea40d7f64139f23d634efd7c2fb2c743a546f.tar.gz stable-diffusion-webui-gfx803-6a9ea40d7f64139f23d634efd7c2fb2c743a546f.tar.bz2 stable-diffusion-webui-gfx803-6a9ea40d7f64139f23d634efd7c2fb2c743a546f.zip |
Move browser and Inspiration into extension
Diffstat (limited to 'modules/devices.py')
-rw-r--r-- | modules/devices.py | 19 |
1 files changed, 15 insertions, 4 deletions
diff --git a/modules/devices.py b/modules/devices.py index eb422583..dc1f3cdd 100644 --- a/modules/devices.py +++ b/modules/devices.py @@ -1,7 +1,6 @@ +import sys, os, shlex import contextlib - import torch - from modules import errors # has_mps is only available in nightly pytorch (for now), `getattr` for compatibility @@ -9,10 +8,22 @@ has_mps = getattr(torch, 'has_mps', False) cpu = torch.device("cpu") +def extract_device_id(args, name): + for x in range(len(args)): + if name in args[x]: return args[x+1] + return None def get_optimal_device(): if torch.cuda.is_available(): - return torch.device("cuda") + from modules import shared + + device_id = shared.cmd_opts.device_id + + if device_id is not None: + cuda_device = f"cuda:{device_id}" + return torch.device(cuda_device) + else: + return torch.device("cuda") if has_mps: return torch.device("mps") @@ -34,7 +45,7 @@ def enable_tf32(): errors.run(enable_tf32, "Enabling TF32") -device = device_interrogate = device_gfpgan = device_bsrgan = device_esrgan = device_scunet = device_codeformer = get_optimal_device() +device = device_interrogate = device_gfpgan = device_bsrgan = device_esrgan = device_scunet = device_codeformer = None dtype = torch.float16 dtype_vae = torch.float16 |