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