diff --git a/README.md b/README.md index 151ca3e..9bb1221 100644 --- a/README.md +++ b/README.md @@ -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. -[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) diff --git a/T5/loader.py b/T5/loader.py index fee4472..70cec7b 100644 --- a/T5/loader.py +++ b/T5/loader.py @@ -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 diff --git a/T5/nodes.py b/T5/nodes.py index a7ef50c..f50ddb3 100644 --- a/T5/nodes.py +++ b/T5/nodes.py @@ -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: diff --git a/T5/t5v11.py b/T5/t5v11.py index 8105a4f..23876d0 100644 --- a/T5/t5v11.py +++ b/T5/t5v11.py @@ -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] # token diff --git a/VAE/loader.py b/VAE/loader.py index c8a59b8..ccd4fb2 100644 --- a/VAE/loader.py +++ b/VAE/loader.py @@ -4,20 +4,13 @@ import comfy.utils from comfy import model_management 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): - 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_scale = model_conf["embed_scale"] self.device = model_management.vae_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) model = None @@ -104,7 +97,7 @@ class EXVAE(comfy.sd.VAE): free_memory = model_management.get_free_memory(self.device) batch_number = int(free_memory / memory_used) 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): 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() diff --git a/VAE/nodes.py b/VAE/nodes.py index a62ea5c..c5c10f9 100644 --- a/VAE/nodes.py +++ b/VAE/nodes.py @@ -3,6 +3,8 @@ import folder_paths from .conf import vae_conf from .loader import EXVAE +from ..utils.dtype import string_to_dtype, dtype_list_short + class ExtraVAELoader: @classmethod def INPUT_TYPES(s): @@ -10,7 +12,7 @@ class ExtraVAELoader: "required": { "vae_name": (folder_paths.get_filename_list("vae"),), "vae_type": (list(vae_conf.keys()), {"default":"kl-f8"}), - "dtype" : (["auto","fp32","fp16","bf16"],), + "dtype" : (dtype_list_short,), } } RETURN_TYPES = ("VAE",) @@ -21,7 +23,7 @@ class ExtraVAELoader: def load_vae(self, vae_name, vae_type, dtype): model_path = folder_paths.get_full_path("vae", vae_name) 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,) NODE_CLASS_MAPPINGS = { diff --git a/utils/dtype.py b/utils/dtype.py new file mode 100644 index 0000000..a5bf5fb --- /dev/null +++ b/utils/dtype.py @@ -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}'") diff --git a/utils/lookup.py b/utils/lookup.py new file mode 100644 index 0000000..65f5124 --- /dev/null +++ b/utils/lookup.py @@ -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}'")