diff --git a/Nodes/Dashboard.py b/Nodes/Dashboard.py index c789cc9..fd68ff8 100644 --- a/Nodes/Dashboard.py +++ b/Nodes/Dashboard.py @@ -895,6 +895,10 @@ class PrimereCKPTLoader: OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE = model_loaders.load_pixart_model(self, ckpt_name, concept_data) case 'AuraFlow': OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE = model_loaders.load_auraflow_model(self, ckpt_name, concept_data) + case 'SANA1024' | 'SANA512': + OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE = model_loaders.load_sana_model(self, ckpt_name, concept_data) + case 'KwaiKolors': + OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE = model_loaders.load_kolors_model(self, ckpt_name, concept_data) return (OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE, MODEL_VERSION_ORIGINAL) @@ -1224,6 +1228,10 @@ class PrimereCLIP: return clipping.encode_flux(clip, positive_text, negative_text, t5xxl_prompt, concept_data, workflow_tuple) case 'PixartSigma': return clipping.encode_pixart_sigma(clip, positive_text, negative_text, workflow_tuple) + case 'SANA1024' | 'SANA512': + return clipping.encode_sana(clip, positive_text, negative_text, t5xxl_prompt, workflow_tuple) + case 'KwaiKolors': + return clipping.encode_kolors(clip, positive_text, negative_text, t5xxl_prompt, workflow_tuple) case _: clip = clipping.apply_clip_overrides(self, clip, workflow_tuple) return clipping.encode_standard(clip, positive_text, negative_text, t5xxl_prompt, adv_encode, token_normalization, weight_interpretation, positive_l, negative_l, width, height, workflow_tuple, advanced_encode) diff --git a/components/clipping.py b/components/clipping.py index 3cde0d9..162217a 100644 --- a/components/clipping.py +++ b/components/clipping.py @@ -1,4 +1,5 @@ import torch +import gc from .long_clip_model import longclip from comfy.sd1_clip import load_embed, ClipTokenWeightEncoder from comfy.sd1_clip import token_weights, escape_important, unescape_important @@ -9,6 +10,8 @@ import comfy_extras.nodes_flux as nodes_flux from . import utility import nodes from ..Nodes.modules import long_clip as long_clip_module +from .sana.diffusion.model.utils import prepare_prompt_ar +from .sana.diffusion.data.datasets.utils import ASPECT_RATIO_1024_TEST class SDLongClipModel(torch.nn.Module, ClipTokenWeightEncoder): LAYERS = [ @@ -690,4 +693,105 @@ def encode_flux(clip, positive_text, negative_text, t5xxl_prompt, concept_data, return (cond_pos, cond_neg, positive_text, negative_text, t5xxl_prompt, "", "", workflow_tuple) else: cond_pos = nodes_flux.CLIPTextEncodeFlux.execute(clip, positive_text, t5xxl_prompt, flux_guidance)[0] - return (cond_pos, cond_pos, positive_text, negative_text, t5xxl_prompt, "", "", workflow_tuple) \ No newline at end of file + return (cond_pos, cond_pos, positive_text, negative_text, t5xxl_prompt, "", "", workflow_tuple) + + +_SANA_MAX_TOKENS = 300 +_SANA_CHI_PROMPT = "\n".join([ + 'Create one detailed perfect prompt from given User Prompt for stable diffusion text-to-image text2image modern DiT models.', + 'Generate only the one enhanced description for the prompt below, avoid including any additional questions comments or evaluations.', + 'User Prompt: ', +]) + + +def _sana_encode_text(tokenizer, text_encoder, text, device): + full_prompt = _SANA_CHI_PROMPT + text + num_chi_tokens = len(tokenizer.encode(_SANA_CHI_PROMPT)) + max_length = num_chi_tokens + _SANA_MAX_TOKENS - 2 + tokens = tokenizer([full_prompt], max_length=max_length, padding="max_length", truncation=True, return_tensors="pt").to(device) + select_idx = [0] + list(range(-_SANA_MAX_TOKENS + 1, 0)) + embs = text_encoder(tokens.input_ids, tokens.attention_mask)[0][:, None][:, :, select_idx] + masks = tokens.attention_mask[:, select_idx] + return embs * masks.unsqueeze(-1) + + +def encode_sana(clip, positive_text, negative_text, t5xxl_prompt, workflow_tuple): + scheduler_name = workflow_tuple.get('scheduler_name', 'flow_dpm-solver') if workflow_tuple else 'flow_dpm-solver' + device = model_management.get_torch_device() + + if scheduler_name == 'flow_dpm-solver' and hasattr(clip, 'text_encoder'): + clip.text_encoder.to(device) + null_token = clip.tokenizer(negative_text, max_length=_SANA_MAX_TOKENS, padding="max_length", truncation=True, return_tensors="pt").to(device) + null_embs = clip.text_encoder(null_token.input_ids, null_token.attention_mask)[0] + with torch.no_grad(): + prompts = [prepare_prompt_ar(positive_text, ASPECT_RATIO_1024_TEST, device=device, show=False)[0].strip()] + num_chi_tokens = len(clip.tokenizer.encode(_SANA_CHI_PROMPT)) + max_length_all = num_chi_tokens + _SANA_MAX_TOKENS - 2 + caption_token = clip.tokenizer([_SANA_CHI_PROMPT + positive_text], max_length=max_length_all, padding="max_length", truncation=True, return_tensors="pt").to(device) + select_index = [0] + list(range(-_SANA_MAX_TOKENS + 1, 0)) + caption_embs = clip.text_encoder(caption_token.input_ids, caption_token.attention_mask)[0][:, None][:, :, select_index] + emb_masks = caption_token.attention_mask[:, select_index] + null_y = null_embs.repeat(len(prompts), 1, 1)[:, None] + clip.text_encoder.to(model_management.text_encoder_offload_device()) + comfy.model_management.soft_empty_cache(True) + return ([[caption_embs, {"emb_masks": emb_masks}]], [[null_y, {}]], positive_text, negative_text, t5xxl_prompt, "", "", workflow_tuple) + else: + tokenizer = clip["tokenizer"] + text_encoder = clip["text_encoder"] + enc_device = text_encoder.device + with torch.no_grad(): + sana_embs_pos = _sana_encode_text(tokenizer, text_encoder, positive_text, enc_device) + sana_embs_neg = _sana_encode_text(tokenizer, text_encoder, negative_text, enc_device) + return ([[sana_embs_pos, {}]], [[sana_embs_neg, {}]], positive_text, negative_text, t5xxl_prompt, "", "", workflow_tuple) + + +def encode_kolors(clip, positive_text, negative_text, t5xxl_prompt, workflow_tuple): + positive_text = utility.DiT_cleaner(positive_text) + negative_text = utility.DiT_cleaner(negative_text) + device = model_management.text_encoder_device() + + try: + model_management.unload_all_models() + model_management.soft_empty_cache() + except Exception: + pass + + tokenizer = clip['tokenizer'] + text_encoder = clip['text_encoder'] + model_management.soft_empty_cache() + + prompt_embeds_dtype = text_encoder.dtype if text_encoder is not None else torch.float16 + try: + text_encoder.to(dtype=prompt_embeds_dtype, device=device) + except Exception: + pass + + text_inputs = tokenizer(positive_text, padding="max_length", max_length=256, truncation=True, return_tensors="pt").to(device) + output = text_encoder(input_ids=text_inputs['input_ids'], attention_mask=text_inputs['attention_mask'], position_ids=text_inputs['position_ids'], output_hidden_states=True) + prompt_embeds = output.hidden_states[-2].permute(1, 0, 2).clone() + text_proj = output.hidden_states[-1][-1, :, :].clone() + bs_embed, seq_len, _ = prompt_embeds.shape + prompt_embeds = prompt_embeds.view(bs_embed, seq_len, -1) + + uncond_input = tokenizer([negative_text], padding="max_length", max_length=prompt_embeds.shape[1], truncation=True, return_tensors="pt").to(device) + output = text_encoder(input_ids=uncond_input['input_ids'], attention_mask=uncond_input['attention_mask'], position_ids=uncond_input['position_ids'], output_hidden_states=True) + negative_prompt_embeds = output.hidden_states[-2].permute(1, 0, 2).clone() + negative_text_proj = output.hidden_states[-1][-1, :, :].clone() + negative_prompt_embeds = negative_prompt_embeds.to(dtype=text_encoder.dtype, device=device).view(1, negative_prompt_embeds.shape[1], -1) + + text_proj = text_proj.view(text_proj.shape[0], -1) + negative_text_proj = negative_text_proj.view(negative_text_proj.shape[0], -1) + + try: + model_management.soft_empty_cache() + except Exception: + pass + gc.collect() + + kolors_embeds = { + 'prompt_embeds': prompt_embeds.half(), + 'negative_prompt_embeds': negative_prompt_embeds.half(), + 'pooled_prompt_embeds': text_proj.half(), + 'negative_pooled_prompt_embeds': negative_text_proj.half(), + } + return (kolors_embeds, None, positive_text, negative_text, t5xxl_prompt, "", "", workflow_tuple) \ No newline at end of file diff --git a/components/models.py b/components/models.py index 122d92a..7aae58f 100644 --- a/components/models.py +++ b/components/models.py @@ -1,4 +1,5 @@ import os +import torch import comfy import comfy.sd import comfy.utils @@ -6,15 +7,33 @@ import folder_paths import nodes import comfy_extras.nodes_sd3 as nodes_sd3 import comfy_extras.nodes_model_advanced as nodes_model_advanced +from comfy import model_management from pathlib import Path from .tree import PRIMERE_ROOT from . import utility from . import nf4_helper +from . import sana_utils from .gguf import nodes as gguf_nodes import difflib import numpy as np +import pyrallis from ComfyUI_ExtraModels.PixArt.loader import load_pixart from ComfyUI_ExtraModels.PixArt.conf import pixart_conf +from diffusers import UNet2DConditionModel, EulerDiscreteScheduler +from .kolors.pipelines.pipeline_stable_diffusion_xl_chatglm_256 import StableDiffusionXLPipeline +from .kolors.models.tokenization_chatglm import ChatGLMTokenizer +from .kolors.models.modeling_chatglm import ChatGLMModel +from ComfyUI_ExtraModels.Sana.conf import sana_conf +from ComfyUI_ExtraModels.Sana.loader import load_sana +from ComfyUI_ExtraModels.VAE.conf import vae_conf +from ComfyUI_ExtraModels.VAE.loader import EXVAE +from ComfyUI_ExtraModels.utils.dtype import string_to_dtype +from transformers import AutoTokenizer, T5Tokenizer, T5EncoderModel, AutoModelForCausalLM, BitsAndBytesConfig +from .sana.diffusion.model.builder import build_model +from .sana.diffusion.model.dc_ae.efficientvit.ae_model_zoo import create_dc_ae_model_cfg +from .sana.diffusion.model.dc_ae.efficientvit.models.efficientvit.dc_ae import DCAE +from .sana.diffusion.utils.config import SanaConfig +from .sana.pipeline.sana_pipeline import SanaPipeline def resolve_symlink(ckpt_name): @@ -302,3 +321,174 @@ def load_lcm_model(loader_self, ckpt_name, concept_data): OUTPUT_MODEL = m return OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE + + +def load_sana_model(loader_self, ckpt_name, concept_data): + encoder_path = concept_data.get('encoder_1', 'gemma-2-2b-it') + weight_dtype_str = concept_data.get('weight_dtype', 'fp16') + vae_name = concept_data.get('vae') + precision = concept_data.get('precision', 'fp16') + scheduler_name = concept_data.get('scheduler_name', 'flow_dpm-solver') + + device = model_management.get_torch_device() + fullpathFile = folder_paths.get_full_path('checkpoints', ckpt_name) + if os.path.islink(str(fullpathFile)): + fullpathFile = Path(str(fullpathFile)).resolve() + + text_encoder_dir = os.path.join(folder_paths.models_dir, 'text_encoders', encoder_path) + if not os.path.exists(text_encoder_dir): + text_encoder_dir = os.path.join(PRIMERE_ROOT, 'Nodes', 'Downloads', 'LLM', encoder_path) + + vae_path = folder_paths.get_full_path("vae", vae_name) + text_encoder_dtype = model_management.text_encoder_dtype(device) + + if scheduler_name == 'flow_dpm-solver': + dtype = utility.get_dtype_by_name(weight_dtype_str) + cfg = create_dc_ae_model_cfg('dc-ae-f32c32-sana-1.0') + vae = DCAE(cfg) + state_dict = comfy.utils.load_torch_file(vae_path, safe_load=True) + vae.load_state_dict(state_dict, strict=False) + vae_dtype = model_management.vae_dtype(device, [torch.float16, torch.bfloat16, torch.float32]) + vae.to(vae_dtype).eval() + + if "T5" in encoder_path: + tokenizer = T5Tokenizer.from_pretrained(str(text_encoder_dir)) + llm_model = None + text_encoder = T5EncoderModel.from_pretrained(str(text_encoder_dir), torch_dtype=text_encoder_dtype) + else: + tokenizer = AutoTokenizer.from_pretrained(text_encoder_dir) + if precision == '8-bit': + quantization_config = BitsAndBytesConfig(load_in_8bit=True) + elif precision == '4-bit': + quantization_config = BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_compute_dtype=text_encoder_dtype) + else: + quantization_config = None + if '-4bit' in encoder_path: + llm_model = AutoModelForCausalLM.from_pretrained(text_encoder_dir, torch_dtype=text_encoder_dtype) + else: + llm_model = AutoModelForCausalLM.from_pretrained(text_encoder_dir, quantization_config=quantization_config, torch_dtype=text_encoder_dtype) + tokenizer.padding_side = "right" + text_encoder = llm_model.get_decoder() + + text_encoder.to(device) + state_dict = comfy.utils.load_torch_file(str(fullpathFile), safe_load=True) + is_1600M = state_dict['final_layer.scale_shift_table'].shape[1] == 2240 + if '512px' in ckpt_name: + config_path = os.path.join(PRIMERE_ROOT, 'components', 'sana', 'configs', 'sana_config', '512ms', 'Sana_1600M_img512.yaml') if is_1600M else os.path.join(PRIMERE_ROOT, 'components', 'sana', 'configs', 'sana_config', '512ms', 'Sana_600M_img512.yaml') + else: + config_path = os.path.join(PRIMERE_ROOT, 'components', 'sana', 'configs', 'sana_config', '1024ms', 'Sana_1600M_img1024_AdamW.yaml') if is_1600M else os.path.join(PRIMERE_ROOT, 'components', 'sana', 'configs', 'sana_config', '1024ms', 'Sana_600M_img1024.yaml') + config = pyrallis.load(SanaConfig, open(config_path)) + + pred_sigma = getattr(config.scheduler, "pred_sigma", True) + learn_sigma = getattr(config.scheduler, "learn_sigma", True) and pred_sigma + image_size = config.model.image_size + latent_size = image_size // config.vae.vae_downsample_rate + model_kwargs = { + "input_size": latent_size, + "pe_interpolation": config.model.pe_interpolation, + "config": config, + "model_max_length": config.text_encoder.model_max_length, + "qk_norm": config.model.qk_norm, + "micro_condition": config.model.micro_condition, + "caption_channels": text_encoder.config.hidden_size, + "y_norm": config.text_encoder.y_norm, + "attn_type": config.model.attn_type, + "ffn_type": config.model.ffn_type, + "mlp_ratio": config.model.mlp_ratio, + "mlp_acts": list(config.model.mlp_acts), + "in_channels": config.vae.vae_latent_dim, + "y_norm_scale_factor": config.text_encoder.y_norm_scale_factor, + "use_pe": config.model.use_pe, + "pred_sigma": pred_sigma, + "learn_sigma": learn_sigma, + "use_fp32_attention": config.model.get("fp32_attention", False) and config.model.mixed_precision != "bf16", + } + unet = build_model(config.model.model, **model_kwargs) + unet.to(dtype) + state_dict = state_dict.get("state_dict", state_dict) + if "pos_embed" in state_dict: + del state_dict["pos_embed"] + unet.load_state_dict(state_dict, strict=False) + del state_dict + unet.eval().to(dtype) + pipe = SanaPipeline(config, vae, dtype, unet) + + SANA_MODEL = { + 'pipe': pipe, + 'unet': unet, + 'text_encoder_model': llm_model, + 'tokenizer': tokenizer, + 'text_encoder': text_encoder, + 'vae': vae, + 'device': device, + } + SANA_VAE = sana_utils.first_stage_model(vae) + SANA_CLIP = sana_utils.cond_stage_model(tokenizer, text_encoder) + else: + model_keys = list(sana_conf.keys()) + model_conf = sana_conf[model_keys[1]] + SANA_MODEL = load_sana(model_path=str(fullpathFile), model_conf=model_conf) + + vae_config = vae_conf['dcae-f32c32-sana-1.0'] + SANA_VAE = EXVAE(vae_path, vae_config, string_to_dtype(weight_dtype_str.upper(), "vae")) + + tokenizer = AutoTokenizer.from_pretrained(text_encoder_dir) + text_encoder_model = AutoModelForCausalLM.from_pretrained(text_encoder_dir, torch_dtype=text_encoder_dtype) + tokenizer.padding_side = "right" + text_encoder = text_encoder_model.get_decoder() + if device != "cpu": + text_encoder = text_encoder.to(device) + + SANA_CLIP = { + "tokenizer": tokenizer, + "text_encoder": text_encoder, + "text_encoder_model": text_encoder_model, + } + + return SANA_MODEL, SANA_CLIP, SANA_VAE + + +def load_kolors_model(loader_self, ckpt_name, concept_data): + weight_dtype_str = concept_data.get('weight_dtype', 'fp16') + precision = concept_data.get('precision', 'quant8') + vae_name = concept_data.get('vae') + + dtype_map = {'bf16': torch.bfloat16, 'fp16': torch.float16, 'fp32': torch.float32} + dtype = dtype_map.get(weight_dtype_str, torch.float16) + + model_name = Path(ckpt_name).stem + fullpathFile = folder_paths.get_full_path('checkpoints', ckpt_name) + if os.path.islink(str(fullpathFile)): + link_path = Path(str(fullpathFile)).resolve() + model_name = Path(link_path.parent.parent).stem + + model_path = os.path.join(folder_paths.models_dir, "diffusers", model_name) + pbar = comfy.utils.ProgressBar(4) + scheduler = EulerDiscreteScheduler.from_pretrained(model_path, subfolder='scheduler') + unet = UNet2DConditionModel.from_pretrained(model_path, subfolder='unet', variant="fp16", revision=None, low_cpu_mem_usage=True).to(dtype).eval() + pipeline = StableDiffusionXLPipeline(unet=unet, scheduler=scheduler) + KOLORS_MODEL = {'pipeline': pipeline, 'dtype': dtype} + pbar.update(1) + + pbar = comfy.utils.ProgressBar(2) + text_encoder_path = os.path.join(model_path, "text_encoder") + text_encoder = ChatGLMModel.from_pretrained(text_encoder_path, torch_dtype=torch.float16) + if precision == 'quant8': + try: + text_encoder.quantize(8) + except Exception: + print('Quantization 8 failed...') + elif precision == 'quant4': + try: + text_encoder.quantize(4) + except Exception: + print('Quantization 4 failed...') + tokenizer = ChatGLMTokenizer.from_pretrained(text_encoder_path) + pbar.update(1) + CHATGLM3_MODEL = {'text_encoder': text_encoder, 'tokenizer': tokenizer} + + if not vae_name: + raise ValueError("KwaiKolors requires an explicit VAE. Set 'vae' in the concept data (e.g. 'kolors\\diffusion_pytorch_model.fp16.safetensors').") + OUTPUT_VAE = utility.vae_loader_class.load_vae(vae_name)[0] + + return KOLORS_MODEL, CHATGLM3_MODEL, OUTPUT_VAE