From b261c66f29b64fe789abf577ad47605273627747 Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Fri, 8 Dec 2023 18:40:25 +0100 Subject: [PATCH] Fix DiT sampling Apply the PixArt fix here as well. Move models to checkpoints folder as that makes more sense in this case. --- DiT/conf.py | 139 +++++++++++++++++++++++++++++++++----------------- DiT/loader.py | 28 ++++++---- DiT/model.py | 1 + DiT/nodes.py | 21 ++------ README.md | 4 +- 5 files changed, 117 insertions(+), 76 deletions(-) diff --git a/DiT/conf.py b/DiT/conf.py index 12f44b3..e3071aa 100644 --- a/DiT/conf.py +++ b/DiT/conf.py @@ -1,77 +1,120 @@ """ List of all DiT model types / settings """ +sampling_settings = { + "beta_schedule" : "sqrt_linear", + "linear_start" : 0.0001, + "linear_end" : 0.02, + "timesteps" : 1000, +} + dit_conf = { "XL/2": { # DiT_XL_2 - "depth" : 28, - "num_heads" : 16, - "patch_size" : 2, - "hidden_size" : 1152, + "unet_config": { + "depth" : 28, + "num_heads" : 16, + "patch_size" : 2, + "hidden_size" : 1152, + }, + "sampling_settings" : sampling_settings, }, "XL/4": { # DiT_XL_4 - "depth" : 28, - "num_heads" : 16, - "patch_size" : 4, - "hidden_size" : 1152, + "unet_config": { + "depth" : 28, + "num_heads" : 16, + "patch_size" : 4, + "hidden_size" : 1152, + }, + "sampling_settings" : sampling_settings, }, "XL/8": { # DiT_XL_8 - "depth" : 28, - "num_heads" : 16, - "patch_size" : 8, - "hidden_size" : 1152, + "unet_config": { + "depth" : 28, + "num_heads" : 16, + "patch_size" : 8, + "hidden_size" : 1152, + }, + "sampling_settings" : sampling_settings, }, "L/2": { # DiT_L_2 - "depth" : 24, - "num_heads" : 16, - "patch_size" : 2, - "hidden_size" : 1024, + "unet_config": { + "depth" : 24, + "num_heads" : 16, + "patch_size" : 2, + "hidden_size" : 1024, + }, + "sampling_settings" : sampling_settings, }, "L/4": { # DiT_L_4 - "depth" : 24, - "num_heads" : 16, - "patch_size" : 4, - "hidden_size" : 1024, + "unet_config": { + "depth" : 24, + "num_heads" : 16, + "patch_size" : 4, + "hidden_size" : 1024, + }, + "sampling_settings" : sampling_settings, }, "L/8": { # DiT_L_8 - "depth" : 24, - "num_heads" : 16, - "patch_size" : 8, - "hidden_size" : 1024, + "unet_config": { + "depth" : 24, + "num_heads" : 16, + "patch_size" : 8, + "hidden_size" : 1024, + }, + "sampling_settings" : sampling_settings, }, "B/2": { # DiT_B_2 - "depth" : 12, - "num_heads" : 12, - "patch_size" : 2, - "hidden_size" : 768, + "unet_config": { + "depth" : 12, + "num_heads" : 12, + "patch_size" : 2, + "hidden_size" : 768, + }, + "sampling_settings" : sampling_settings, }, "B/4": { # DiT_B_4 - "depth" : 12, - "num_heads" : 12, - "patch_size" : 4, - "hidden_size" : 768, + "unet_config": { + "depth" : 12, + "num_heads" : 12, + "patch_size" : 4, + "hidden_size" : 768, + }, + "sampling_settings" : sampling_settings, }, "B/8": { # DiT_B_8 - "depth" : 12, - "num_heads" : 12, - "patch_size" : 8, - "hidden_size" : 768, + "unet_config": { + "depth" : 12, + "num_heads" : 12, + "patch_size" : 8, + "hidden_size" : 768, + }, + "sampling_settings" : sampling_settings, }, "S/2": { # DiT_S_2 - "depth" : 12, - "num_heads" : 6, - "patch_size" : 2, - "hidden_size" : 384, + "unet_config": { + "depth" : 12, + "num_heads" : 6, + "patch_size" : 2, + "hidden_size" : 384, + }, + "sampling_settings" : sampling_settings, }, "S/4": { # DiT_S_4 - "depth" : 12, - "num_heads" : 6, - "patch_size" : 4, - "hidden_size" : 384, + "unet_config": { + "depth" : 12, + "num_heads" : 6, + "patch_size" : 4, + "hidden_size" : 384, + }, + "sampling_settings" : sampling_settings, }, "S/8": { # DiT_S_8 - "depth" : 12, - "num_heads" : 6, - "patch_size" : 8, - "hidden_size" : 384, + "unet_config": { + "depth" : 12, + "num_heads" : 6, + "patch_size" : 8, + "hidden_size" : 384, + }, + "sampling_settings" : sampling_settings, }, } diff --git a/DiT/loader.py b/DiT/loader.py index adafe94..daad506 100644 --- a/DiT/loader.py +++ b/DiT/loader.py @@ -1,5 +1,4 @@ import comfy.supported_models_base -import comfy.supported_models import comfy.latent_formats import comfy.model_patcher import comfy.model_base @@ -7,11 +6,17 @@ import comfy.utils import torch from comfy import model_management -from .model import DiT - -class EXMDiT(comfy.supported_models.SD15): +class EXM_DiT(comfy.supported_models_base.BASE): unet_config = {} unet_extra_config = {} + latent_format = comfy.latent_formats.SD15 + + def __init__(self, model_conf): + self.unet_config = model_conf.get("unet_config", {}) + self.sampling_settings = model_conf.get("sampling_settings", {}) + self.latent_format = self.latent_format() + # UNET is handled by extension + self.unet_config["disable_unet_model_creation"] = True def model_type(self, state_dict, prefix=""): return comfy.model_base.ModelType.EPS @@ -22,18 +27,21 @@ def load_dit(model_path, model_conf): parameters = comfy.utils.calculate_parameters(state_dict) unet_dtype = model_management.unet_dtype(model_params=parameters) - offload_device = model_management.unet_offload_device() + model_conf["unet_config"]["num_classes"] = state_dict["y_embedder.embedding_table.weight"].shape[0] - 1 # adj. for empty + + model_conf = EXM_DiT(model_conf) model = comfy.model_base.BaseModel( - EXMDiT({"disable_unet_model_creation" : True }), + model_conf, model_type=comfy.model_base.ModelType.EPS, device=model_management.get_torch_device() ) - model_conf["num_classes"] = state_dict["y_embedder.embedding_table.weight"].shape[0] - 1 # adj. for empty - model.dit_config = model_conf - model.diffusion_model = DiT(**model_conf).eval() + + from .model import DiT + model.diffusion_model = DiT(**model_conf.unet_config) + model.diffusion_model.load_state_dict(state_dict) - model.diffusion_model.eval() model.diffusion_model.dtype = unet_dtype + model.diffusion_model.eval() model.diffusion_model.to(unet_dtype) model_patcher = comfy.model_patcher.ModelPatcher( diff --git a/DiT/model.py b/DiT/model.py index 74103ee..e08200d 100644 --- a/DiT/model.py +++ b/DiT/model.py @@ -158,6 +158,7 @@ class DiT(nn.Module): class_dropout_prob=0.1, num_classes=1000, learn_sigma=True, + **kwargs, ): super().__init__() self.learn_sigma = learn_sigma diff --git a/DiT/nodes.py b/DiT/nodes.py index 1b05f79..0ad00ef 100644 --- a/DiT/nodes.py +++ b/DiT/nodes.py @@ -6,23 +6,12 @@ import folder_paths from .conf import dit_conf from .loader import load_dit -# initialize custom folder path -# TODO: integrate with `extra_model_paths.yaml` -os.makedirs( - os.path.join(folder_paths.models_dir,"dit"), - exist_ok = True, -) -folder_paths.folder_names_and_paths["dit"] = ( - [os.path.join(folder_paths.models_dir,"dit")], - folder_paths.supported_pt_extensions -) - class DitCheckpointLoader: @classmethod def INPUT_TYPES(s): return { "required": { - "ckpt_name": (folder_paths.get_filename_list("dit"),), + "ckpt_name": (folder_paths.get_filename_list("checkpoints"),), "model": (list(dit_conf.keys()),), "image_size": ([256, 512],), # "num_classes": ("INT", {"default": 1000, "min": 0,}), @@ -35,10 +24,10 @@ class DitCheckpointLoader: TITLE = "DitCheckpointLoader" def load_checkpoint(self, ckpt_name, model, image_size): - ckpt_path = folder_paths.get_full_path("dit", ckpt_name) + ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name) model_conf = dit_conf[model] - model_conf["input_size"] = image_size // 8 - # model_conf["num_classes"] = num_classes + model_conf["unet_config"]["input_size"] = image_size // 8 + # model_conf["unet_config"]["num_classes"] = num_classes dit = load_dit( model_path = ckpt_path, model_conf = model_conf, @@ -98,7 +87,7 @@ class DiTCondLabelEmpty: def cond_empty(self, model): # [ID of last class + 1] == [num_classes] - y_null = model.model.dit_config["num_classes"] + y_null = model.model.model_config.unet_config["num_classes"] y = torch.tensor([[y_null]]).to(torch.int) return ([[y, {}]], ) diff --git a/README.md b/README.md index 6ad2e51..6e744a3 100644 --- a/README.md +++ b/README.md @@ -85,13 +85,13 @@ Limitations: ### Usage 1. Download the original model weights from the [DiT Repo](https://github.com/facebookresearch/DiT) or the converted [FP16 safetensor ones from Huggingface](https://huggingface.co/city96/DiT/tree/main). -2. Place them in `ComfyUI\models\dit` (created on first run after installing the extension) +2. Place them in your checkpoints folder. (You may need to move them if you had them in `ComfyUI\models\dit` before) 3. Load the model and select the class labels as shown in the image below 4. **Make sure to use the Empty label conditioning for the Negative input of the KSampler!** ConditioningCombine nodes *should* work for combining multiple labels. The area ones don't since the model currently can't handle dynamic input dimensions. -[Image with sample workflow](https://github.com/city96/ComfyUI_ExtraModels/assets/125218114/33bfb812-23ea-4bb0-b1e2-082756e53010) +[Sample workflow here](https://github.com/city96/ComfyUI_ExtraModels/files/13619259/DiTV2.json) ![DIT_WORKFLOW_IMG](https://github.com/city96/ComfyUI_ExtraModels/assets/125218114/cdd4ec94-b0eb-436a-bf23-a3bcef8d7b90)