dtype cleanup
This commit is contained in:
+15
-2
@@ -4,7 +4,7 @@ import torch
|
|||||||
import folder_paths
|
import folder_paths
|
||||||
|
|
||||||
from .loader import load_t5
|
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
|
# initialize custom folder path
|
||||||
# TODO: integrate with `extra_model_paths.yaml`
|
# TODO: integrate with `extra_model_paths.yaml`
|
||||||
@@ -17,6 +17,19 @@ folder_paths.folder_names_and_paths["t5"] = (
|
|||||||
folder_paths.supported_pt_extensions
|
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:
|
class T5v11Loader:
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(s):
|
def INPUT_TYPES(s):
|
||||||
@@ -26,7 +39,7 @@ class T5v11Loader:
|
|||||||
"t5v11_ver": (["xxl"],),
|
"t5v11_ver": (["xxl"],),
|
||||||
"path_type": (["folder", "file"],),
|
"path_type": (["folder", "file"],),
|
||||||
"device": (["auto", "cpu", "gpu"],{"default":"cpu"}),
|
"device": (["auto", "cpu", "gpu"],{"default":"cpu"}),
|
||||||
"dtype": (dtype_list,),
|
"dtype": (dtypes,),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
RETURN_TYPES = ("T5",)
|
RETURN_TYPES = ("T5",)
|
||||||
|
|||||||
+9
-2
@@ -3,7 +3,14 @@ import folder_paths
|
|||||||
from .conf import vae_conf
|
from .conf import vae_conf
|
||||||
from .loader import EXVAE
|
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:
|
class ExtraVAELoader:
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -12,7 +19,7 @@ class ExtraVAELoader:
|
|||||||
"required": {
|
"required": {
|
||||||
"vae_name": (folder_paths.get_filename_list("vae"),),
|
"vae_name": (folder_paths.get_filename_list("vae"),),
|
||||||
"vae_type": (list(vae_conf.keys()), {"default":"kl-f8"}),
|
"vae_type": (list(vae_conf.keys()), {"default":"kl-f8"}),
|
||||||
"dtype" : ([*dtype_list_short, "BF16"],),
|
"dtype" : (dtypes,),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
RETURN_TYPES = ("VAE",)
|
RETURN_TYPES = ("VAE",)
|
||||||
|
|||||||
@@ -1,26 +1,6 @@
|
|||||||
import torch
|
import torch
|
||||||
from comfy import model_management
|
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):
|
def string_to_dtype(s="none", mode=None):
|
||||||
s = s.lower().strip()
|
s = s.lower().strip()
|
||||||
if s in ["default", "as-is"]:
|
if s in ["default", "as-is"]:
|
||||||
|
|||||||
@@ -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}'")
|
|
||||||
Reference in New Issue
Block a user