diff --git a/Sana/conf.py b/Sana/conf.py deleted file mode 100644 index 3719c0f..0000000 --- a/Sana/conf.py +++ /dev/null @@ -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"], -}) diff --git a/Sana/loader.py b/Sana/loader.py index 806ca64..60eea7a 100644 --- a/Sana/loader.py +++ b/Sana/loader.py @@ -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 SanaConfig(comfy.supported_models_base.BASE): + unet_class = SanaMS + unet_config = {} + unet_extra_config = {} -class EXM_Sana(comfy.supported_models_base.BASE): - unet_config = {} - unet_extra_config = {} - latent_format = SanaLatent + latent_format = SanaLatent + sampling_settings = { + "shift": 3.0, + } - 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 + def model_type(self, state_dict, prefix=""): + return comfy.model_base.ModelType.FLOW - 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 SanaModel(comfy.model_base.BaseModel): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) -class EXM_Sana_Model(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] + + 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 - 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) - 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 +# 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"], +} diff --git a/Sana/nodes.py b/Sana/nodes.py index abce9ba..ca1628b 100644 --- a/Sana/nodes.py +++ b/Sana/nodes.py @@ -2,151 +2,66 @@ 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" - - 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,) + CATEGORY = "ExtraModels/Sana" + 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 SanaTextEncode: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "text": ("STRING", {"multiline": True}), - "GEMMA": ("GEMMA",), - } - } + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "text": ("STRING", {"multiline": True}), + "GEMMA": ("GEMMA",), + } + } - RETURN_TYPES = ("CONDITIONING",) - FUNCTION = "encode" - CATEGORY = "ExtraModels/Sana" - TITLE = "Sana Text Encode" + RETURN_TYPES = ("CONDITIONING",) + FUNCTION = "encode" + CATEGORY = "ExtraModels/Sana" + TITLE = "Sana Text Encode" - def encode(self, text, GEMMA=None): - tokenizer = GEMMA["tokenizer"] - text_encoder = GEMMA["text_encoder"] - - with torch.no_grad(): - chi_prompt = "\n".join(preset_te_prompt) - full_prompt = chi_prompt + text - num_chi_tokens = len(tokenizer.encode(chi_prompt)) - max_length = num_chi_tokens + 300 - 2 - - tokens = tokenizer( - [full_prompt], - max_length=max_length, - padding="max_length", - truncation=True, - return_tensors="pt" - ).to(text_encoder.device) - - select_idx = [0] + list(range(-300 + 1, 0)) - embs = text_encoder(tokens.input_ids, tokens.attention_mask)[0][:, None][:, :, select_idx] - emb_masks = tokens.attention_mask[:, select_idx] - embs = embs * emb_masks.unsqueeze(-1) - - return ([[embs, {}]], ) + def encode(self, text, GEMMA=None): + tokenizer = GEMMA["tokenizer"] + text_encoder = GEMMA["text_encoder"] + + with torch.no_grad(): + chi_prompt = "\n".join(preset_te_prompt) + full_prompt = chi_prompt + text + num_chi_tokens = len(tokenizer.encode(chi_prompt)) + max_length = num_chi_tokens + 300 - 2 + + tokens = tokenizer( + [full_prompt], + max_length=max_length, + padding="max_length", + truncation=True, + return_tensors="pt" + ).to(text_encoder.device) + + select_idx = [0] + list(range(-300 + 1, 0)) + embs = text_encoder(tokens.input_ids, tokens.attention_mask)[0][:, None][:, :, select_idx] + emb_masks = tokens.attention_mask[:, select_idx] + embs = embs * emb_masks.unsqueeze(-1) + + return ([[embs, {}]], ) 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:', - '- 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.', - '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 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:', - '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 already detailed, refine and enhance the existing details slightly without overcomplicating.', + '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 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:', + 'User Prompt: ' ] NODE_CLASS_MAPPINGS = { - "SanaCheckpointLoader" : SanaCheckpointLoader, - "SanaResolutionSelect" : SanaResolutionSelect, - "SanaTextEncode" : SanaTextEncode, - "SanaResolutionCond" : SanaResolutionCond, - "EmptySanaLatentImage": EmptySanaLatentImage, + "SanaTextEncode" : SanaTextEncode, + "EmptySanaLatentImage": EmptySanaLatentImage, } diff --git a/nodes.py b/nodes.py index eab0d3d..31d7f53 100644 --- a/nodes.py +++ b/nodes.py @@ -2,9 +2,11 @@ import folder_paths import comfy.utils from .PixArt.loader import load_pixart_state_dict +from .Sana.loader import load_sana_state_dict loaders = { "PixArt": load_pixart_state_dict, + "Sana": load_sana_state_dict, } class EXMUnetLoader: