diff --git a/modules/devices.py b/modules/devices.py index 7511e1dc325..058a5e00114 100644 --- a/modules/devices.py +++ b/modules/devices.py @@ -39,10 +39,17 @@ def torch_gc(): def enable_tf32(): if torch.cuda.is_available(): + for devid in range(0,torch.cuda.device_count()): + if torch.cuda.get_device_capability(devid) == (7, 5): + shd = True + if shd: + torch.backends.cudnn.benchmark = True + torch.backends.cudnn.enabled = True torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cudnn.allow_tf32 = True + errors.run(enable_tf32, "Enabling TF32") device = device_interrogate = device_gfpgan = device_swinir = device_esrgan = device_scunet = device_codeformer = None