Fix DiT sampling

Apply the PixArt fix here as well.

Move models to checkpoints folder as that makes more sense in this case.
This commit is contained in:
City
2023-12-08 18:40:25 +01:00
parent 80d5d9299b
commit b261c66f29
5 changed files with 117 additions and 76 deletions
+91 -48
View File
@@ -1,77 +1,120 @@
""" """
List of all DiT model types / settings 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 = { dit_conf = {
"XL/2": { # DiT_XL_2 "XL/2": { # DiT_XL_2
"depth" : 28, "unet_config": {
"num_heads" : 16, "depth" : 28,
"patch_size" : 2, "num_heads" : 16,
"hidden_size" : 1152, "patch_size" : 2,
"hidden_size" : 1152,
},
"sampling_settings" : sampling_settings,
}, },
"XL/4": { # DiT_XL_4 "XL/4": { # DiT_XL_4
"depth" : 28, "unet_config": {
"num_heads" : 16, "depth" : 28,
"patch_size" : 4, "num_heads" : 16,
"hidden_size" : 1152, "patch_size" : 4,
"hidden_size" : 1152,
},
"sampling_settings" : sampling_settings,
}, },
"XL/8": { # DiT_XL_8 "XL/8": { # DiT_XL_8
"depth" : 28, "unet_config": {
"num_heads" : 16, "depth" : 28,
"patch_size" : 8, "num_heads" : 16,
"hidden_size" : 1152, "patch_size" : 8,
"hidden_size" : 1152,
},
"sampling_settings" : sampling_settings,
}, },
"L/2": { # DiT_L_2 "L/2": { # DiT_L_2
"depth" : 24, "unet_config": {
"num_heads" : 16, "depth" : 24,
"patch_size" : 2, "num_heads" : 16,
"hidden_size" : 1024, "patch_size" : 2,
"hidden_size" : 1024,
},
"sampling_settings" : sampling_settings,
}, },
"L/4": { # DiT_L_4 "L/4": { # DiT_L_4
"depth" : 24, "unet_config": {
"num_heads" : 16, "depth" : 24,
"patch_size" : 4, "num_heads" : 16,
"hidden_size" : 1024, "patch_size" : 4,
"hidden_size" : 1024,
},
"sampling_settings" : sampling_settings,
}, },
"L/8": { # DiT_L_8 "L/8": { # DiT_L_8
"depth" : 24, "unet_config": {
"num_heads" : 16, "depth" : 24,
"patch_size" : 8, "num_heads" : 16,
"hidden_size" : 1024, "patch_size" : 8,
"hidden_size" : 1024,
},
"sampling_settings" : sampling_settings,
}, },
"B/2": { # DiT_B_2 "B/2": { # DiT_B_2
"depth" : 12, "unet_config": {
"num_heads" : 12, "depth" : 12,
"patch_size" : 2, "num_heads" : 12,
"hidden_size" : 768, "patch_size" : 2,
"hidden_size" : 768,
},
"sampling_settings" : sampling_settings,
}, },
"B/4": { # DiT_B_4 "B/4": { # DiT_B_4
"depth" : 12, "unet_config": {
"num_heads" : 12, "depth" : 12,
"patch_size" : 4, "num_heads" : 12,
"hidden_size" : 768, "patch_size" : 4,
"hidden_size" : 768,
},
"sampling_settings" : sampling_settings,
}, },
"B/8": { # DiT_B_8 "B/8": { # DiT_B_8
"depth" : 12, "unet_config": {
"num_heads" : 12, "depth" : 12,
"patch_size" : 8, "num_heads" : 12,
"hidden_size" : 768, "patch_size" : 8,
"hidden_size" : 768,
},
"sampling_settings" : sampling_settings,
}, },
"S/2": { # DiT_S_2 "S/2": { # DiT_S_2
"depth" : 12, "unet_config": {
"num_heads" : 6, "depth" : 12,
"patch_size" : 2, "num_heads" : 6,
"hidden_size" : 384, "patch_size" : 2,
"hidden_size" : 384,
},
"sampling_settings" : sampling_settings,
}, },
"S/4": { # DiT_S_4 "S/4": { # DiT_S_4
"depth" : 12, "unet_config": {
"num_heads" : 6, "depth" : 12,
"patch_size" : 4, "num_heads" : 6,
"hidden_size" : 384, "patch_size" : 4,
"hidden_size" : 384,
},
"sampling_settings" : sampling_settings,
}, },
"S/8": { # DiT_S_8 "S/8": { # DiT_S_8
"depth" : 12, "unet_config": {
"num_heads" : 6, "depth" : 12,
"patch_size" : 8, "num_heads" : 6,
"hidden_size" : 384, "patch_size" : 8,
"hidden_size" : 384,
},
"sampling_settings" : sampling_settings,
}, },
} }
+18 -10
View File
@@ -1,5 +1,4 @@
import comfy.supported_models_base import comfy.supported_models_base
import comfy.supported_models
import comfy.latent_formats import comfy.latent_formats
import comfy.model_patcher import comfy.model_patcher
import comfy.model_base import comfy.model_base
@@ -7,11 +6,17 @@ import comfy.utils
import torch import torch
from comfy import model_management from comfy import model_management
from .model import DiT class EXM_DiT(comfy.supported_models_base.BASE):
class EXMDiT(comfy.supported_models.SD15):
unet_config = {} unet_config = {}
unet_extra_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=""): def model_type(self, state_dict, prefix=""):
return comfy.model_base.ModelType.EPS return comfy.model_base.ModelType.EPS
@@ -22,18 +27,21 @@ def load_dit(model_path, model_conf):
parameters = comfy.utils.calculate_parameters(state_dict) parameters = comfy.utils.calculate_parameters(state_dict)
unet_dtype = model_management.unet_dtype(model_params=parameters) 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( model = comfy.model_base.BaseModel(
EXMDiT({"disable_unet_model_creation" : True }), model_conf,
model_type=comfy.model_base.ModelType.EPS, model_type=comfy.model_base.ModelType.EPS,
device=model_management.get_torch_device() 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 from .model import DiT
model.diffusion_model = DiT(**model_conf).eval() model.diffusion_model = DiT(**model_conf.unet_config)
model.diffusion_model.load_state_dict(state_dict) model.diffusion_model.load_state_dict(state_dict)
model.diffusion_model.eval()
model.diffusion_model.dtype = unet_dtype model.diffusion_model.dtype = unet_dtype
model.diffusion_model.eval()
model.diffusion_model.to(unet_dtype) model.diffusion_model.to(unet_dtype)
model_patcher = comfy.model_patcher.ModelPatcher( model_patcher = comfy.model_patcher.ModelPatcher(
+1
View File
@@ -158,6 +158,7 @@ class DiT(nn.Module):
class_dropout_prob=0.1, class_dropout_prob=0.1,
num_classes=1000, num_classes=1000,
learn_sigma=True, learn_sigma=True,
**kwargs,
): ):
super().__init__() super().__init__()
self.learn_sigma = learn_sigma self.learn_sigma = learn_sigma
+5 -16
View File
@@ -6,23 +6,12 @@ import folder_paths
from .conf import dit_conf from .conf import dit_conf
from .loader import load_dit 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: class DitCheckpointLoader:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
return { return {
"required": { "required": {
"ckpt_name": (folder_paths.get_filename_list("dit"),), "ckpt_name": (folder_paths.get_filename_list("checkpoints"),),
"model": (list(dit_conf.keys()),), "model": (list(dit_conf.keys()),),
"image_size": ([256, 512],), "image_size": ([256, 512],),
# "num_classes": ("INT", {"default": 1000, "min": 0,}), # "num_classes": ("INT", {"default": 1000, "min": 0,}),
@@ -35,10 +24,10 @@ class DitCheckpointLoader:
TITLE = "DitCheckpointLoader" TITLE = "DitCheckpointLoader"
def load_checkpoint(self, ckpt_name, model, image_size): 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 = dit_conf[model]
model_conf["input_size"] = image_size // 8 model_conf["unet_config"]["input_size"] = image_size // 8
# model_conf["num_classes"] = num_classes # model_conf["unet_config"]["num_classes"] = num_classes
dit = load_dit( dit = load_dit(
model_path = ckpt_path, model_path = ckpt_path,
model_conf = model_conf, model_conf = model_conf,
@@ -98,7 +87,7 @@ class DiTCondLabelEmpty:
def cond_empty(self, model): def cond_empty(self, model):
# [ID of last class + 1] == [num_classes] # [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) y = torch.tensor([[y_null]]).to(torch.int)
return ([[y, {}]], ) return ([[y, {}]], )
+2 -2
View File
@@ -85,13 +85,13 @@ Limitations:
### Usage ### 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). 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 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!** 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. 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) ![DIT_WORKFLOW_IMG](https://github.com/city96/ComfyUI_ExtraModels/assets/125218114/cdd4ec94-b0eb-436a-bf23-a3bcef8d7b90)