Merge pull request #3377 from Extraltodeus/cuda-device-id-selection
Implementation of CUDA device id selection (--device-id 0/1/2)
This commit is contained in:
commit
e80bdcab91
2 changed files with 19 additions and 3 deletions
|
@ -1,7 +1,6 @@
|
||||||
|
import sys, os, shlex
|
||||||
import contextlib
|
import contextlib
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from modules import errors
|
from modules import errors
|
||||||
|
|
||||||
# has_mps is only available in nightly pytorch (for now), `getattr` for compatibility
|
# has_mps is only available in nightly pytorch (for now), `getattr` for compatibility
|
||||||
|
@ -9,10 +8,26 @@ has_mps = getattr(torch, 'has_mps', False)
|
||||||
|
|
||||||
cpu = torch.device("cpu")
|
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():
|
def get_optimal_device():
|
||||||
if torch.cuda.is_available():
|
if torch.cuda.is_available():
|
||||||
return torch.device("cuda")
|
# CUDA device selection support:
|
||||||
|
if "shared" not in sys.modules:
|
||||||
|
commandline_args = os.environ.get('COMMANDLINE_ARGS', "") #re-parse the commandline arguments because using the shared.py module creates an import loop.
|
||||||
|
sys.argv += shlex.split(commandline_args)
|
||||||
|
device_id = extract_device_id(sys.argv, '--device-id')
|
||||||
|
else:
|
||||||
|
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:
|
if has_mps:
|
||||||
return torch.device("mps")
|
return torch.device("mps")
|
||||||
|
|
|
@ -79,6 +79,7 @@ parser.add_argument('--vae-path', type=str, help='Path to Variational Autoencode
|
||||||
parser.add_argument("--disable-safe-unpickle", action='store_true', help="disable checking pytorch models for malicious code", default=False)
|
parser.add_argument("--disable-safe-unpickle", action='store_true', help="disable checking pytorch models for malicious code", default=False)
|
||||||
parser.add_argument("--api", action='store_true', help="use api=True to launch the api with the webui")
|
parser.add_argument("--api", action='store_true', help="use api=True to launch the api with the webui")
|
||||||
parser.add_argument("--nowebui", action='store_true', help="use api=True to launch the api instead of the webui")
|
parser.add_argument("--nowebui", action='store_true', help="use api=True to launch the api instead of the webui")
|
||||||
|
parser.add_argument("--device-id", type=str, help="Select the default CUDA device to use (export CUDA_VISIBLE_DEVICES=0,1,etc might be needed before)", default=None)
|
||||||
parser.add_argument("--browse-all-images", action='store_true', help="Allow browsing all images by Image Browser", default=False)
|
parser.add_argument("--browse-all-images", action='store_true', help="Allow browsing all images by Image Browser", default=False)
|
||||||
|
|
||||||
cmd_opts = parser.parse_args()
|
cmd_opts = parser.parse_args()
|
||||||
|
|
Loading…
Reference in a new issue