diff --git a/PixArt/loader.py b/PixArt/loader.py index 63294a5..0659bca 100644 --- a/PixArt/loader.py +++ b/PixArt/loader.py @@ -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() @@ -115,3 +120,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, + } + } diff --git a/PixArt/models/PixArt.py b/PixArt/models/PixArt.py index 7850833..4d6cf93 100644 --- a/PixArt/models/PixArt.py +++ b/PixArt/models/PixArt.py @@ -72,6 +72,7 @@ class PixArt(nn.Module): drop_path: float = 0., caption_channels=4096, pe_interpolation=1.0, + pe_precision=None, config=None, model_max_length=120, qk_norm=False, @@ -85,6 +86,7 @@ class PixArt(nn.Module): self.patch_size = patch_size self.num_heads = num_heads self.pe_interpolation = pe_interpolation + self.pe_precision = pe_precision self.depth = depth self.x_embedder = PatchEmbed(input_size, patch_size, in_channels, hidden_size, bias=True) diff --git a/PixArt/models/PixArtMS.py b/PixArt/models/PixArtMS.py index 957811e..908589c 100644 --- a/PixArt/models/PixArtMS.py +++ b/PixArt/models/PixArtMS.py @@ -98,7 +98,8 @@ class PixArtMS(PixArt): pred_sigma=True, drop_path: float = 0., caption_channels=4096, - pe_interpolation=1., + pe_interpolation=None, + pe_precision=None, config=None, model_max_length=120, micro_condition=True, @@ -168,10 +169,16 @@ class PixArtMS(PixArt): x = x.to(self.dtype) timestep = t.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 pos_embed = torch.from_numpy( 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 ) ).unsqueeze(0).to(device=x.device, dtype=self.dtype) diff --git a/PixArt/nodes.py b/PixArt/nodes.py index 87ff74f..ba4d0ee 100644 --- a/PixArt/nodes.py +++ b/PixArt/nodes.py @@ -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,