Add dtype selection to T5
This commit is contained in:
+16
-8
@@ -8,7 +8,7 @@ import folder_paths
|
||||
from .t5v11 import T5v11Model, T5v11Tokenizer
|
||||
|
||||
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:
|
||||
return
|
||||
|
||||
@@ -17,22 +17,27 @@ class EXM_T5v11:
|
||||
self.load_device = model_management.text_encoder_device()
|
||||
self.offload_device = model_management.text_encoder_offload_device()
|
||||
self.init_device = "cpu"
|
||||
elif device == "bnb8bit":
|
||||
elif dtype == "bnb8bit":
|
||||
# BNB doesn't support size enum
|
||||
size = 12.4 * (1024**3)
|
||||
# Or moving between devices
|
||||
self.load_device = model_management.get_torch_device()
|
||||
self.offload_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?
|
||||
size = 6.2 * (1024**3)
|
||||
self.load_device = model_management.get_torch_device()
|
||||
self.offload_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:
|
||||
size = 0
|
||||
self.load_device = device
|
||||
self.load_device = model_management.get_torch_device()
|
||||
self.offload_device = "cpu"
|
||||
self.init_device="cpu"
|
||||
|
||||
@@ -40,6 +45,7 @@ class EXM_T5v11:
|
||||
textmodel_ver = textmodel_ver,
|
||||
textmodel_path = textmodel_path,
|
||||
device = device,
|
||||
dtype = dtype,
|
||||
)
|
||||
self.tokenizer = T5v11Tokenizer(embedding_directory=embedding_directory)
|
||||
self.patcher = comfy.model_patcher.ModelPatcher(
|
||||
@@ -86,11 +92,13 @@ class EXM_T5v11:
|
||||
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
|
||||
model_args = {}
|
||||
model_args["textmodel_ver"] = model_ver
|
||||
model_args["device"] = device
|
||||
model_args = {
|
||||
"textmodel_ver" : model_ver,
|
||||
"device" : device,
|
||||
"dtype" : dtype,
|
||||
}
|
||||
|
||||
if path_type == "folder":
|
||||
# pass directly to transformers and initialize there
|
||||
|
||||
+11
-2
@@ -4,6 +4,7 @@ import torch
|
||||
import folder_paths
|
||||
|
||||
from .loader import load_t5
|
||||
from ..utils.dtype import string_to_dtype, dtype_list
|
||||
|
||||
# initialize custom folder path
|
||||
# TODO: integrate with `extra_model_paths.yaml`
|
||||
@@ -24,7 +25,8 @@ class T5v11Loader:
|
||||
"t5v11_name": (folder_paths.get_filename_list("t5"),),
|
||||
"t5v11_ver": (["xxl"],),
|
||||
"path_type": (["folder", "file"],),
|
||||
"device": (["auto", "cpu", "bnb8bit", "bnb4bit"],{"default":"cpu"})
|
||||
"device": (["auto", "cpu", "gpu"],{"default":"cpu"}),
|
||||
"dtype": (dtype_list,),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("T5",)
|
||||
@@ -32,13 +34,20 @@ class T5v11Loader:
|
||||
CATEGORY = "ExtraModels/T5"
|
||||
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(
|
||||
model_type = "t5v11",
|
||||
model_ver = t5v11_ver,
|
||||
model_path = folder_paths.get_full_path("t5", t5v11_name),
|
||||
path_type = path_type,
|
||||
device = device,
|
||||
dtype = dtype,
|
||||
),)
|
||||
|
||||
class T5TextEncode:
|
||||
|
||||
+5
-6
@@ -20,16 +20,19 @@ class T5v11Model(torch.nn.Module):
|
||||
|
||||
self.num_layers = 24
|
||||
self.max_length = max_length
|
||||
self.bnb = False
|
||||
|
||||
if textmodel_path is not None:
|
||||
model_args = {}
|
||||
model_args["low_cpu_mem_usage"] = True # Don't take 2x system ram on cpu
|
||||
if device == "bnb8bit":
|
||||
if dtype == "bnb8bit":
|
||||
self.bnb = True
|
||||
model_args["load_in_8bit"] = True
|
||||
elif device == "bnb4bit":
|
||||
elif dtype == "bnb4bit":
|
||||
self.bnb = True
|
||||
model_args["load_in_4bit"] = True
|
||||
else:
|
||||
if dtype: model_args["torch_dtype"] = dtype
|
||||
self.bnb = False
|
||||
# TODO: custom device map?
|
||||
print(f"Loading T5 from '{textmodel_path}'")
|
||||
@@ -46,10 +49,6 @@ class T5v11Model(torch.nn.Module):
|
||||
with modeling_utils.no_init_weights():
|
||||
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:
|
||||
self.freeze()
|
||||
self.empty_tokens = [[0] * self.max_length] # <pad> token
|
||||
|
||||
Reference in New Issue
Block a user