From 06a936813f37001c519ecbcfa5f089a44a4bf8ca Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Wed, 11 Dec 2024 17:32:57 +0100 Subject: [PATCH] Consolidate loader logic --- PixArt/config.py | 134 ---------------------------------------- PixArt/loader.py | 155 ++++++++++++++++++++++++++++++++++++++--------- utils/loader.py | 35 +++++++++++ 3 files changed, 162 insertions(+), 162 deletions(-) delete mode 100644 PixArt/config.py create mode 100644 utils/loader.py diff --git a/PixArt/config.py b/PixArt/config.py deleted file mode 100644 index 43082a1..0000000 --- a/PixArt/config.py +++ /dev/null @@ -1,134 +0,0 @@ -import math -import logging - -import comfy.supported_models_base -import comfy.supported_models -import comfy.latent_formats -import comfy.model_base - -from .model.pixart import PixArt -from .model.pixartms import PixArtMS -from ..text_encoders.pixart.tenc import PixArtTokenizer, PixArtT5XXL - -class PixArtConfig(comfy.supported_models_base.BASE): - unet_class = PixArtMS - unet_config = {} - unet_extra_config = {} - - latent_format = comfy.latent_formats.SD15 - sampling_settings = { - "beta_schedule" : "sqrt_linear", - "linear_start" : 0.0001, - "linear_end" : 0.02, - "timesteps" : 1000, - } - - def model_type(self, state_dict, prefix=""): - return comfy.model_base.ModelType.EPS - - def get_model(self, state_dict, prefix="", device=None): - return PixArtModel(model_config=self, unet_model=self.unet_class, device=device) - - def clip_target(self, state_dict={}): - return comfy.supported_models_base.ClipTarget(PixArtTokenizer, PixArtT5XXL) - -class PixArtModel(comfy.model_base.BaseModel): - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) - - def extra_conds(self, **kwargs): - out = super().extra_conds(**kwargs) - return out - -def model_config_from_unet(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 - - model_config = PixArtModel - if config["model_max_length"] == 300: - # Sigma - model_class = PixArtMS - model_config.latent_format = comfy.latent_formats.SDXL - 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? - logging.warn(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) - model_class = PixArtMS - config["micro_condition"] = True - if "input_size" not in config: - config["input_size"] = 1024//8 - config["pe_interpolation"] = 2 - else: - # PixArt - model_class = PixArt - if "input_size" not in config: - config["input_size"] = 512//8 - config["pe_interpolation"] = 1 - model_config = PixArtConfig(config) - model_config.unet_class = model_class - logging.debug(f"PixArt config: {model_class}\n{config}") - return model_config - -resolutions = { - "PixArt 512": { - 0.25: [256,1024], 0.26: [256, 992], 0.27: [256, 960], 0.28: [256, 928], - 0.32: [288, 896], 0.33: [288, 864], 0.35: [288, 832], 0.40: [320, 800], - 0.42: [320, 768], 0.48: [352, 736], 0.50: [352, 704], 0.52: [352, 672], - 0.57: [384, 672], 0.60: [384, 640], 0.68: [416, 608], 0.72: [416, 576], - 0.78: [448, 576], 0.82: [448, 544], 0.88: [480, 544], 0.94: [480, 512], - 1.00: [512, 512], 1.07: [512, 480], 1.13: [544, 480], 1.21: [544, 448], - 1.29: [576, 448], 1.38: [576, 416], 1.46: [608, 416], 1.67: [640, 384], - 1.75: [672, 384], 2.00: [704, 352], 2.09: [736, 352], 2.40: [768, 320], - 2.50: [800, 320], 2.89: [832, 288], 3.00: [864, 288], 3.11: [896, 288], - 3.62: [928, 256], 3.75: [960, 256], 3.88: [992, 256], 4.00: [1024,256] - }, - "PixArt 1024": { - 0.25: [512, 2048], 0.26: [512, 1984], 0.27: [512, 1920], 0.28: [512, 1856], - 0.32: [576, 1792], 0.33: [576, 1728], 0.35: [576, 1664], 0.40: [640, 1600], - 0.42: [640, 1536], 0.48: [704, 1472], 0.50: [704, 1408], 0.52: [704, 1344], - 0.57: [768, 1344], 0.60: [768, 1280], 0.68: [832, 1216], 0.72: [832, 1152], - 0.78: [896, 1152], 0.82: [896, 1088], 0.88: [960, 1088], 0.94: [960, 1024], - 1.00: [1024,1024], 1.07: [1024, 960], 1.13: [1088, 960], 1.21: [1088, 896], - 1.29: [1152, 896], 1.38: [1152, 832], 1.46: [1216, 832], 1.67: [1280, 768], - 1.75: [1344, 768], 2.00: [1408, 704], 2.09: [1472, 704], 2.40: [1536, 640], - 2.50: [1600, 640], 2.89: [1664, 576], 3.00: [1728, 576], 3.11: [1792, 576], - 3.62: [1856, 512], 3.75: [1920, 512], 3.88: [1984, 512], 4.00: [2048, 512], - }, - "PixArt 2K": { - 0.25: [1024, 4096], 0.26: [1024, 3968], 0.27: [1024, 3840], 0.28: [1024, 3712], - 0.32: [1152, 3584], 0.33: [1152, 3456], 0.35: [1152, 3328], 0.40: [1280, 3200], - 0.42: [1280, 3072], 0.48: [1408, 2944], 0.50: [1408, 2816], 0.52: [1408, 2688], - 0.57: [1536, 2688], 0.60: [1536, 2560], 0.68: [1664, 2432], 0.72: [1664, 2304], - 0.78: [1792, 2304], 0.82: [1792, 2176], 0.88: [1920, 2176], 0.94: [1920, 2048], - 1.00: [2048, 2048], 1.07: [2048, 1920], 1.13: [2176, 1920], 1.21: [2176, 1792], - 1.29: [2304, 1792], 1.38: [2304, 1664], 1.46: [2432, 1664], 1.67: [2560, 1536], - 1.75: [2688, 1536], 2.00: [2816, 1408], 2.09: [2944, 1408], 2.40: [3072, 1280], - 2.50: [3200, 1280], 2.89: [3328, 1152], 3.00: [3456, 1152], 3.11: [3584, 1152], - 3.62: [3712, 1024], 3.75: [3840, 1024], 3.88: [3968, 1024], 4.00: [4096, 1024] - } -} diff --git a/PixArt/loader.py b/PixArt/loader.py index 9692a71..bc0ca00 100644 --- a/PixArt/loader.py +++ b/PixArt/loader.py @@ -1,13 +1,49 @@ +import math import logging import comfy.utils import comfy.model_base -import comfy.model_patcher import comfy.model_detection -from comfy import model_management +import comfy.supported_models_base +import comfy.supported_models +import comfy.latent_formats + +from .model.pixart import PixArt +from .model.pixartms import PixArtMS from .diffusers_convert import convert_state_dict -from .config import model_config_from_unet +from ..utils.loader import load_state_dict_from_config +from ..text_encoders.pixart.tenc import PixArtTokenizer, PixArtT5XXL + +class PixArtConfig(comfy.supported_models_base.BASE): + unet_class = PixArtMS + unet_config = {} + unet_extra_config = {} + + latent_format = comfy.latent_formats.SD15 + sampling_settings = { + "beta_schedule" : "sqrt_linear", + "linear_start" : 0.0001, + "linear_end" : 0.02, + "timesteps" : 1000, + } + + def model_type(self, state_dict, prefix=""): + return comfy.model_base.ModelType.EPS + + def get_model(self, state_dict, prefix="", device=None): + return PixArtModel(model_config=self, unet_model=self.unet_class, device=device) + + def clip_target(self, state_dict={}): + return comfy.supported_models_base.ClipTarget(PixArtTokenizer, PixArtT5XXL) + +class PixArtModel(comfy.model_base.BaseModel): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + + def extra_conds(self, **kwargs): + out = super().extra_conds(**kwargs) + return out def load_pixart_state_dict(sd, model_options={}): # prefix / format @@ -23,34 +59,97 @@ def load_pixart_state_dict(sd, model_options={}): # model config model_config = model_config_from_unet(sd) + return load_state_dict_from_config(model_config, sd, model_options) - # TODO: move lines below to utils - parameters = comfy.utils.calculate_parameters(sd) - load_device = model_management.get_torch_device() - offload_device = comfy.model_management.unet_offload_device() +def model_config_from_unet(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 - dtype = model_options.get("dtype", None) - weight_dtype = comfy.utils.weight_dtype(sd) - unet_weight_dtype = list(model_config.supported_inference_dtypes) + 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 weight_dtype is not None and model_config.scaled_fp8 is None: - unet_weight_dtype.append(weight_dtype) + 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 - if dtype is None: - unet_dtype = model_management.unet_dtype(model_params=parameters, supported_dtypes=unet_weight_dtype) + model_config = PixArtModel + if config["model_max_length"] == 300: + # Sigma + model_class = PixArtMS + model_config.latent_format = comfy.latent_formats.SDXL + 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? + logging.warn(f"PixArt: diffusers weights - 2K model will be broken, use manual loading!") + config["input_size"] = 1024//8 else: - unet_dtype = dtype + # Alpha + if "csize_embedder.mlp.0.weight" in sd: + # MS (microconds) + model_class = PixArtMS + config["micro_condition"] = True + if "input_size" not in config: + config["input_size"] = 1024//8 + config["pe_interpolation"] = 2 + else: + # PixArt + model_class = PixArt + if "input_size" not in config: + config["input_size"] = 512//8 + config["pe_interpolation"] = 1 + model_config = PixArtConfig(config) + model_config.unet_class = model_class + logging.debug(f"PixArt config: {model_class}\n{config}") + return model_config - manual_cast_dtype = model_management.unet_manual_cast(unet_dtype, load_device, model_config.supported_inference_dtypes) - model_config.set_inference_dtype(unet_dtype, manual_cast_dtype) - model_config.custom_operations = model_options.get("custom_operations", model_config.custom_operations) - if model_options.get("fp8_optimizations", False): - model_config.optimizations["fp8"] = True - - model = model_config.get_model(sd, "") - model = model.to(offload_device).eval() - model.load_model_weights(sd, "") - left_over = sd.keys() - if len(left_over) > 0: - logging.info("left over keys in unet: {}".format(left_over)) - return comfy.model_patcher.ModelPatcher(model, load_device=load_device, offload_device=offload_device) +resolutions = { + "PixArt 512": { + 0.25: [256,1024], 0.26: [256, 992], 0.27: [256, 960], 0.28: [256, 928], + 0.32: [288, 896], 0.33: [288, 864], 0.35: [288, 832], 0.40: [320, 800], + 0.42: [320, 768], 0.48: [352, 736], 0.50: [352, 704], 0.52: [352, 672], + 0.57: [384, 672], 0.60: [384, 640], 0.68: [416, 608], 0.72: [416, 576], + 0.78: [448, 576], 0.82: [448, 544], 0.88: [480, 544], 0.94: [480, 512], + 1.00: [512, 512], 1.07: [512, 480], 1.13: [544, 480], 1.21: [544, 448], + 1.29: [576, 448], 1.38: [576, 416], 1.46: [608, 416], 1.67: [640, 384], + 1.75: [672, 384], 2.00: [704, 352], 2.09: [736, 352], 2.40: [768, 320], + 2.50: [800, 320], 2.89: [832, 288], 3.00: [864, 288], 3.11: [896, 288], + 3.62: [928, 256], 3.75: [960, 256], 3.88: [992, 256], 4.00: [1024,256] + }, + "PixArt 1024": { + 0.25: [512, 2048], 0.26: [512, 1984], 0.27: [512, 1920], 0.28: [512, 1856], + 0.32: [576, 1792], 0.33: [576, 1728], 0.35: [576, 1664], 0.40: [640, 1600], + 0.42: [640, 1536], 0.48: [704, 1472], 0.50: [704, 1408], 0.52: [704, 1344], + 0.57: [768, 1344], 0.60: [768, 1280], 0.68: [832, 1216], 0.72: [832, 1152], + 0.78: [896, 1152], 0.82: [896, 1088], 0.88: [960, 1088], 0.94: [960, 1024], + 1.00: [1024,1024], 1.07: [1024, 960], 1.13: [1088, 960], 1.21: [1088, 896], + 1.29: [1152, 896], 1.38: [1152, 832], 1.46: [1216, 832], 1.67: [1280, 768], + 1.75: [1344, 768], 2.00: [1408, 704], 2.09: [1472, 704], 2.40: [1536, 640], + 2.50: [1600, 640], 2.89: [1664, 576], 3.00: [1728, 576], 3.11: [1792, 576], + 3.62: [1856, 512], 3.75: [1920, 512], 3.88: [1984, 512], 4.00: [2048, 512], + }, + "PixArt 2K": { + 0.25: [1024, 4096], 0.26: [1024, 3968], 0.27: [1024, 3840], 0.28: [1024, 3712], + 0.32: [1152, 3584], 0.33: [1152, 3456], 0.35: [1152, 3328], 0.40: [1280, 3200], + 0.42: [1280, 3072], 0.48: [1408, 2944], 0.50: [1408, 2816], 0.52: [1408, 2688], + 0.57: [1536, 2688], 0.60: [1536, 2560], 0.68: [1664, 2432], 0.72: [1664, 2304], + 0.78: [1792, 2304], 0.82: [1792, 2176], 0.88: [1920, 2176], 0.94: [1920, 2048], + 1.00: [2048, 2048], 1.07: [2048, 1920], 1.13: [2176, 1920], 1.21: [2176, 1792], + 1.29: [2304, 1792], 1.38: [2304, 1664], 1.46: [2432, 1664], 1.67: [2560, 1536], + 1.75: [2688, 1536], 2.00: [2816, 1408], 2.09: [2944, 1408], 2.40: [3072, 1280], + 2.50: [3200, 1280], 2.89: [3328, 1152], 3.00: [3456, 1152], 3.11: [3584, 1152], + 3.62: [3712, 1024], 3.75: [3840, 1024], 3.88: [3968, 1024], 4.00: [4096, 1024] + } +} diff --git a/utils/loader.py b/utils/loader.py new file mode 100644 index 0000000..f3f1b77 --- /dev/null +++ b/utils/loader.py @@ -0,0 +1,35 @@ +import logging +import comfy.utils +import comfy.model_patcher +from comfy import model_management + +def load_state_dict_from_config(model_config, sd, model_options={}): + parameters = comfy.utils.calculate_parameters(sd) + load_device = model_management.get_torch_device() + offload_device = comfy.model_management.unet_offload_device() + + dtype = model_options.get("dtype", None) + weight_dtype = comfy.utils.weight_dtype(sd) + unet_weight_dtype = list(model_config.supported_inference_dtypes) + + if weight_dtype is not None and model_config.scaled_fp8 is None: + unet_weight_dtype.append(weight_dtype) + + if dtype is None: + unet_dtype = model_management.unet_dtype(model_params=parameters, supported_dtypes=unet_weight_dtype) + else: + unet_dtype = dtype + + manual_cast_dtype = model_management.unet_manual_cast(unet_dtype, load_device, model_config.supported_inference_dtypes) + model_config.set_inference_dtype(unet_dtype, manual_cast_dtype) + model_config.custom_operations = model_options.get("custom_operations", model_config.custom_operations) + if model_options.get("fp8_optimizations", False): + model_config.optimizations["fp8"] = True + + model = model_config.get_model(sd, "") + model = model.to(offload_device).eval() + model.load_model_weights(sd, "") + left_over = sd.keys() + if len(left_over) > 0: + logging.info("left over keys in unet: {}".format(left_over)) + return comfy.model_patcher.ModelPatcher(model, load_device=load_device, offload_device=offload_device)