Merge branch 'city96:main' into get-depth-programatically

This commit is contained in:
codeMODE
2024-06-23 19:43:25 +01:00
committed by GitHub
4 changed files with 94 additions and 3 deletions
+67 -1
View File
@@ -5,6 +5,7 @@ import comfy.model_base
import comfy.utils import comfy.utils
import comfy.conds import comfy.conds
import torch import torch
import math
from comfy import model_management from comfy import model_management
from .diffusers_convert import convert_state_dict from .diffusers_convert import convert_state_dict
@@ -45,7 +46,7 @@ class EXM_PixArt_Model(comfy.model_base.BaseModel):
return out 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 = comfy.utils.load_torch_file(model_path)
state_dict = state_dict.get("model", state_dict) 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: if "adaln_single.linear.weight" in state_dict:
state_dict = convert_state_dict(state_dict) # Diffusers 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) 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)
load_device = comfy.model_management.get_torch_device() load_device = comfy.model_management.get_torch_device()
@@ -115,3 +120,64 @@ def load_pixart(model_path, model_conf):
current_device = "cpu", current_device = "cpu",
) )
return model_patcher 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,
}
}
+2
View File
@@ -72,6 +72,7 @@ class PixArt(nn.Module):
drop_path: float = 0., drop_path: float = 0.,
caption_channels=4096, caption_channels=4096,
pe_interpolation=1.0, pe_interpolation=1.0,
pe_precision=None,
config=None, config=None,
model_max_length=120, model_max_length=120,
qk_norm=False, qk_norm=False,
@@ -85,6 +86,7 @@ class PixArt(nn.Module):
self.patch_size = patch_size self.patch_size = patch_size
self.num_heads = num_heads self.num_heads = num_heads
self.pe_interpolation = pe_interpolation self.pe_interpolation = pe_interpolation
self.pe_precision = pe_precision
self.depth = depth self.depth = depth
self.x_embedder = PatchEmbed(input_size, patch_size, in_channels, hidden_size, bias=True) self.x_embedder = PatchEmbed(input_size, patch_size, in_channels, hidden_size, bias=True)
+9 -2
View File
@@ -98,7 +98,8 @@ class PixArtMS(PixArt):
pred_sigma=True, pred_sigma=True,
drop_path: float = 0., drop_path: float = 0.,
caption_channels=4096, caption_channels=4096,
pe_interpolation=1., pe_interpolation=None,
pe_precision=None,
config=None, config=None,
model_max_length=120, model_max_length=120,
micro_condition=True, micro_condition=True,
@@ -168,10 +169,16 @@ class PixArtMS(PixArt):
x = x.to(self.dtype) x = x.to(self.dtype)
timestep = t.to(self.dtype) timestep = t.to(self.dtype)
y = y.to(self.dtype) y = y.to(self.dtype)
pe_interpolation = self.pe_interpolation
if pe_interpolation is None or self.pe_precision is not None:
# calculate pe_interpolation on-the-fly
pe_interpolation = round((x.shape[-1]+x.shape[-2])/2.0 / (512/8.0), self.pe_precision or 0)
self.h, self.w = x.shape[-2]//self.patch_size, x.shape[-1]//self.patch_size self.h, self.w = x.shape[-2]//self.patch_size, x.shape[-1]//self.patch_size
pos_embed = torch.from_numpy( pos_embed = torch.from_numpy(
get_2d_sincos_pos_embed( get_2d_sincos_pos_embed(
self.pos_embed.shape[-1], (self.h, self.w), pe_interpolation=self.pe_interpolation, self.pos_embed.shape[-1], (self.h, self.w), pe_interpolation=pe_interpolation,
base_size=self.base_size base_size=self.base_size
) )
).unsqueeze(0).to(device=x.device, dtype=self.dtype) ).unsqueeze(0).to(device=x.device, dtype=self.dtype)
+16
View File
@@ -32,6 +32,21 @@ class PixArtCheckpointLoader:
) )
return (model,) 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(): class PixArtResolutionSelect():
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
@@ -245,6 +260,7 @@ class PixArtT5FromSD3CLIP:
NODE_CLASS_MAPPINGS = { NODE_CLASS_MAPPINGS = {
"PixArtCheckpointLoader" : PixArtCheckpointLoader, "PixArtCheckpointLoader" : PixArtCheckpointLoader,
"PixArtCheckpointLoaderSimple" : PixArtCheckpointLoaderSimple,
"PixArtResolutionSelect" : PixArtResolutionSelect, "PixArtResolutionSelect" : PixArtResolutionSelect,
"PixArtLoraLoader" : PixArtLoraLoader, "PixArtLoraLoader" : PixArtLoraLoader,
"PixArtT5TextEncode" : PixArtT5TextEncode, "PixArtT5TextEncode" : PixArtT5TextEncode,