Add dtype selection to T5
This commit is contained in:
@@ -84,7 +84,7 @@ PixArt uses the same T5v1.1-xxl text encoder as DeepFloyd, so the T5 section of
|
|||||||
>
|
>
|
||||||
> Installing `xformers` is optional but strongly recommended as torch SDP is only partially implemented, if that.
|
> Installing `xformers` is optional but strongly recommended as torch SDP is only partially implemented, if that.
|
||||||
|
|
||||||
[Sample workflow here](https://github.com/city96/ComfyUI_ExtraModels/files/13192747/PixArt.json)
|
[Sample workflow here](https://github.com/city96/ComfyUI_ExtraModels/files/13481704/PixArtV2.json)
|
||||||
|
|
||||||

|

|
||||||
|
|
||||||
|
|||||||
+16
-8
@@ -8,7 +8,7 @@ import folder_paths
|
|||||||
from .t5v11 import T5v11Model, T5v11Tokenizer
|
from .t5v11 import T5v11Model, T5v11Tokenizer
|
||||||
|
|
||||||
class EXM_T5v11:
|
class EXM_T5v11:
|
||||||
def __init__(self, textmodel_ver="xxl", embedding_directory=None, textmodel_path=None, no_init=False, device="cpu"):
|
def __init__(self, textmodel_ver="xxl", embedding_directory=None, textmodel_path=None, no_init=False, device="cpu", dtype=None):
|
||||||
if no_init:
|
if no_init:
|
||||||
return
|
return
|
||||||
|
|
||||||
@@ -17,22 +17,27 @@ class EXM_T5v11:
|
|||||||
self.load_device = model_management.text_encoder_device()
|
self.load_device = model_management.text_encoder_device()
|
||||||
self.offload_device = model_management.text_encoder_offload_device()
|
self.offload_device = model_management.text_encoder_offload_device()
|
||||||
self.init_device = "cpu"
|
self.init_device = "cpu"
|
||||||
elif device == "bnb8bit":
|
elif dtype == "bnb8bit":
|
||||||
# BNB doesn't support size enum
|
# BNB doesn't support size enum
|
||||||
size = 12.4 * (1024**3)
|
size = 12.4 * (1024**3)
|
||||||
# Or moving between devices
|
# Or moving between devices
|
||||||
self.load_device = model_management.get_torch_device()
|
self.load_device = model_management.get_torch_device()
|
||||||
self.offload_device = self.load_device
|
self.offload_device = self.load_device
|
||||||
self.init_device = self.load_device
|
self.init_device = self.load_device
|
||||||
elif device == "bnb4bit":
|
elif dtype == "bnb4bit":
|
||||||
# This seems to use the same VRAM as 8bit on Pascal?
|
# This seems to use the same VRAM as 8bit on Pascal?
|
||||||
size = 6.2 * (1024**3)
|
size = 6.2 * (1024**3)
|
||||||
self.load_device = model_management.get_torch_device()
|
self.load_device = model_management.get_torch_device()
|
||||||
self.offload_device = self.load_device
|
self.offload_device = self.load_device
|
||||||
self.init_device = self.load_device
|
self.init_device = self.load_device
|
||||||
|
elif device == "cpu":
|
||||||
|
size = 0
|
||||||
|
self.load_device = "cpu"
|
||||||
|
self.offload_device = "cpu"
|
||||||
|
self.init_device="cpu"
|
||||||
else:
|
else:
|
||||||
size = 0
|
size = 0
|
||||||
self.load_device = device
|
self.load_device = model_management.get_torch_device()
|
||||||
self.offload_device = "cpu"
|
self.offload_device = "cpu"
|
||||||
self.init_device="cpu"
|
self.init_device="cpu"
|
||||||
|
|
||||||
@@ -40,6 +45,7 @@ class EXM_T5v11:
|
|||||||
textmodel_ver = textmodel_ver,
|
textmodel_ver = textmodel_ver,
|
||||||
textmodel_path = textmodel_path,
|
textmodel_path = textmodel_path,
|
||||||
device = device,
|
device = device,
|
||||||
|
dtype = dtype,
|
||||||
)
|
)
|
||||||
self.tokenizer = T5v11Tokenizer(embedding_directory=embedding_directory)
|
self.tokenizer = T5v11Tokenizer(embedding_directory=embedding_directory)
|
||||||
self.patcher = comfy.model_patcher.ModelPatcher(
|
self.patcher = comfy.model_patcher.ModelPatcher(
|
||||||
@@ -86,11 +92,13 @@ class EXM_T5v11:
|
|||||||
return self.patcher.get_key_patches()
|
return self.patcher.get_key_patches()
|
||||||
|
|
||||||
|
|
||||||
def load_t5(model_type, model_ver, model_path, path_type="file", device="cpu"):
|
def load_t5(model_type, model_ver, model_path, path_type="file", device="cpu", dtype=None):
|
||||||
assert model_type in ["t5v11"] # Only supported model for now
|
assert model_type in ["t5v11"] # Only supported model for now
|
||||||
model_args = {}
|
model_args = {
|
||||||
model_args["textmodel_ver"] = model_ver
|
"textmodel_ver" : model_ver,
|
||||||
model_args["device"] = device
|
"device" : device,
|
||||||
|
"dtype" : dtype,
|
||||||
|
}
|
||||||
|
|
||||||
if path_type == "folder":
|
if path_type == "folder":
|
||||||
# pass directly to transformers and initialize there
|
# pass directly to transformers and initialize there
|
||||||
|
|||||||
+11
-2
@@ -4,6 +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
|
||||||
|
|
||||||
# initialize custom folder path
|
# initialize custom folder path
|
||||||
# TODO: integrate with `extra_model_paths.yaml`
|
# TODO: integrate with `extra_model_paths.yaml`
|
||||||
@@ -24,7 +25,8 @@ class T5v11Loader:
|
|||||||
"t5v11_name": (folder_paths.get_filename_list("t5"),),
|
"t5v11_name": (folder_paths.get_filename_list("t5"),),
|
||||||
"t5v11_ver": (["xxl"],),
|
"t5v11_ver": (["xxl"],),
|
||||||
"path_type": (["folder", "file"],),
|
"path_type": (["folder", "file"],),
|
||||||
"device": (["auto", "cpu", "bnb8bit", "bnb4bit"],{"default":"cpu"})
|
"device": (["auto", "cpu", "gpu"],{"default":"cpu"}),
|
||||||
|
"dtype": (dtype_list,),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
RETURN_TYPES = ("T5",)
|
RETURN_TYPES = ("T5",)
|
||||||
@@ -32,13 +34,20 @@ class T5v11Loader:
|
|||||||
CATEGORY = "ExtraModels/T5"
|
CATEGORY = "ExtraModels/T5"
|
||||||
TITLE = "T5v1.1 Loader"
|
TITLE = "T5v1.1 Loader"
|
||||||
|
|
||||||
def load_model(self, t5v11_name, t5v11_ver, path_type, device):
|
def load_model(self, t5v11_name, t5v11_ver, path_type, device, dtype):
|
||||||
|
if "bnb" in dtype:
|
||||||
|
assert device == "gpu", "BitsAndBytes only works on CUDA! Set device to 'gpu'."
|
||||||
|
dtype = string_to_dtype(dtype, "text_encoder")
|
||||||
|
if device == "cpu":
|
||||||
|
assert dtype in [None, torch.float32], f"Can't use dtype '{dtype}' with CPU! Set dtype to 'default'."
|
||||||
|
|
||||||
return (load_t5(
|
return (load_t5(
|
||||||
model_type = "t5v11",
|
model_type = "t5v11",
|
||||||
model_ver = t5v11_ver,
|
model_ver = t5v11_ver,
|
||||||
model_path = folder_paths.get_full_path("t5", t5v11_name),
|
model_path = folder_paths.get_full_path("t5", t5v11_name),
|
||||||
path_type = path_type,
|
path_type = path_type,
|
||||||
device = device,
|
device = device,
|
||||||
|
dtype = dtype,
|
||||||
),)
|
),)
|
||||||
|
|
||||||
class T5TextEncode:
|
class T5TextEncode:
|
||||||
|
|||||||
+5
-6
@@ -20,16 +20,19 @@ class T5v11Model(torch.nn.Module):
|
|||||||
|
|
||||||
self.num_layers = 24
|
self.num_layers = 24
|
||||||
self.max_length = max_length
|
self.max_length = max_length
|
||||||
|
self.bnb = False
|
||||||
|
|
||||||
if textmodel_path is not None:
|
if textmodel_path is not None:
|
||||||
model_args = {}
|
model_args = {}
|
||||||
model_args["low_cpu_mem_usage"] = True # Don't take 2x system ram on cpu
|
model_args["low_cpu_mem_usage"] = True # Don't take 2x system ram on cpu
|
||||||
if device == "bnb8bit":
|
if dtype == "bnb8bit":
|
||||||
self.bnb = True
|
self.bnb = True
|
||||||
model_args["load_in_8bit"] = True
|
model_args["load_in_8bit"] = True
|
||||||
elif device == "bnb4bit":
|
elif dtype == "bnb4bit":
|
||||||
self.bnb = True
|
self.bnb = True
|
||||||
model_args["load_in_4bit"] = True
|
model_args["load_in_4bit"] = True
|
||||||
else:
|
else:
|
||||||
|
if dtype: model_args["torch_dtype"] = dtype
|
||||||
self.bnb = False
|
self.bnb = False
|
||||||
# TODO: custom device map?
|
# TODO: custom device map?
|
||||||
print(f"Loading T5 from '{textmodel_path}'")
|
print(f"Loading T5 from '{textmodel_path}'")
|
||||||
@@ -46,10 +49,6 @@ class T5v11Model(torch.nn.Module):
|
|||||||
with modeling_utils.no_init_weights():
|
with modeling_utils.no_init_weights():
|
||||||
self.transformer = T5EncoderModel(config)
|
self.transformer = T5EncoderModel(config)
|
||||||
|
|
||||||
if dtype is not None and not self.bnb:
|
|
||||||
self.transformer.to(dtype)
|
|
||||||
self.transformer.encoder.embeddings.embed_tokens.to(torch.float32)
|
|
||||||
|
|
||||||
if freeze:
|
if freeze:
|
||||||
self.freeze()
|
self.freeze()
|
||||||
self.empty_tokens = [[0] * self.max_length] # <pad> token
|
self.empty_tokens = [[0] * self.max_length] # <pad> token
|
||||||
|
|||||||
+3
-10
@@ -4,20 +4,13 @@ import comfy.utils
|
|||||||
from comfy import model_management
|
from comfy import model_management
|
||||||
from comfy import diffusers_convert
|
from comfy import diffusers_convert
|
||||||
|
|
||||||
vae_dtype_dict = {
|
|
||||||
"auto" : model_management.vae_device(),
|
|
||||||
"fp32" : torch.float32,
|
|
||||||
"fp16" : torch.float16,
|
|
||||||
"bf16" : torch.bfloat16,
|
|
||||||
}
|
|
||||||
|
|
||||||
class EXVAE(comfy.sd.VAE):
|
class EXVAE(comfy.sd.VAE):
|
||||||
def __init__(self, model_path, model_conf, dtype=None):
|
def __init__(self, model_path, model_conf, dtype=torch.float32):
|
||||||
self.latent_dim = model_conf["embed_dim"]
|
self.latent_dim = model_conf["embed_dim"]
|
||||||
self.latent_scale = model_conf["embed_scale"]
|
self.latent_scale = model_conf["embed_scale"]
|
||||||
self.device = model_management.vae_device()
|
self.device = model_management.vae_device()
|
||||||
self.offload_device = model_management.vae_offload_device()
|
self.offload_device = model_management.vae_offload_device()
|
||||||
self.vae_dtype = vae_dtype_dict.get(dtype, "auto")
|
self.vae_dtype = dtype
|
||||||
|
|
||||||
sd = comfy.utils.load_torch_file(model_path)
|
sd = comfy.utils.load_torch_file(model_path)
|
||||||
model = None
|
model = None
|
||||||
@@ -104,7 +97,7 @@ class EXVAE(comfy.sd.VAE):
|
|||||||
free_memory = model_management.get_free_memory(self.device)
|
free_memory = model_management.get_free_memory(self.device)
|
||||||
batch_number = int(free_memory / memory_used)
|
batch_number = int(free_memory / memory_used)
|
||||||
batch_number = max(1, batch_number)
|
batch_number = max(1, batch_number)
|
||||||
samples = torch.empty((pixel_samples.shape[0], self.first_stage_model.embed_dim, round(pixel_samples.shape[2] // self.latent_scale), round(pixel_samples.shape[3] // self.latent_scale)), device="cpu")
|
samples = torch.empty((pixel_samples.shape[0], self.latent_dim, round(pixel_samples.shape[2] // self.latent_scale), round(pixel_samples.shape[3] // self.latent_scale)), device="cpu")
|
||||||
for x in range(0, pixel_samples.shape[0], batch_number):
|
for x in range(0, pixel_samples.shape[0], batch_number):
|
||||||
pixels_in = (2. * pixel_samples[x:x+batch_number] - 1.).to(self.vae_dtype).to(self.device)
|
pixels_in = (2. * pixel_samples[x:x+batch_number] - 1.).to(self.vae_dtype).to(self.device)
|
||||||
samples[x:x+batch_number] = self.first_stage_model.encode(pixels_in).cpu().float()
|
samples[x:x+batch_number] = self.first_stage_model.encode(pixels_in).cpu().float()
|
||||||
|
|||||||
+4
-2
@@ -3,6 +3,8 @@ 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
|
||||||
|
|
||||||
class ExtraVAELoader:
|
class ExtraVAELoader:
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(s):
|
def INPUT_TYPES(s):
|
||||||
@@ -10,7 +12,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" : (["auto","fp32","fp16","bf16"],),
|
"dtype" : (dtype_list_short,),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
RETURN_TYPES = ("VAE",)
|
RETURN_TYPES = ("VAE",)
|
||||||
@@ -21,7 +23,7 @@ class ExtraVAELoader:
|
|||||||
def load_vae(self, vae_name, vae_type, dtype):
|
def load_vae(self, vae_name, vae_type, dtype):
|
||||||
model_path = folder_paths.get_full_path("vae", vae_name)
|
model_path = folder_paths.get_full_path("vae", vae_name)
|
||||||
model_conf = vae_conf[vae_type]
|
model_conf = vae_conf[vae_type]
|
||||||
vae = EXVAE(model_path, model_conf, dtype)
|
vae = EXVAE(model_path, model_conf, string_to_dtype(dtype, "vae"))
|
||||||
return (vae,)
|
return (vae,)
|
||||||
|
|
||||||
NODE_CLASS_MAPPINGS = {
|
NODE_CLASS_MAPPINGS = {
|
||||||
|
|||||||
@@ -0,0 +1,56 @@
|
|||||||
|
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"]:
|
||||||
|
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}'")
|
||||||
@@ -0,0 +1,56 @@
|
|||||||
|
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