Consolidate loader logic

This commit is contained in:
City
2024-12-11 17:32:57 +01:00
parent 6a88140002
commit 06a936813f
3 changed files with 162 additions and 162 deletions
-134
View File
@@ -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]
}
}
+127 -28
View File
@@ -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]
}
}
+35
View File
@@ -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)