Merge branch 'city96:main' into get-depth-programatically
This commit is contained in:
+67
-1
@@ -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,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
Reference in New Issue
Block a user