Sana loader logic
This commit is contained in:
@@ -1,98 +0,0 @@
|
||||
"""
|
||||
List of all Sana model types / settings
|
||||
"""
|
||||
|
||||
sampling_settings = {
|
||||
"shift": 3.0,
|
||||
}
|
||||
|
||||
sana_conf = {
|
||||
"SanaMS_600M_P1_D28": {
|
||||
"target": "SanaMS",
|
||||
"unet_config": {
|
||||
"in_channels": 32,
|
||||
"depth": 28,
|
||||
"hidden_size": 1152,
|
||||
"patch_size": 1,
|
||||
"num_heads": 16,
|
||||
"linear_head_dim": 32,
|
||||
"model_max_length": 300,
|
||||
"y_norm": True,
|
||||
"attn_type": "linear",
|
||||
"ffn_type": "glumbconv",
|
||||
"mlp_ratio": 2.5,
|
||||
"mlp_acts": ["silu", "silu", None],
|
||||
"use_pe": False,
|
||||
"pred_sigma": False,
|
||||
"learn_sigma": False,
|
||||
"fp32_attention": True,
|
||||
},
|
||||
"sampling_settings" : sampling_settings,
|
||||
},
|
||||
"SanaMS_1600M_P1_D20": {
|
||||
"target": "SanaMS",
|
||||
"unet_config": {
|
||||
"in_channels": 32,
|
||||
"depth": 20,
|
||||
"hidden_size": 2240,
|
||||
"patch_size": 1,
|
||||
"num_heads": 20,
|
||||
"linear_head_dim": 32,
|
||||
"model_max_length": 300,
|
||||
"y_norm": True,
|
||||
"attn_type": "linear",
|
||||
"ffn_type": "glumbconv",
|
||||
"mlp_ratio": 2.5,
|
||||
"mlp_acts": ["silu", "silu", None],
|
||||
"use_pe": False,
|
||||
"pred_sigma": False,
|
||||
"learn_sigma": False,
|
||||
"fp32_attention": True,
|
||||
},
|
||||
"sampling_settings" : sampling_settings,
|
||||
},
|
||||
}
|
||||
|
||||
sana_res = {
|
||||
"1024px": { # models/SanaMS 1024x1024
|
||||
'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],
|
||||
},
|
||||
"512px": { # models/SanaMS 512x512
|
||||
'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]
|
||||
},
|
||||
"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]
|
||||
}
|
||||
}
|
||||
# These should be the same
|
||||
sana_res.update({
|
||||
"SanaMS_600M_P1_D28": sana_res["1024px"],
|
||||
"SanaMS_1600M_P1_D20": sana_res["1024px"],
|
||||
})
|
||||
+80
-76
@@ -1,100 +1,104 @@
|
||||
import comfy.supported_models_base
|
||||
import comfy.latent_formats
|
||||
import comfy.model_patcher
|
||||
import comfy.model_base
|
||||
import logging
|
||||
|
||||
import comfy.utils
|
||||
import comfy.conds
|
||||
import torch
|
||||
import math
|
||||
from comfy import model_management
|
||||
from comfy.latent_formats import LatentFormat
|
||||
import comfy.model_base
|
||||
import comfy.model_detection
|
||||
|
||||
import comfy.supported_models_base
|
||||
import comfy.supported_models
|
||||
import comfy.latent_formats
|
||||
|
||||
from .models.sana import Sana
|
||||
from .models.sana_multi_scale import SanaMS
|
||||
from .diffusers_convert import convert_state_dict
|
||||
from ..utils.loader import load_state_dict_from_config
|
||||
|
||||
|
||||
class SanaLatent(LatentFormat):
|
||||
class SanaLatent(comfy.latent_formats.LatentFormat):
|
||||
scale_factor = 0.41407
|
||||
latent_channels = 32
|
||||
def __init__(self):
|
||||
self.scale_factor = 0.41407
|
||||
|
||||
|
||||
class EXM_Sana(comfy.supported_models_base.BASE):
|
||||
class SanaConfig(comfy.supported_models_base.BASE):
|
||||
unet_class = SanaMS
|
||||
unet_config = {}
|
||||
unet_extra_config = {}
|
||||
latent_format = SanaLatent
|
||||
|
||||
def __init__(self, model_conf):
|
||||
self.model_target = model_conf.get("target")
|
||||
self.unet_config = model_conf.get("unet_config", {})
|
||||
self.sampling_settings = model_conf.get("sampling_settings", {})
|
||||
self.latent_format = self.latent_format()
|
||||
# UNET is handled by extension
|
||||
self.unet_config["disable_unet_model_creation"] = True
|
||||
latent_format = SanaLatent
|
||||
sampling_settings = {
|
||||
"shift": 3.0,
|
||||
}
|
||||
|
||||
def model_type(self, state_dict, prefix=""):
|
||||
return comfy.model_base.ModelType.FLOW
|
||||
|
||||
def get_model(self, state_dict, prefix="", device=None):
|
||||
return SanaModel(
|
||||
model_config=self,
|
||||
model_type=comfy.model_base.ModelType.FLOW,
|
||||
unet_model=self.unet_class,
|
||||
device=device
|
||||
)
|
||||
|
||||
class EXM_Sana_Model(comfy.model_base.BaseModel):
|
||||
class SanaModel(comfy.model_base.BaseModel):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
def extra_conds(self, **kwargs):
|
||||
out = super().extra_conds(**kwargs)
|
||||
def load_sana_state_dict(sd, model_options={}):
|
||||
# prefix / format
|
||||
sd = sd.get("model", sd) # ref ckpt
|
||||
diffusion_model_prefix = comfy.model_detection.unet_prefix_from_state_dict(sd)
|
||||
temp_sd = comfy.utils.state_dict_prefix_replace(sd, {diffusion_model_prefix: ""}, filter_keys=True)
|
||||
if len(temp_sd) > 0:
|
||||
sd = temp_sd
|
||||
|
||||
cn_hint = kwargs.get("cn_hint", None)
|
||||
if cn_hint is not None:
|
||||
out["cn_hint"] = comfy.conds.CONDRegular(cn_hint)
|
||||
# diffusers convert
|
||||
if "adaln_single.linear.weight" in sd:
|
||||
sd = convert_state_dict(sd)
|
||||
|
||||
return out
|
||||
# model config
|
||||
model_config = model_config_from_unet(sd)
|
||||
return load_state_dict_from_config(model_config, sd, model_options)
|
||||
|
||||
def model_config_from_unet(sd):
|
||||
"""
|
||||
Guess config based on (converted) state dict.
|
||||
"""
|
||||
# shared settings that match between all models
|
||||
# TODO: some can (should) be enumerated
|
||||
config = {
|
||||
"in_channels": 32,
|
||||
"linear_head_dim": 32,
|
||||
"model_max_length": 300,
|
||||
"y_norm": True,
|
||||
"attn_type": "linear",
|
||||
"ffn_type": "glumbconv",
|
||||
"mlp_ratio": 2.5,
|
||||
"mlp_acts": ["silu", "silu", None],
|
||||
"use_pe": False,
|
||||
"pred_sigma": False,
|
||||
"learn_sigma": False,
|
||||
"fp32_attention": True,
|
||||
"patch_size": 1,
|
||||
}
|
||||
config["depth"] = sum([key.endswith(".point_conv.conv.weight") for key in sd.keys()]) or 28
|
||||
|
||||
def load_sana(model_path, model_conf):
|
||||
state_dict = comfy.utils.load_torch_file(model_path)
|
||||
state_dict = state_dict.get("model", state_dict)
|
||||
if "x_embedder.proj.bias" in sd:
|
||||
config["hidden_size"] = sd["x_embedder.proj.bias"].shape[0]
|
||||
|
||||
# prefix
|
||||
for prefix in ["model.diffusion_model.",]:
|
||||
if any(True for x in state_dict if x.startswith(prefix)):
|
||||
state_dict = {k[len(prefix):]:v for k,v in state_dict.items()}
|
||||
|
||||
# diffusers
|
||||
if "adaln_single.linear.weight" in state_dict:
|
||||
state_dict = convert_state_dict(state_dict) # Diffusers
|
||||
|
||||
parameters = comfy.utils.calculate_parameters(state_dict)
|
||||
unet_dtype = comfy.model_management.unet_dtype()
|
||||
load_device = comfy.model_management.get_torch_device()
|
||||
offload_device = comfy.model_management.unet_offload_device()
|
||||
|
||||
# ignore fp8/etc and use directly for now
|
||||
manual_cast_dtype = model_management.unet_manual_cast(unet_dtype, load_device)
|
||||
if manual_cast_dtype:
|
||||
print(f"Sana: falling back to {manual_cast_dtype}")
|
||||
unet_dtype = manual_cast_dtype
|
||||
|
||||
model_conf = EXM_Sana(model_conf) # convert to object
|
||||
model = EXM_Sana_Model( # same as comfy.model_base.BaseModel
|
||||
model_conf,
|
||||
model_type=comfy.model_base.ModelType.FLOW,
|
||||
device=model_management.get_torch_device()
|
||||
)
|
||||
|
||||
if model_conf.model_target == "SanaMS":
|
||||
from .models.sana_multi_scale import SanaMS
|
||||
model.diffusion_model = SanaMS(**model_conf.unet_config)
|
||||
if config["hidden_size"] == 1152:
|
||||
config["num_heads"] = 16
|
||||
elif config["hidden_size"] == 2240:
|
||||
config["num_heads"] = 20
|
||||
else:
|
||||
raise NotImplementedError(f"Unknown model target '{model_conf.model_target}'")
|
||||
raise RuntimeError(f"Unknown model config.")
|
||||
|
||||
m, u = model.diffusion_model.load_state_dict(state_dict, strict=False)
|
||||
if len(m) > 0: print("Missing UNET keys", m)
|
||||
if len(u) > 0: print("Leftover UNET keys", u)
|
||||
model.diffusion_model.dtype = unet_dtype
|
||||
model.diffusion_model.eval()
|
||||
model.diffusion_model.to(unet_dtype)
|
||||
model_config = SanaConfig(config)
|
||||
logging.debug(f"Sana config:\n{config}")
|
||||
return model_config
|
||||
|
||||
model_patcher = comfy.model_patcher.ModelPatcher(
|
||||
model,
|
||||
load_device = load_device,
|
||||
offload_device = offload_device,
|
||||
)
|
||||
return model_patcher
|
||||
# 512/1024/2K match, TODO: 4K is new, add on release
|
||||
from ..PixArt.loader import resolutions as pixart_res
|
||||
resolutions = {
|
||||
"Sana 512": pixart_res["PixArt 512"],
|
||||
"Sana 1024": pixart_res["PixArt 1024"],
|
||||
"Sana 2K": pixart_res["PixArt 2K"],
|
||||
}
|
||||
|
||||
@@ -2,41 +2,6 @@ import torch
|
||||
import folder_paths
|
||||
from nodes import EmptyLatentImage
|
||||
|
||||
from .conf import sana_conf, sana_res
|
||||
from .loader import load_sana
|
||||
|
||||
dtypes = [
|
||||
"auto",
|
||||
"FP32",
|
||||
"FP16",
|
||||
"BF16"
|
||||
]
|
||||
|
||||
class SanaCheckpointLoader:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"ckpt_name": (folder_paths.get_filename_list("checkpoints"),),
|
||||
"model": (list(sana_conf.keys()),),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
RETURN_NAMES = ("model",)
|
||||
FUNCTION = "load_checkpoint"
|
||||
CATEGORY = "ExtraModels/Sana"
|
||||
TITLE = "Sana Checkpoint Loader"
|
||||
|
||||
def load_checkpoint(self, ckpt_name, model):
|
||||
ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name)
|
||||
model_conf = sana_conf[model]
|
||||
model = load_sana(
|
||||
model_path = ckpt_path,
|
||||
model_conf = model_conf,
|
||||
)
|
||||
return (model,)
|
||||
|
||||
|
||||
class EmptySanaLatentImage(EmptyLatentImage):
|
||||
CATEGORY = "ExtraModels/Sana"
|
||||
TITLE = "Empty Sana Latent Image"
|
||||
@@ -45,53 +10,6 @@ class EmptySanaLatentImage(EmptyLatentImage):
|
||||
latent = torch.zeros([batch_size, 32, height // 32, width // 32], device=self.device)
|
||||
return ({"samples":latent}, )
|
||||
|
||||
|
||||
class SanaResolutionSelect():
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": (list(sana_res.keys()),),
|
||||
"ratio": (list(sana_res["1024px"].keys()),{"default":"1.00"}),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("INT","INT")
|
||||
RETURN_NAMES = ("width","height")
|
||||
FUNCTION = "get_res"
|
||||
CATEGORY = "ExtraModels/Sana"
|
||||
TITLE = "Sana Resolution Select"
|
||||
|
||||
def get_res(self, model, ratio):
|
||||
width, height = sana_res[model][ratio]
|
||||
return (width,height)
|
||||
|
||||
|
||||
class SanaResolutionCond:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"cond": ("CONDITIONING", ),
|
||||
"width": ("INT", {"default": 1024.0, "min": 0, "max": 8192}),
|
||||
"height": ("INT", {"default": 1024.0, "min": 0, "max": 8192}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING",)
|
||||
RETURN_NAMES = ("cond",)
|
||||
FUNCTION = "add_cond"
|
||||
CATEGORY = "ExtraModels/Sana"
|
||||
TITLE = "Sana Resolution Conditioning"
|
||||
|
||||
def add_cond(self, cond, width, height):
|
||||
for c in range(len(cond)):
|
||||
cond[c][1].update({
|
||||
"img_hw": [[height, width]],
|
||||
"aspect_ratio": [[height/width]],
|
||||
})
|
||||
return (cond,)
|
||||
|
||||
|
||||
class SanaTextEncode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -144,9 +62,6 @@ preset_te_prompt = [
|
||||
]
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"SanaCheckpointLoader" : SanaCheckpointLoader,
|
||||
"SanaResolutionSelect" : SanaResolutionSelect,
|
||||
"SanaTextEncode" : SanaTextEncode,
|
||||
"SanaResolutionCond" : SanaResolutionCond,
|
||||
"EmptySanaLatentImage": EmptySanaLatentImage,
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user