From 451c7588c7d0188f03c24b729d5e83ec058794cd Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Fri, 22 Dec 2023 15:54:16 +0100 Subject: [PATCH] dtype cleanup --- T5/nodes.py | 17 +++++++++++++-- VAE/nodes.py | 11 ++++++++-- utils/dtype.py | 20 ------------------ utils/lookup.py | 56 ------------------------------------------------- 4 files changed, 24 insertions(+), 80 deletions(-) delete mode 100644 utils/lookup.py diff --git a/T5/nodes.py b/T5/nodes.py index f50ddb3..d5436c8 100644 --- a/T5/nodes.py +++ b/T5/nodes.py @@ -4,7 +4,7 @@ import torch import folder_paths from .loader import load_t5 -from ..utils.dtype import string_to_dtype, dtype_list +from ..utils.dtype import string_to_dtype # initialize custom folder path # TODO: integrate with `extra_model_paths.yaml` @@ -17,6 +17,19 @@ folder_paths.folder_names_and_paths["t5"] = ( folder_paths.supported_pt_extensions ) +dtypes = [ + "default", + "auto (comfy)", + "FP32", + "FP16", + # Note: remove these at some point + "bnb8bit", + "bnb4bit", +] +try: torch.float8_e5m2 +except AttributeError: print("Torch version too old for FP8") +else: dtypes += ["FP8 E4M3", "FP8 E5M2"] + class T5v11Loader: @classmethod def INPUT_TYPES(s): @@ -26,7 +39,7 @@ class T5v11Loader: "t5v11_ver": (["xxl"],), "path_type": (["folder", "file"],), "device": (["auto", "cpu", "gpu"],{"default":"cpu"}), - "dtype": (dtype_list,), + "dtype": (dtypes,), } } RETURN_TYPES = ("T5",) diff --git a/VAE/nodes.py b/VAE/nodes.py index b5ad1ca..b3639ae 100644 --- a/VAE/nodes.py +++ b/VAE/nodes.py @@ -3,7 +3,14 @@ import folder_paths from .conf import vae_conf from .loader import EXVAE -from ..utils.dtype import string_to_dtype, dtype_list_short +from ..utils.dtype import string_to_dtype + +dtypes = [ + "auto", + "FP32", + "FP16", + "BF16" +] class ExtraVAELoader: @classmethod @@ -12,7 +19,7 @@ class ExtraVAELoader: "required": { "vae_name": (folder_paths.get_filename_list("vae"),), "vae_type": (list(vae_conf.keys()), {"default":"kl-f8"}), - "dtype" : ([*dtype_list_short, "BF16"],), + "dtype" : (dtypes,), } } RETURN_TYPES = ("VAE",) diff --git a/utils/dtype.py b/utils/dtype.py index 03121cd..7ad6fcb 100644 --- a/utils/dtype.py +++ b/utils/dtype.py @@ -1,26 +1,6 @@ import torch from comfy import model_management -dtype_list_torch = [ - "FP32", - "FP16", - # "FP8 E4M3", # ToDo: add these - # "FP8 E5M2", -] - -dtype_list = [ - "default", - "auto (comfy)", - *dtype_list_torch, - "bnb8bit", - "bnb4bit" -] - -dtype_list_short = [ - "auto", - *dtype_list_torch, -] - def string_to_dtype(s="none", mode=None): s = s.lower().strip() if s in ["default", "as-is"]: diff --git a/utils/lookup.py b/utils/lookup.py deleted file mode 100644 index 65f5124..0000000 --- a/utils/lookup.py +++ /dev/null @@ -1,56 +0,0 @@ -import torch -from comfy import model_management - -dtype_list_torch = [ - "FP32", - "FP16", - # "FP8 E4M3", # ToDo: add these - # "FP8 E5M2", -] - -dtype_list_short = [ - "auto", - *dtype_list_torch, -] - -dtype_list = [ - "default", - "auto (comfy)", - *dtype_list_torch, - "bnb8bit", - "bnb4bit" -] - -def string_to_dtype(s="none", mode=None): - s = s.lower().strip() - if s in ["default", "as-is"]: - return None - elif s in ["auto", "auto (comfy)"]: - if mode == "vae": - return model_management.vae_device() - elif mode == "text_encoder": - return model_management.text_encoder_dtype() - elif mode == "unet": - return model_management.unet_dtype() - else: - raise NotImplementedError(f"Unknown dtype mode '{mode}'") - elif s in ["none", "auto (hf)", "auto (hf/bnb)"]: - return None - elif s in ["fp32", "float32", "float"]: - return torch.float32 - elif s in ["fp16", "float16", "half"]: - return torch.float16 - elif "fp8" in s or "float8" in s: - if "e5m2" in s: - return torch.float8_e5m2 - elif "e4m3" in s: - return torch.float8_e4m3fn - else: - raise NotImplementedError(f"Unknown 8bit dtype '{s}'") - elif "bnb" in s: - assert s in ["bnb8bit", "bnb4bit"], f"Unknown bnb mode '{s}'" - return s - elif s is None: - return None - else: - raise NotImplementedError(f"Unknown dtype '{s}'")