diff --git a/Nodes/Dashboard.py b/Nodes/Dashboard.py index b9d75c8..785f718 100644 --- a/Nodes/Dashboard.py +++ b/Nodes/Dashboard.py @@ -18,6 +18,7 @@ from pathlib import Path import re from ..components import hypernetwork from ..components import clipping +from ..components import models as model_loaders from ..components import nf4_helper from ..components import sana_utils import comfy.sd @@ -872,231 +873,17 @@ class PrimereCKPTLoader: match model_concept: case 'SD1' | 'SD2' | 'SDXL' | 'Illustrious': - if os.path.isfile(ModelConfigFullPath) and use_yaml == True: - ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name) - try: - LOADED_CHECKPOINT = comfy.sd.load_checkpoint(ModelConfigFullPath, ckpt_path, True, True, None, None, None) - except Exception: - LOADED_CHECKPOINT = nodes.CheckpointLoaderSimple.load_checkpoint(self, ckpt_name) - else: - LOADED_CHECKPOINT = nodes.CheckpointLoaderSimple.load_checkpoint(self, ckpt_name) - OUTPUT_MODEL = LOADED_CHECKPOINT[0] - OUTPUT_CLIP = LOADED_CHECKPOINT[1] - - vae_selection = concept_data.get('vae_selection', True) - vae_name = concept_data.get('vae', None) - if not vae_selection and vae_name: - OUTPUT_VAE = utility.vae_loader_class.load_vae(vae_name)[0] - elif len(LOADED_CHECKPOINT) >= 3 and type(LOADED_CHECKPOINT[2]).__name__ == 'VAE': - OUTPUT_VAE = LOADED_CHECKPOINT[2] - elif vae_name: - OUTPUT_VAE = utility.vae_loader_class.load_vae(vae_name)[0] - else: - OUTPUT_VAE = LOADED_CHECKPOINT[2] - - return (OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE, MODEL_VERSION_ORIGINAL) - + OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE = model_loaders.load_sd_model(self, ckpt_name, use_yaml, ModelConfigFullPath, concept_data) case 'SD3': - sd3_gguf = False - fullpathFile = folder_paths.get_full_path('checkpoints', ckpt_name) - is_link = os.path.islink(str(fullpathFile)) - if is_link: - File_link = Path(str(fullpathFile)).resolve() - model_ext = os.path.splitext(File_link)[1].lower() - if model_ext == '.gguf': - linkName_U = str(folder_paths.folder_names_and_paths["diffusion_models"][0][0]) - linkName_D = str(folder_paths.folder_names_and_paths["diffusion_models"][0][1]) - sd3_gguf = str(File_link).replace(linkName_U + '\\', '').replace(linkName_D + '\\', '') - if str(Path(linkName_U).stem) in sd3_gguf: - sd3_gguf = sd3_gguf.split(Path(linkName_U).stem + '\\', 1)[1] - if str(Path(linkName_D).stem) in sd3_gguf: - sd3_gguf = sd3_gguf.split(Path(linkName_D).stem + '\\', 1)[1] - - if sd3_gguf: - OUTPUT_MODEL = gguf_nodes.UnetLoaderGGUF.load_unet(self, sd3_gguf)[0] - else: - LOADED_CHECKPOINT = nodes.CheckpointLoaderSimple.load_checkpoint(self, ckpt_name) - OUTPUT_MODEL = LOADED_CHECKPOINT[0] - - clip_selection = concept_data.get('clip_selection', True) - clip_from_ckpt = len(LOADED_CHECKPOINT) >= 2 and type(LOADED_CHECKPOINT[1]).__name__ == 'CLIP' - if not clip_selection or not clip_from_ckpt: - encoder_1 = concept_data.get('encoder_1', None) - encoder_2 = concept_data.get('encoder_2', None) - encoder_3 = concept_data.get('encoder_3', None) - OUTPUT_CLIP = nodes_sd3.TripleCLIPLoader.execute(encoder_3, encoder_2, encoder_1)[0] - else: - OUTPUT_CLIP = LOADED_CHECKPOINT[1] - - vae_name = concept_data.get('vae', None) - if len(LOADED_CHECKPOINT) >= 3 and type(LOADED_CHECKPOINT[2]).__name__ == 'VAE': - OUTPUT_VAE = LOADED_CHECKPOINT[2] - elif vae_name: - OUTPUT_VAE = utility.vae_loader_class.load_vae(vae_name)[0] - else: - OUTPUT_VAE = LOADED_CHECKPOINT[2] - - if concept_data.get('speed_lora') == True: - speed_lora_name = concept_data.get('speed_lora_name', None) - speed_lora_strength = concept_data.get('speed_lora_strength', 1) - if speed_lora_name and speed_lora_strength != 0: - lora_path = os.path.join(PRIMERE_ROOT, 'Nodes', 'Downloads', speed_lora_name) - if os.path.exists(lora_path): - lora = None - if self.loaded_lora is not None: - if self.loaded_lora[0] == lora_path: - lora = self.loaded_lora[1] - else: - temp = self.loaded_lora - self.loaded_lora = None - del temp - if lora is None: - lora = comfy.utils.load_torch_file(lora_path, safe_load=True) - self.loaded_lora = (lora_path, lora) - - OUTPUT_MODEL = comfy.sd.load_lora_for_models(OUTPUT_MODEL, None, lora, speed_lora_strength, 0)[0] - - return (OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE, MODEL_VERSION_ORIGINAL) - + OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE = model_loaders.load_sd3_model(self, ckpt_name, concept_data) case 'StableCascade': - stage_b = concept_data.get('encoder_1', None) - cascade_clip = concept_data.get('encoder_3', None) - vae_name = concept_data.get('vae', None) - - OUTPUT_CLIP = nodes.CLIPLoader.load_clip(self, cascade_clip, 'stable_cascade')[0] - OUTPUT_VAE = utility.vae_loader_class.load_vae(vae_name)[0] - MODEL_B = nodes.UNETLoader.load_unet(self, stage_b, 'default')[0] - - fullpathFile = folder_paths.get_full_path('checkpoints', ckpt_name) - is_link = os.path.islink(str(fullpathFile)) - if not is_link: - MODEL_C = nodes.UNETLoader.load_unet(self, ckpt_name, 'default')[0] - else: - File_link = Path(str(fullpathFile)).resolve() - linkName_U = str(folder_paths.folder_names_and_paths["diffusion_models"][0][0]) - linkName_D = str(folder_paths.folder_names_and_paths["diffusion_models"][0][1]) - linkedFileName = str(File_link).replace(linkName_U + '\\', '').replace(linkName_D + '\\', '') - if str(Path(linkName_U).stem) in linkedFileName: - linkedFileName = linkedFileName.split(Path(linkName_U).stem + '\\', 1)[1] - if str(Path(linkName_D).stem) in linkedFileName: - linkedFileName = linkedFileName.split(Path(linkName_D).stem + '\\', 1)[1] - MODEL_C = nodes.UNETLoader.load_unet(self, linkedFileName, 'default')[0] - - OUTPUT_MODEL = [MODEL_B, MODEL_C] - return (OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE, MODEL_VERSION_ORIGINAL) - + OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE = model_loaders.load_stable_cascade_model(self, ckpt_name, concept_data) case 'Z-Image': - weight_dtype = concept_data.get('weight_dtype', 'default') - if 'e4m3fn' in ckpt_name: - weight_dtype = 'fp8_e4m3fn' - if 'e5m2' in ckpt_name: - weight_dtype = 'fp8_e5m2' - - fullpathFile = folder_paths.get_full_path('checkpoints', ckpt_name) - is_link = os.path.islink(str(fullpathFile)) - if is_link: - File_link = Path(str(fullpathFile)).resolve() - linkName_U = str(folder_paths.folder_names_and_paths["diffusion_models"][0][0]) - linkName_D = str(folder_paths.folder_names_and_paths["diffusion_models"][0][1]) - linkedFileName = str(File_link).replace(linkName_U + '\\', '').replace(linkName_D + '\\', '') - model_ext = os.path.splitext(linkedFileName)[1].lower() - if str(Path(linkName_U).stem) in linkedFileName: - linkedFileName = linkedFileName.split(Path(linkName_U).stem + '\\', 1)[1] - if str(Path(linkName_D).stem) in linkedFileName: - linkedFileName = linkedFileName.split(Path(linkName_D).stem + '\\', 1)[1] - - if 'diffusion_models' in str(File_link): - if model_ext == '.gguf': - OUTPUT_MODEL = gguf_nodes.UnetLoaderGGUF.load_unet(self, linkedFileName)[0] - else: - OUTPUT_MODEL = nodes.UNETLoader.load_unet(self, linkedFileName, weight_dtype)[0] - elif 'unet' in str(File_link): - if model_ext == '.gguf': - OUTPUT_MODEL = gguf_nodes.UnetLoaderGGUF.load_unet(self, linkedFileName)[0] - else: - try: - OUTPUT_MODEL = nodes.UNETLoader.load_unet(self, linkedFileName, weight_dtype)[0] - except Exception: - OUTPUT_MODEL = nf4_helper.UNETLoaderNF4.load_nf4unet(linkedFileName)[0] - else: - OUTPUT_MODEL = nodes.CheckpointLoaderSimple.load_checkpoint(self, ckpt_name)[0] - - encoder_1 = concept_data.get('encoder_1', None) - clip_ext = os.path.splitext(encoder_1)[1].lower() if encoder_1 else '' - if clip_ext == '.gguf': - OUTPUT_CLIP = gguf_nodes.CLIPLoaderGGUF.load_clip(self, encoder_1, 'qwen_image')[0] - else: - OUTPUT_CLIP = nodes.CLIPLoader.load_clip(self, encoder_1, 'flux2')[0] - - vae_name = concept_data.get('vae', None) - OUTPUT_VAE = utility.vae_loader_class.load_vae(vae_name)[0] - return (OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE, MODEL_VERSION_ORIGINAL) - + OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE = model_loaders.load_zimage_model(self, ckpt_name, concept_data) case 'Flux': - weight_dtype = concept_data.get('weight_dtype', 'default') - is_gguf_model = False - fullpathFile = folder_paths.get_full_path('checkpoints', ckpt_name) - is_link = os.path.islink(str(fullpathFile)) - if is_link: - File_link = Path(str(fullpathFile)).resolve() - model_ext = os.path.splitext(File_link)[1].lower() - linkName_U = str(folder_paths.folder_names_and_paths["diffusion_models"][0][0]) - linkName_D = str(folder_paths.folder_names_and_paths["diffusion_models"][0][1]) - linkedFileName = str(File_link).replace(linkName_U + '\\', '').replace(linkName_D + '\\', '') - if str(Path(linkName_U).stem) in linkedFileName: - linkedFileName = linkedFileName.split(Path(linkName_U).stem + '\\', 1)[1] - if str(Path(linkName_D).stem) in linkedFileName: - linkedFileName = linkedFileName.split(Path(linkName_D).stem + '\\', 1)[1] - if 'diffusion_models' in str(File_link) or 'unet' in str(File_link): - if model_ext == '.gguf': - OUTPUT_MODEL = gguf_nodes.UnetLoaderGGUF.load_unet(self, linkedFileName)[0] - is_gguf_model = True - else: - try: - OUTPUT_MODEL = nodes.UNETLoader.load_unet(self, linkedFileName, weight_dtype)[0] - except Exception: - OUTPUT_MODEL = nf4_helper.UNETLoaderNF4.load_nf4unet(linkedFileName)[0] - else: - OUTPUT_MODEL = nodes.CheckpointLoaderSimple.load_checkpoint(self, linkedFileName)[0] - else: - OUTPUT_MODEL = nodes.CheckpointLoaderSimple.load_checkpoint(self, ckpt_name)[0] - encoder_1 = concept_data.get('encoder_1', None) - encoder_2 = concept_data.get('encoder_2', None) - clip_ext_1 = os.path.splitext(encoder_1)[1].lower() if encoder_1 else '' - clip_ext_2 = os.path.splitext(encoder_2)[1].lower() if encoder_2 else '' - use_gguf_clip = is_gguf_model or clip_ext_1 == '.gguf' or clip_ext_2 == '.gguf' - if use_gguf_clip: - OUTPUT_CLIP = gguf_nodes.DualCLIPLoaderGGUF.load_clip(self, encoder_2, encoder_1, 'flux')[0] - else: - OUTPUT_CLIP = nodes.DualCLIPLoader.load_clip(self, encoder_2, encoder_1, 'flux')[0] - vae_name = concept_data.get('vae', None) - OUTPUT_VAE = utility.vae_loader_class.load_vae(vae_name)[0] + OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE = model_loaders.load_flux_model(self, ckpt_name, concept_data) - if concept_data.get('speed_lora') == True: - flux_lora_name = concept_data.get('speed_lora_name', None) - flux_lora_strength = concept_data.get('speed_lora_strength', 1) - elif concept_data.get('srpo_lora') == True: - flux_lora_name = concept_data.get('srpo_lora_name', None) - flux_lora_strength = concept_data.get('srpo_lora_strength', 1) - else: - flux_lora_name = None - flux_lora_strength = 0 - if flux_lora_name and flux_lora_strength != 0: - lora_path = os.path.join(PRIMERE_ROOT, 'Nodes', 'Downloads', flux_lora_name) - if os.path.exists(lora_path): - lora = None - if self.loaded_lora is not None: - if self.loaded_lora[0] == lora_path: - lora = self.loaded_lora[1] - else: - temp = self.loaded_lora - self.loaded_lora = None - del temp - if lora is None: - lora = comfy.utils.load_torch_file(lora_path, safe_load=True) - self.loaded_lora = (lora_path, lora) - OUTPUT_MODEL = comfy.sd.load_lora_for_models(OUTPUT_MODEL, None, lora, flux_lora_strength, 0)[0] - return (OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE, MODEL_VERSION_ORIGINAL) + return (OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE, MODEL_VERSION_ORIGINAL) class AnyType(str): def __ne__(self, __value: object) -> bool: @@ -1578,98 +1365,14 @@ class PrimereCLIP: if clip_model != 'Default' and clip_mode == True: clip_path = folder_paths.get_full_path("clip", clip_model) if clip_path is not None: - concept_type = 'stable_diffusion' - clip = nodes.CLIPLoader.load_clip(self, clip_model, concept_type)[0] - - if adv_encode == True: - tokens_p = clip.tokenize(positive_text) - tokens_n = clip.tokenize(negative_text) - if 'l' not in tokens_p or 'g' not in tokens_p or 'l' not in tokens_n or 'g' not in tokens_n: - embeddings_final_pos, pooled_pos = advanced_encode(clip, positive_text, token_normalization, weight_interpretation, w_max=1.0, apply_to_pooled=True) - embeddings_final_neg, pooled_neg = advanced_encode(clip, negative_text, token_normalization, weight_interpretation, w_max=1.0, apply_to_pooled=True) - return ([[embeddings_final_pos, {"pooled_output": pooled_pos}]], [[embeddings_final_neg, {"pooled_output": pooled_neg}]], positive_text, negative_text, t5xxl_prompt, "", "", workflow_tuple) - else: - # tokens_p = clip.tokenize(positive_text) - if 'l' in clip.tokenize(positive_l): - tokens_p["l"] = clip.tokenize(positive_l)["l"] - if len(tokens_p["l"]) != len(tokens_p["g"]): - empty = clip.tokenize("") - while len(tokens_p["l"]) < len(tokens_p["g"]): - tokens_p["l"] += empty["l"] - while len(tokens_p["l"]) > len(tokens_p["g"]): - tokens_p["g"] += empty["g"] - - # tokens_n = clip.tokenize(negative_text) - if 'l' in clip.tokenize(negative_l): - tokens_n["l"] = clip.tokenize(negative_l)["l"] - - if len(tokens_n["l"]) != len(tokens_n["g"]): - empty = clip.tokenize("") - while len(tokens_n["l"]) < len(tokens_n["g"]): - tokens_n["l"] += empty["l"] - while len(tokens_n["l"]) > len(tokens_n["g"]): - tokens_n["g"] += empty["g"] - - cond_p, pooled_p = clip.encode_from_tokens(tokens_p, return_pooled=True) - cond_n, pooled_n = clip.encode_from_tokens(tokens_n, return_pooled=True) - - return ([[cond_p, {"pooled_output": pooled_p, "width": width, "height": height, "crop_w": 0, "crop_h": 0, "target_width": width, "target_height": height}]], [[cond_n, {"pooled_output": pooled_n, "width": width, "height": height, "crop_w": 0, "crop_h": 0, "target_width": width, "target_height": height}]], positive_text, negative_text, "", positive_l, negative_l, workflow_tuple) - - else: - tokens_pos = clip.tokenize(positive_text) - tokens_neg = clip.tokenize(negative_text) - - try: - comfy.model_management.soft_empty_cache() - except Exception: - print('No need to clear cache...') - - out_pos = clip.encode_from_tokens(tokens_pos, return_pooled=True, return_dict=True) - out_neg = clip.encode_from_tokens(tokens_neg, return_pooled=True, return_dict=True) - - cond_pos = out_pos.pop("cond") - cond_neg = out_neg.pop("cond") - - return ([[cond_pos, out_pos]], [[cond_neg, out_neg]], positive_text, negative_text, t5xxl_prompt, "", "", workflow_tuple) - + clip = nodes.CLIPLoader.load_clip(self, clip_model, 'stable_diffusion')[0] + 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) case 'SD3': - if t5xxl_prompt: - pos_out = nodes_sd3.CLIPTextEncodeSD3.execute(clip, positive_text, positive_text, t5xxl_prompt, 'none') - tokens_neg = clip.tokenize(negative_text) - out_neg = clip.encode_from_tokens(tokens_neg, return_pooled=True, return_dict=True) - cond_neg = out_neg.pop("cond") - return (pos_out[0], [[cond_neg, out_neg]], positive_text, negative_text, t5xxl_prompt, "", "", workflow_tuple) - else: - tokens_pos = clip.tokenize(positive_text) - tokens_neg = clip.tokenize(negative_text) - out_pos = clip.encode_from_tokens(tokens_pos, return_pooled=True, return_dict=True) - out_neg = clip.encode_from_tokens(tokens_neg, return_pooled=True, return_dict=True) - cond_pos = out_pos.pop("cond") - cond_neg = out_neg.pop("cond") - return ([[cond_pos, out_pos]], [[cond_neg, out_neg]], positive_text, negative_text, "", "", "", workflow_tuple) - + return clipping.encode_sd3(clip, positive_text, negative_text, t5xxl_prompt, workflow_tuple) case 'StableCascade': - positive_text = utility.DiT_cleaner(positive_text) - negative_text = utility.DiT_cleaner(negative_text) - tokens_pos = clip.tokenize(positive_text) - tokens_neg = clip.tokenize(negative_text) - cond_pos, pooled_pos = clip.encode_from_tokens(tokens_pos, return_pooled=True) - cond_neg, pooled_neg = clip.encode_from_tokens(tokens_neg, return_pooled=True) - return ([[cond_pos, {"pooled_output": pooled_pos}]], [[cond_neg, {"pooled_output": pooled_neg}]], positive_text, negative_text, "", "", "", workflow_tuple) - + return clipping.encode_stable_cascade(clip, positive_text, negative_text, workflow_tuple) case 'Flux': - FLUX_SAMPLER = concept_data.get('sampler', 'ksampler') if concept_data else 'ksampler' - FLUX_GUIDANCE = float(concept_data.get('guidance', 2.0)) if concept_data else 2.0 - if FLUX_SAMPLER == 'ksampler': - CONDITIONING_POS = nodes_flux.CLIPTextEncodeFlux.execute(clip, positive_text, t5xxl_prompt, FLUX_GUIDANCE)[0] - if concept_data is not None and float(concept_data.get('cfg', 2.0)) < 1.2: - CONDITIONING_NEG = CONDITIONING_POS - else: - CONDITIONING_NEG = nodes_flux.CLIPTextEncodeFlux.execute(clip, negative_text, negative_text, FLUX_GUIDANCE)[0] - return (CONDITIONING_POS, CONDITIONING_NEG, positive_text, negative_text, t5xxl_prompt, "", "", workflow_tuple) - else: - CONDITIONING_POS = nodes_flux.CLIPTextEncodeFlux.execute(clip, positive_text, t5xxl_prompt, FLUX_GUIDANCE)[0] - return (CONDITIONING_POS, CONDITIONING_POS, positive_text, negative_text, t5xxl_prompt, "", "", workflow_tuple) + return clipping.encode_flux(clip, positive_text, negative_text, t5xxl_prompt, concept_data, workflow_tuple) class PrimereResolution: RETURN_TYPES = ("INT", "INT", "INT", "STRING") diff --git a/components/clipping.py b/components/clipping.py index 66de535..ab56d2e 100644 --- a/components/clipping.py +++ b/components/clipping.py @@ -4,6 +4,9 @@ from comfy.sd1_clip import load_embed, ClipTokenWeightEncoder from comfy.sd1_clip import token_weights, escape_important, unescape_important from comfy import model_management import comfy +import comfy_extras.nodes_sd3 as nodes_sd3 +import comfy_extras.nodes_flux as nodes_flux +from . import utility class SDLongClipModel(torch.nn.Module, ClipTokenWeightEncoder): LAYERS = [ @@ -436,4 +439,88 @@ def HunyuanClipping(self, text, text_t5, CLIP, T5): "context_mask": clip_mask.float(), "context_t5_mask": t5_mask.float() } - ]],) \ No newline at end of file + ]],) + + +def 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_fn): + if adv_encode: + tokens_p = clip.tokenize(positive_text) + tokens_n = clip.tokenize(negative_text) + if 'l' not in tokens_p or 'g' not in tokens_p or 'l' not in tokens_n or 'g' not in tokens_n: + embeddings_final_pos, pooled_pos = advanced_encode_fn(clip, positive_text, token_normalization, weight_interpretation, w_max=1.0, apply_to_pooled=True) + embeddings_final_neg, pooled_neg = advanced_encode_fn(clip, negative_text, token_normalization, weight_interpretation, w_max=1.0, apply_to_pooled=True) + return ([[embeddings_final_pos, {"pooled_output": pooled_pos}]], [[embeddings_final_neg, {"pooled_output": pooled_neg}]], positive_text, negative_text, t5xxl_prompt, "", "", workflow_tuple) + else: + if 'l' in clip.tokenize(positive_l): + tokens_p["l"] = clip.tokenize(positive_l)["l"] + if len(tokens_p["l"]) != len(tokens_p["g"]): + empty = clip.tokenize("") + while len(tokens_p["l"]) < len(tokens_p["g"]): + tokens_p["l"] += empty["l"] + while len(tokens_p["l"]) > len(tokens_p["g"]): + tokens_p["g"] += empty["g"] + if 'l' in clip.tokenize(negative_l): + tokens_n["l"] = clip.tokenize(negative_l)["l"] + if len(tokens_n["l"]) != len(tokens_n["g"]): + empty = clip.tokenize("") + while len(tokens_n["l"]) < len(tokens_n["g"]): + tokens_n["l"] += empty["l"] + while len(tokens_n["l"]) > len(tokens_n["g"]): + tokens_n["g"] += empty["g"] + cond_p, pooled_p = clip.encode_from_tokens(tokens_p, return_pooled=True) + cond_n, pooled_n = clip.encode_from_tokens(tokens_n, return_pooled=True) + return ([[cond_p, {"pooled_output": pooled_p, "width": width, "height": height, "crop_w": 0, "crop_h": 0, "target_width": width, "target_height": height}]], [[cond_n, {"pooled_output": pooled_n, "width": width, "height": height, "crop_w": 0, "crop_h": 0, "target_width": width, "target_height": height}]], positive_text, negative_text, "", positive_l, negative_l, workflow_tuple) + else: + tokens_pos = clip.tokenize(positive_text) + tokens_neg = clip.tokenize(negative_text) + try: + comfy.model_management.soft_empty_cache() + except Exception: + pass + out_pos = clip.encode_from_tokens(tokens_pos, return_pooled=True, return_dict=True) + out_neg = clip.encode_from_tokens(tokens_neg, return_pooled=True, return_dict=True) + cond_pos = out_pos.pop("cond") + cond_neg = out_neg.pop("cond") + return ([[cond_pos, out_pos]], [[cond_neg, out_neg]], positive_text, negative_text, t5xxl_prompt, "", "", workflow_tuple) + + +def encode_sd3(clip, positive_text, negative_text, t5xxl_prompt, workflow_tuple): + if t5xxl_prompt: + pos_out = nodes_sd3.CLIPTextEncodeSD3.execute(clip, positive_text, positive_text, t5xxl_prompt, 'none') + tokens_neg = clip.tokenize(negative_text) + out_neg = clip.encode_from_tokens(tokens_neg, return_pooled=True, return_dict=True) + cond_neg = out_neg.pop("cond") + return (pos_out[0], [[cond_neg, out_neg]], positive_text, negative_text, t5xxl_prompt, "", "", workflow_tuple) + else: + tokens_pos = clip.tokenize(positive_text) + tokens_neg = clip.tokenize(negative_text) + out_pos = clip.encode_from_tokens(tokens_pos, return_pooled=True, return_dict=True) + out_neg = clip.encode_from_tokens(tokens_neg, return_pooled=True, return_dict=True) + cond_pos = out_pos.pop("cond") + cond_neg = out_neg.pop("cond") + return ([[cond_pos, out_pos]], [[cond_neg, out_neg]], positive_text, negative_text, "", "", "", workflow_tuple) + + +def encode_stable_cascade(clip, positive_text, negative_text, workflow_tuple): + positive_text = utility.DiT_cleaner(positive_text) + negative_text = utility.DiT_cleaner(negative_text) + tokens_pos = clip.tokenize(positive_text) + tokens_neg = clip.tokenize(negative_text) + cond_pos, pooled_pos = clip.encode_from_tokens(tokens_pos, return_pooled=True) + cond_neg, pooled_neg = clip.encode_from_tokens(tokens_neg, return_pooled=True) + return ([[cond_pos, {"pooled_output": pooled_pos}]], [[cond_neg, {"pooled_output": pooled_neg}]], positive_text, negative_text, "", "", "", workflow_tuple) + + +def encode_flux(clip, positive_text, negative_text, t5xxl_prompt, concept_data, workflow_tuple): + flux_sampler = concept_data.get('sampler', 'ksampler') if concept_data else 'ksampler' + flux_guidance = float(concept_data.get('guidance', 2.0)) if concept_data else 2.0 + if flux_sampler == 'ksampler': + cond_pos = nodes_flux.CLIPTextEncodeFlux.execute(clip, positive_text, t5xxl_prompt, flux_guidance)[0] + if concept_data is not None and float(concept_data.get('cfg', 2.0)) < 1.2: + cond_neg = cond_pos + else: + cond_neg = nodes_flux.CLIPTextEncodeFlux.execute(clip, negative_text, negative_text, flux_guidance)[0] + 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 diff --git a/components/models.py b/components/models.py new file mode 100644 index 0000000..c8507c8 --- /dev/null +++ b/components/models.py @@ -0,0 +1,182 @@ +import os +import comfy +import comfy.sd +import comfy.utils +import folder_paths +import nodes +import comfy_extras.nodes_sd3 as nodes_sd3 +from pathlib import Path +from .tree import PRIMERE_ROOT +from . import utility +from . import nf4_helper +from .gguf import nodes as gguf_nodes + + +def resolve_symlink(ckpt_name): + fullpathFile = folder_paths.get_full_path('checkpoints', ckpt_name) + if not os.path.islink(str(fullpathFile)): + return None, None, None + File_link = Path(str(fullpathFile)).resolve() + model_ext = os.path.splitext(File_link)[1].lower() + linkName_U = str(folder_paths.folder_names_and_paths["diffusion_models"][0][0]) + linkName_D = str(folder_paths.folder_names_and_paths["diffusion_models"][0][1]) + linkedFileName = str(File_link).replace(linkName_U + '\\', '').replace(linkName_D + '\\', '') + if str(Path(linkName_U).stem) in linkedFileName: + linkedFileName = linkedFileName.split(Path(linkName_U).stem + '\\', 1)[1] + if str(Path(linkName_D).stem) in linkedFileName: + linkedFileName = linkedFileName.split(Path(linkName_D).stem + '\\', 1)[1] + return File_link, linkedFileName, model_ext + + +def apply_lora(loader_self, model, lora_path, strength): + if not os.path.exists(lora_path) or strength == 0: + return model + lora = None + if loader_self.loaded_lora is not None: + if loader_self.loaded_lora[0] == lora_path: + lora = loader_self.loaded_lora[1] + else: + temp = loader_self.loaded_lora + loader_self.loaded_lora = None + del temp + if lora is None: + lora = comfy.utils.load_torch_file(lora_path, safe_load=True) + loader_self.loaded_lora = (lora_path, lora) + return comfy.sd.load_lora_for_models(model, None, lora, strength, 0)[0] + + +def pick_lora(concept_data): + if concept_data.get('speed_lora') == True: + return concept_data.get('speed_lora_name'), concept_data.get('speed_lora_strength', 1) + if concept_data.get('srpo_lora') == True: + return concept_data.get('srpo_lora_name'), concept_data.get('srpo_lora_strength', 1) + return None, 0 + + +def load_sd_model(loader_self, ckpt_name, use_yaml, model_config_full_path, concept_data): + if os.path.isfile(model_config_full_path) and use_yaml: + ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name) + try: + LOADED_CHECKPOINT = comfy.sd.load_checkpoint(model_config_full_path, ckpt_path, True, True, None, None, None) + except Exception: + LOADED_CHECKPOINT = nodes.CheckpointLoaderSimple.load_checkpoint(loader_self, ckpt_name) + else: + LOADED_CHECKPOINT = nodes.CheckpointLoaderSimple.load_checkpoint(loader_self, ckpt_name) + OUTPUT_MODEL = LOADED_CHECKPOINT[0] + OUTPUT_CLIP = LOADED_CHECKPOINT[1] + vae_selection = concept_data.get('vae_selection', True) + vae_name = concept_data.get('vae', None) + if not vae_selection and vae_name: + OUTPUT_VAE = utility.vae_loader_class.load_vae(vae_name)[0] + elif len(LOADED_CHECKPOINT) >= 3 and type(LOADED_CHECKPOINT[2]).__name__ == 'VAE': + OUTPUT_VAE = LOADED_CHECKPOINT[2] + elif vae_name: + OUTPUT_VAE = utility.vae_loader_class.load_vae(vae_name)[0] + else: + OUTPUT_VAE = LOADED_CHECKPOINT[2] + return OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE + + +def load_sd3_model(loader_self, ckpt_name, concept_data): + LOADED_CHECKPOINT = [] + File_link, linkedFileName, model_ext = resolve_symlink(ckpt_name) + if File_link and model_ext == '.gguf': + OUTPUT_MODEL = gguf_nodes.UnetLoaderGGUF.load_unet(loader_self, linkedFileName)[0] + else: + LOADED_CHECKPOINT = nodes.CheckpointLoaderSimple.load_checkpoint(loader_self, ckpt_name) + OUTPUT_MODEL = LOADED_CHECKPOINT[0] + clip_selection = concept_data.get('clip_selection', True) + clip_from_ckpt = len(LOADED_CHECKPOINT) >= 2 and type(LOADED_CHECKPOINT[1]).__name__ == 'CLIP' + if not clip_selection or not clip_from_ckpt: + OUTPUT_CLIP = nodes_sd3.TripleCLIPLoader.execute(concept_data.get('encoder_3'), concept_data.get('encoder_2'), concept_data.get('encoder_1'))[0] + else: + OUTPUT_CLIP = LOADED_CHECKPOINT[1] + vae_name = concept_data.get('vae', None) + if len(LOADED_CHECKPOINT) >= 3 and type(LOADED_CHECKPOINT[2]).__name__ == 'VAE': + OUTPUT_VAE = LOADED_CHECKPOINT[2] + elif vae_name: + OUTPUT_VAE = utility.vae_loader_class.load_vae(vae_name)[0] + else: + OUTPUT_VAE = LOADED_CHECKPOINT[2] + lora_name, lora_strength = pick_lora(concept_data) + if lora_name: + OUTPUT_MODEL = apply_lora(loader_self, OUTPUT_MODEL, os.path.join(PRIMERE_ROOT, 'Nodes', 'Downloads', lora_name), lora_strength) + return OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE + + +def load_stable_cascade_model(loader_self, ckpt_name, concept_data): + stage_b = concept_data.get('encoder_1', None) + cascade_clip = concept_data.get('encoder_3', None) + vae_name = concept_data.get('vae', None) + OUTPUT_CLIP = nodes.CLIPLoader.load_clip(loader_self, cascade_clip, 'stable_cascade')[0] + OUTPUT_VAE = utility.vae_loader_class.load_vae(vae_name)[0] + MODEL_B = nodes.UNETLoader.load_unet(loader_self, stage_b, 'default')[0] + File_link, linkedFileName, _ = resolve_symlink(ckpt_name) + if File_link: + MODEL_C = nodes.UNETLoader.load_unet(loader_self, linkedFileName, 'default')[0] + else: + MODEL_C = nodes.UNETLoader.load_unet(loader_self, ckpt_name, 'default')[0] + return [MODEL_B, MODEL_C], OUTPUT_CLIP, OUTPUT_VAE + + +def load_zimage_model(loader_self, ckpt_name, concept_data): + weight_dtype = concept_data.get('weight_dtype', 'default') + if 'e4m3fn' in ckpt_name: + weight_dtype = 'fp8_e4m3fn' + if 'e5m2' in ckpt_name: + weight_dtype = 'fp8_e5m2' + File_link, linkedFileName, model_ext = resolve_symlink(ckpt_name) + if File_link: + if 'diffusion_models' in str(File_link) or 'unet' in str(File_link): + if model_ext == '.gguf': + OUTPUT_MODEL = gguf_nodes.UnetLoaderGGUF.load_unet(loader_self, linkedFileName)[0] + else: + try: + OUTPUT_MODEL = nodes.UNETLoader.load_unet(loader_self, linkedFileName, weight_dtype)[0] + except Exception: + OUTPUT_MODEL = nf4_helper.UNETLoaderNF4.load_nf4unet(linkedFileName)[0] + else: + OUTPUT_MODEL = nodes.CheckpointLoaderSimple.load_checkpoint(loader_self, ckpt_name)[0] + else: + OUTPUT_MODEL = nodes.CheckpointLoaderSimple.load_checkpoint(loader_self, ckpt_name)[0] + encoder_1 = concept_data.get('encoder_1', None) + clip_ext = os.path.splitext(encoder_1)[1].lower() if encoder_1 else '' + if clip_ext == '.gguf': + OUTPUT_CLIP = gguf_nodes.CLIPLoaderGGUF.load_clip(loader_self, encoder_1, 'qwen_image')[0] + else: + OUTPUT_CLIP = nodes.CLIPLoader.load_clip(loader_self, encoder_1, 'flux2')[0] + OUTPUT_VAE = utility.vae_loader_class.load_vae(concept_data.get('vae', None))[0] + return OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE + + +def load_flux_model(loader_self, ckpt_name, concept_data): + weight_dtype = concept_data.get('weight_dtype', 'default') + is_gguf_model = False + File_link, linkedFileName, model_ext = resolve_symlink(ckpt_name) + if File_link: + if 'diffusion_models' in str(File_link) or 'unet' in str(File_link): + if model_ext == '.gguf': + OUTPUT_MODEL = gguf_nodes.UnetLoaderGGUF.load_unet(loader_self, linkedFileName)[0] + is_gguf_model = True + else: + try: + OUTPUT_MODEL = nodes.UNETLoader.load_unet(loader_self, linkedFileName, weight_dtype)[0] + except Exception: + OUTPUT_MODEL = nf4_helper.UNETLoaderNF4.load_nf4unet(linkedFileName)[0] + else: + OUTPUT_MODEL = nodes.CheckpointLoaderSimple.load_checkpoint(loader_self, linkedFileName)[0] + else: + OUTPUT_MODEL = nodes.CheckpointLoaderSimple.load_checkpoint(loader_self, ckpt_name)[0] + encoder_1 = concept_data.get('encoder_1', None) + encoder_2 = concept_data.get('encoder_2', None) + clip_ext_1 = os.path.splitext(encoder_1)[1].lower() if encoder_1 else '' + clip_ext_2 = os.path.splitext(encoder_2)[1].lower() if encoder_2 else '' + if is_gguf_model or clip_ext_1 == '.gguf' or clip_ext_2 == '.gguf': + OUTPUT_CLIP = gguf_nodes.DualCLIPLoaderGGUF.load_clip(loader_self, encoder_2, encoder_1, 'flux')[0] + else: + OUTPUT_CLIP = nodes.DualCLIPLoader.load_clip(loader_self, encoder_2, encoder_1, 'flux')[0] + OUTPUT_VAE = utility.vae_loader_class.load_vae(concept_data.get('vae', None))[0] + lora_name, lora_strength = pick_lora(concept_data) + if lora_name: + OUTPUT_MODEL = apply_lora(loader_self, OUTPUT_MODEL, os.path.join(PRIMERE_ROOT, 'Nodes', 'Downloads', lora_name), lora_strength) + return OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE