Sana loader logic

This commit is contained in:
City
2024-12-11 18:55:17 +01:00
parent 06a936813f
commit 5505ab4f40
4 changed files with 142 additions and 319 deletions
-98
View File
@@ -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"],
})
+89 -85
View File
@@ -1,100 +1,104 @@
import comfy.supported_models_base import logging
import comfy.latent_formats
import comfy.model_patcher
import comfy.model_base
import comfy.utils import comfy.utils
import comfy.conds import comfy.model_base
import torch import comfy.model_detection
import math
from comfy import model_management import comfy.supported_models_base
from comfy.latent_formats import LatentFormat 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 .diffusers_convert import convert_state_dict
from ..utils.loader import load_state_dict_from_config
class SanaLatent(comfy.latent_formats.LatentFormat):
class SanaLatent(LatentFormat): scale_factor = 0.41407
latent_channels = 32 latent_channels = 32
def __init__(self):
self.scale_factor = 0.41407
class SanaConfig(comfy.supported_models_base.BASE):
unet_class = SanaMS
unet_config = {}
unet_extra_config = {}
class EXM_Sana(comfy.supported_models_base.BASE): latent_format = SanaLatent
unet_config = {} sampling_settings = {
unet_extra_config = {} "shift": 3.0,
latent_format = SanaLatent }
def __init__(self, model_conf): def model_type(self, state_dict, prefix=""):
self.model_target = model_conf.get("target") return comfy.model_base.ModelType.FLOW
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
def model_type(self, state_dict, prefix=""): def get_model(self, state_dict, prefix="", device=None):
return comfy.model_base.ModelType.FLOW return SanaModel(
model_config=self,
model_type=comfy.model_base.ModelType.FLOW,
unet_model=self.unet_class,
device=device
)
class SanaModel(comfy.model_base.BaseModel):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
class EXM_Sana_Model(comfy.model_base.BaseModel): def load_sana_state_dict(sd, model_options={}):
def __init__(self, *args, **kwargs): # prefix / format
super().__init__(*args, **kwargs) sd = sd.get("model", sd) # ref ckpt
diffusion_model_prefix = comfy.model_detection.unet_prefix_from_state_dict(sd)
def extra_conds(self, **kwargs): temp_sd = comfy.utils.state_dict_prefix_replace(sd, {diffusion_model_prefix: ""}, filter_keys=True)
out = super().extra_conds(**kwargs) if len(temp_sd) > 0:
sd = temp_sd
cn_hint = kwargs.get("cn_hint", None) # diffusers convert
if cn_hint is not None: if "adaln_single.linear.weight" in sd:
out["cn_hint"] = comfy.conds.CONDRegular(cn_hint) 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): if "x_embedder.proj.bias" in sd:
state_dict = comfy.utils.load_torch_file(model_path) config["hidden_size"] = sd["x_embedder.proj.bias"].shape[0]
state_dict = state_dict.get("model", state_dict)
if config["hidden_size"] == 1152:
config["num_heads"] = 16
elif config["hidden_size"] == 2240:
config["num_heads"] = 20
else:
raise RuntimeError(f"Unknown model config.")
model_config = SanaConfig(config)
logging.debug(f"Sana config:\n{config}")
return model_config
# prefix # 512/1024/2K match, TODO: 4K is new, add on release
for prefix in ["model.diffusion_model.",]: from ..PixArt.loader import resolutions as pixart_res
if any(True for x in state_dict if x.startswith(prefix)): resolutions = {
state_dict = {k[len(prefix):]:v for k,v in state_dict.items()} "Sana 512": pixart_res["PixArt 512"],
"Sana 1024": pixart_res["PixArt 1024"],
# diffusers "Sana 2K": pixart_res["PixArt 2K"],
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)
else:
raise NotImplementedError(f"Unknown model target '{model_conf.model_target}'")
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_patcher = comfy.model_patcher.ModelPatcher(
model,
load_device = load_device,
offload_device = offload_device,
)
return model_patcher
+51 -136
View File
@@ -2,151 +2,66 @@ import torch
import folder_paths import folder_paths
from nodes import EmptyLatentImage 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): class EmptySanaLatentImage(EmptyLatentImage):
CATEGORY = "ExtraModels/Sana" CATEGORY = "ExtraModels/Sana"
TITLE = "Empty Sana Latent Image" TITLE = "Empty Sana Latent Image"
def generate(self, width, height, batch_size=1):
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,)
def generate(self, width, height, batch_size=1):
latent = torch.zeros([batch_size, 32, height // 32, width // 32], device=self.device)
return ({"samples":latent}, )
class SanaTextEncode: class SanaTextEncode:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
return { return {
"required": { "required": {
"text": ("STRING", {"multiline": True}), "text": ("STRING", {"multiline": True}),
"GEMMA": ("GEMMA",), "GEMMA": ("GEMMA",),
} }
} }
RETURN_TYPES = ("CONDITIONING",) RETURN_TYPES = ("CONDITIONING",)
FUNCTION = "encode" FUNCTION = "encode"
CATEGORY = "ExtraModels/Sana" CATEGORY = "ExtraModels/Sana"
TITLE = "Sana Text Encode" TITLE = "Sana Text Encode"
def encode(self, text, GEMMA=None): def encode(self, text, GEMMA=None):
tokenizer = GEMMA["tokenizer"] tokenizer = GEMMA["tokenizer"]
text_encoder = GEMMA["text_encoder"] text_encoder = GEMMA["text_encoder"]
with torch.no_grad(): with torch.no_grad():
chi_prompt = "\n".join(preset_te_prompt) chi_prompt = "\n".join(preset_te_prompt)
full_prompt = chi_prompt + text full_prompt = chi_prompt + text
num_chi_tokens = len(tokenizer.encode(chi_prompt)) num_chi_tokens = len(tokenizer.encode(chi_prompt))
max_length = num_chi_tokens + 300 - 2 max_length = num_chi_tokens + 300 - 2
tokens = tokenizer( tokens = tokenizer(
[full_prompt], [full_prompt],
max_length=max_length, max_length=max_length,
padding="max_length", padding="max_length",
truncation=True, truncation=True,
return_tensors="pt" return_tensors="pt"
).to(text_encoder.device) ).to(text_encoder.device)
select_idx = [0] + list(range(-300 + 1, 0)) select_idx = [0] + list(range(-300 + 1, 0))
embs = text_encoder(tokens.input_ids, tokens.attention_mask)[0][:, None][:, :, select_idx] embs = text_encoder(tokens.input_ids, tokens.attention_mask)[0][:, None][:, :, select_idx]
emb_masks = tokens.attention_mask[:, select_idx] emb_masks = tokens.attention_mask[:, select_idx]
embs = embs * emb_masks.unsqueeze(-1) embs = embs * emb_masks.unsqueeze(-1)
return ([[embs, {}]], ) return ([[embs, {}]], )
preset_te_prompt = [ preset_te_prompt = [
'Given a user prompt, generate an "Enhanced prompt" that provides detailed visual descriptions suitable for image generation. Evaluate the level of detail in the user prompt:', 'Given a user prompt, generate an "Enhanced prompt" that provides detailed visual descriptions suitable for image generation. Evaluate the level of detail in the user prompt:',
'- If the prompt is simple, focus on adding specifics about colors, shapes, sizes, textures, and spatial relationships to create vivid and concrete scenes.', '- If the prompt is simple, focus on adding specifics about colors, shapes, sizes, textures, and spatial relationships to create vivid and concrete scenes.',
'- If the prompt is already detailed, refine and enhance the existing details slightly without overcomplicating.', '- If the prompt is already detailed, refine and enhance the existing details slightly without overcomplicating.',
'Here are examples of how to transform or refine prompts:', 'Here are examples of how to transform or refine prompts:',
'- User Prompt: A cat sleeping -> Enhanced: A small, fluffy white cat curled up in a round shape, sleeping peacefully on a warm sunny windowsill, surrounded by pots of blooming red flowers.', '- User Prompt: A cat sleeping -> Enhanced: A small, fluffy white cat curled up in a round shape, sleeping peacefully on a warm sunny windowsill, surrounded by pots of blooming red flowers.',
'- User Prompt: A busy city street -> Enhanced: A bustling city street scene at dusk, featuring glowing street lamps, a diverse crowd of people in colorful clothing, and a double-decker bus passing by towering glass skyscrapers.', '- User Prompt: A busy city street -> Enhanced: A bustling city street scene at dusk, featuring glowing street lamps, a diverse crowd of people in colorful clothing, and a double-decker bus passing by towering glass skyscrapers.',
'Please generate only the enhanced description for the prompt below and avoid including any additional commentary or evaluations:', 'Please generate only the enhanced description for the prompt below and avoid including any additional commentary or evaluations:',
'User Prompt: ' 'User Prompt: '
] ]
NODE_CLASS_MAPPINGS = { NODE_CLASS_MAPPINGS = {
"SanaCheckpointLoader" : SanaCheckpointLoader, "SanaTextEncode" : SanaTextEncode,
"SanaResolutionSelect" : SanaResolutionSelect, "EmptySanaLatentImage": EmptySanaLatentImage,
"SanaTextEncode" : SanaTextEncode,
"SanaResolutionCond" : SanaResolutionCond,
"EmptySanaLatentImage": EmptySanaLatentImage,
} }
+2
View File
@@ -2,9 +2,11 @@ import folder_paths
import comfy.utils import comfy.utils
from .PixArt.loader import load_pixart_state_dict from .PixArt.loader import load_pixart_state_dict
from .Sana.loader import load_sana_state_dict
loaders = { loaders = {
"PixArt": load_pixart_state_dict, "PixArt": load_pixart_state_dict,
"Sana": load_sana_state_dict,
} }
class EXMUnetLoader: class EXMUnetLoader: