From 14480f926ae2c1c98708ef49c70bdeb285bbcd10 Mon Sep 17 00:00:00 2001 From: Kijai <40791699+kijai@users.noreply.github.com> Date: Tue, 19 Mar 2024 16:25:54 +0200 Subject: [PATCH] Default to fp16 unet with 'auto' --- nodes.py | 8 ++++---- nodes_v2.py | 11 ++++++----- 2 files changed, 10 insertions(+), 9 deletions(-) diff --git a/nodes.py b/nodes.py index 0afd352..2b9a78c 100644 --- a/nodes.py +++ b/nodes.py @@ -186,14 +186,14 @@ class SUPIR_Upscale: if diffusion_dtype == 'auto': try: + if mm.should_use_fp16(): + print("Diffusion using fp16") + dtype = torch.float16 + model_dtype = 'fp16' if mm.should_use_bf16(): print("Diffusion using bf16") dtype = torch.bfloat16 model_dtype = 'bf16' - elif mm.should_use_fp16(): - print("Diffusion using using fp16") - dtype = torch.float16 - model_dtype = 'fp16' else: print("Diffusion using using fp32") dtype = torch.float32 diff --git a/nodes_v2.py b/nodes_v2.py index d579097..cd563ef 100644 --- a/nodes_v2.py +++ b/nodes_v2.py @@ -259,6 +259,7 @@ class SUPIR_first_stage: mm.unload_all_models() if encoder_dtype == 'auto': try: + if mm.should_use_bf16(): print("Encoder using bf16") vae_dtype = 'bf16' @@ -610,14 +611,14 @@ class SUPIR_model_loader: if diffusion_dtype == 'auto': try: - if mm.should_use_bf16(): + if mm.should_use_fp16(): + print("Diffusion using fp16") + dtype = torch.float16 + model_dtype = 'fp16' + elif mm.should_use_bf16(): print("Diffusion using bf16") dtype = torch.bfloat16 model_dtype = 'bf16' - elif mm.should_use_fp16(): - print("Diffusion using using fp16") - dtype = torch.float16 - model_dtype = 'fp16' else: print("Diffusion using using fp32") dtype = torch.float32