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"],
|
|
||||||
})
|
|
||||||
+89
-85
@@ -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
@@ -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,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:
|
||||||
|
|||||||
Reference in New Issue
Block a user