implement CUDA device selection by ID
This commit is contained in:
parent
f49c08ea56
commit
57eb54b838
1 changed files with 18 additions and 3 deletions
|
@ -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,26 @@ 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")
|
||||
# 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:
|
||||
return torch.device("mps")
|
||||
|
|
Loading…
Reference in a new issue