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:
+91
-48
@@ -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
@@ -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(
|
||||||
|
|||||||
@@ -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
@@ -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, {}]], )
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||

|

|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user