PixArt auto checkpoint loader
Needs better logic for diffusers weights
This commit is contained in:
+67
-1
@@ -5,6 +5,7 @@ import comfy.model_base
|
||||
import comfy.utils
|
||||
import comfy.conds
|
||||
import torch
|
||||
import math
|
||||
from comfy import model_management
|
||||
from .diffusers_convert import convert_state_dict
|
||||
|
||||
@@ -45,7 +46,7 @@ class EXM_PixArt_Model(comfy.model_base.BaseModel):
|
||||
|
||||
return out
|
||||
|
||||
def load_pixart(model_path, model_conf):
|
||||
def load_pixart(model_path, model_conf=None):
|
||||
state_dict = comfy.utils.load_torch_file(model_path)
|
||||
state_dict = state_dict.get("model", state_dict)
|
||||
|
||||
@@ -58,6 +59,10 @@ def load_pixart(model_path, model_conf):
|
||||
if "adaln_single.linear.weight" in state_dict:
|
||||
state_dict = convert_state_dict(state_dict) # Diffusers
|
||||
|
||||
# guess auto config
|
||||
if model_conf is None:
|
||||
model_conf = guess_pixart_config(state_dict)
|
||||
|
||||
parameters = comfy.utils.calculate_parameters(state_dict)
|
||||
unet_dtype = model_management.unet_dtype(model_params=parameters)
|
||||
load_device = comfy.model_management.get_torch_device()
|
||||
@@ -113,3 +118,64 @@ def load_pixart(model_path, model_conf):
|
||||
current_device = "cpu",
|
||||
)
|
||||
return model_patcher
|
||||
|
||||
def guess_pixart_config(sd):
|
||||
"""
|
||||
Guess config based on converted state dict.
|
||||
"""
|
||||
# Shared settings based on DiT_XL_2 - could be enumerated
|
||||
config = {
|
||||
"num_heads" : 16, # get from attention
|
||||
"patch_size" : 2, # final layer I guess?
|
||||
"hidden_size" : 1152, # pos_embed.shape[2]
|
||||
}
|
||||
config["depth"] = sum([key.endswith(".attn.proj.weight") for key in sd.keys()]) or 28
|
||||
|
||||
try:
|
||||
# this is not present in the diffusers version for sigma?
|
||||
config["model_max_length"] = sd["y_embedder.y_embedding"].shape[0]
|
||||
except KeyError:
|
||||
# need better logic to guess this
|
||||
config["model_max_length"] = 300
|
||||
|
||||
if "pos_embed" in sd:
|
||||
config["input_size"] = int(math.sqrt(sd["pos_embed"].shape[1])) * config["patch_size"]
|
||||
config["pe_interpolation"] = config["input_size"] // (512//8) # dumb guess
|
||||
|
||||
target_arch = "PixArtMS"
|
||||
if config["model_max_length"] == 300:
|
||||
# Sigma
|
||||
target_arch = "PixArtMSSigma"
|
||||
config["micro_condition"] = False
|
||||
if "input_size" not in config:
|
||||
# The diffusers weights for 1K/2K are exactly the same...?
|
||||
# replace patch embed logic with HyDiT?
|
||||
print(f"PixArt: diffusers weights - 2K model will be broken, use manual loading!")
|
||||
config["input_size"] = 1024//8
|
||||
else:
|
||||
# Alpha
|
||||
if "csize_embedder.mlp.0.weight" in sd:
|
||||
# MS (microconds)
|
||||
target_arch = "PixArtMS"
|
||||
config["micro_condition"] = True
|
||||
if "input_size" not in config:
|
||||
config["input_size"] = 1024//8
|
||||
config["pe_interpolation"] = 2
|
||||
else:
|
||||
# PixArt
|
||||
target_arch = "PixArt"
|
||||
if "input_size" not in config:
|
||||
config["input_size"] = 512//8
|
||||
config["pe_interpolation"] = 1
|
||||
|
||||
print("PixArt guessed config:", target_arch, config)
|
||||
return {
|
||||
"target": target_arch,
|
||||
"unet_config": config,
|
||||
"sampling_settings": {
|
||||
"beta_schedule" : "sqrt_linear",
|
||||
"linear_start" : 0.0001,
|
||||
"linear_end" : 0.02,
|
||||
"timesteps" : 1000,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -32,6 +32,21 @@ class PixArtCheckpointLoader:
|
||||
)
|
||||
return (model,)
|
||||
|
||||
class PixArtCheckpointLoaderSimple(PixArtCheckpointLoader):
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"ckpt_name": (folder_paths.get_filename_list("checkpoints"),),
|
||||
}
|
||||
}
|
||||
TITLE = "PixArt Checkpoint Loader (auto)"
|
||||
|
||||
def load_checkpoint(self, ckpt_name):
|
||||
ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name)
|
||||
model = load_pixart(model_path=ckpt_path)
|
||||
return (model,)
|
||||
|
||||
class PixArtResolutionSelect():
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -245,6 +260,7 @@ class PixArtT5FromSD3CLIP:
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"PixArtCheckpointLoader" : PixArtCheckpointLoader,
|
||||
"PixArtCheckpointLoaderSimple" : PixArtCheckpointLoaderSimple,
|
||||
"PixArtResolutionSelect" : PixArtResolutionSelect,
|
||||
"PixArtLoraLoader" : PixArtLoraLoader,
|
||||
"PixArtT5TextEncode" : PixArtT5TextEncode,
|
||||
|
||||
Reference in New Issue
Block a user