V 2.0.0 - Auto config #13 - more types
This commit is contained in:
@@ -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)
|
||||
|
||||
+105
-1
@@ -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)
|
||||
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)
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user