Add dtype selection to T5

This commit is contained in:
City
2023-11-28 00:49:06 +01:00
parent 33e5db52a3
commit 389c16f2f5
8 changed files with 152 additions and 29 deletions
+1 -1
View File
@@ -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)
![PixArtT10](https://github.com/city96/ComfyUI_ExtraModels/assets/125218114/e709da33-d313-43a0-bdbf-b6e84bde14e7) ![PixArtT10](https://github.com/city96/ComfyUI_ExtraModels/assets/125218114/e709da33-d313-43a0-bdbf-b6e84bde14e7)
+16 -8
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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 = {
+56
View File
@@ -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}'")
+56
View File
@@ -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}'")