mirror of
https://github.com/invoke-ai/InvokeAI
synced 2024-08-30 20:32:17 +00:00
629ca09fda
* Fix conflicts with main branch changes * Fix logic error in choose_autocast_device() that was causing crashes on CUDA systems.
18 lines
591 B
Python
18 lines
591 B
Python
import torch
|
|
|
|
def choose_torch_device() -> str:
|
|
'''Convenience routine for guessing which GPU device to run model on'''
|
|
if torch.cuda.is_available():
|
|
return 'cuda'
|
|
if hasattr(torch.backends, 'mps') and torch.backends.mps.is_available():
|
|
return 'mps'
|
|
return 'cpu'
|
|
|
|
def choose_autocast_device(device) -> str:
|
|
'''Returns an autocast compatible device from a torch device'''
|
|
device_type = device.type # this returns 'mps' on M1
|
|
# autocast only supports cuda or cpu
|
|
if device_type not in ('cuda','cpu'):
|
|
return 'cpu'
|
|
return device_type
|