dtype cleanup

This commit is contained in:
City
2023-12-22 15:54:16 +01:00
parent 1a968d7bc8
commit 451c7588c7
4 changed files with 24 additions and 80 deletions
+15 -2
View File
@@ -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",)
+9 -2
View File
@@ -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",)
-20
View File
@@ -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"]:
-56
View File
@@ -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}'")