diff --git a/Nodes/Dashboard.py b/Nodes/Dashboard.py index 0067026..f2bceca 100644 --- a/Nodes/Dashboard.py +++ b/Nodes/Dashboard.py @@ -154,24 +154,34 @@ class PrimereLCMSelector: return (sampler_name, scheduler_name, steps, cfg_scale, model_concept,) class PrimereModelConceptSelector: - RETURN_TYPES = (comfy.samplers.KSampler.SAMPLERS, comfy.samplers.KSampler.SCHEDULERS, "INT", "FLOAT", "STRING", "STRING", "INT") - RETURN_NAMES = ("SAMPLER_NAME", "SCHEDULER_NAME", "STEPS", "CFG", "MODEL_CONCEPT", "LIGHTNING_SELECTOR", "LIGHTNING_MODEL_STEP") + RETURN_TYPES = (comfy.samplers.KSampler.SAMPLERS, comfy.samplers.KSampler.SCHEDULERS, "INT", "FLOAT", "STRING", "STRING", "INT", "STRING", "STRING", "STRING", "STRING") + RETURN_NAMES = ("SAMPLER_NAME", "SCHEDULER_NAME", "STEPS", "CFG", "MODEL_CONCEPT", "LIGHTNING_SELECTOR", "LIGHTNING_MODEL_STEP", "CASCADE_STAGE_A", "CASCADE_STAGE_B", "CASCADE_STAGE_C", "CASCADE_CLIP") FUNCTION = "select_model_concept" CATEGORY = TREE_DASHBOARD + UNETLIST = folder_paths.get_filename_list("unet") + VAELIST = folder_paths.get_filename_list("vae") + CLIPLIST = folder_paths.get_filename_list("clip") + @classmethod - def INPUT_TYPES(s): + def INPUT_TYPES(cls): return { "required": { - "model_concept":(["Normal", "LCM", "Turbo", "Cascade", "Lightning"], {"default": "Normal"}), - "lightning_selector": (["UNET", "LORA", "SAFETENSOR", "CUSTOM"], {"default": "SAFETENSOR"}), - "lightning_model_step": ([1, 2, 4, 8], {"default": 8}), - "lightning_sampler": ("BOOLEAN", {"default": False, "label_on": "Set by model", "label_off": "Custom (external)"}), - "normal_sampler_name": (comfy.samplers.KSampler.SAMPLERS, {"forceInput": True, "default": "euler"}), "normal_scheduler_name": (comfy.samplers.KSampler.SCHEDULERS, {"forceInput": True, "default": "normal"}), "normal_cfg_scale": ('FLOAT', {"forceInput": True, "default": 7}), "normal_steps": ('INT', {"forceInput": True, "default": 12}), + + "model_concept": (["Normal", "LCM", "Turbo", "Cascade", "Lightning"], {"default": "Normal"}), + "lightning_selector": (["UNET", "LORA", "SAFETENSOR", "CUSTOM"], {"default": "SAFETENSOR"}), + "lightning_model_step": ([1, 2, 4, 8], {"default": 8}), + "lightning_sampler": ( + "BOOLEAN", {"default": False, "label_on": "Set by model", "label_off": "Custom (external)"}), + + "cascade_stage_a": (cls.VAELIST,), + "cascade_stage_b": (cls.UNETLIST,), + "cascade_stage_c": (cls.UNETLIST,), + "cascade_clip": (cls.CLIPLIST,), }, "optional": { "lcm_sampler_name": (comfy.samplers.KSampler.SAMPLERS, {"forceInput": True, "default": "lcm"}), @@ -196,7 +206,8 @@ class PrimereModelConceptSelector: } } - def select_model_concept(self, model_concept = 'Normal', lightning_selector = "SAFETENSOR", lightning_model_step = 8, lightning_sampler = False, + def select_model_concept(self, cascade_stage_a, cascade_stage_b, cascade_stage_c, cascade_clip, + model_concept = 'Normal', lightning_selector = "SAFETENSOR", lightning_model_step = 8, lightning_sampler = False, normal_sampler_name = 'euler', normal_scheduler_name = 'normal', normal_cfg_scale = 7, normal_steps = 12, lcm_sampler_name = 'lcm', lcm_scheduler_name = 'sgm_uniform', lcm_cfg_scale = 1.2, lcm_steps = 6, turbo_sampler_name = 'dpmpp_sde', turbo_scheduler_name = "karras", turbo_cfg_scale = 1.15, turbo_steps = 2, @@ -244,7 +255,13 @@ class PrimereModelConceptSelector: lightning_selector = None lightning_model_step = None - return (sampler_name, scheduler_name, steps, round(cfg_scale, 2), model_concept, lightning_selector, lightning_model_step) + if model_concept != 'Cascade': + cascade_stage_a = None + cascade_stage_b = None + cascade_stage_c = None + cascade_clip = None + + return (sampler_name, scheduler_name, steps, round(cfg_scale, 2), model_concept, lightning_selector, lightning_model_step, cascade_stage_a, cascade_stage_b, cascade_stage_c, cascade_clip) class PrimereCKPTLoader: RETURN_TYPES = ("MODEL", "CLIP", "VAE", "STRING",) @@ -256,7 +273,7 @@ class PrimereCKPTLoader: self.loaded_lora = None @classmethod - def INPUT_TYPES(s): + def INPUT_TYPES(cls): return { "required": { "ckpt_name": ("CHECKPOINT_NAME",), @@ -274,13 +291,41 @@ class PrimereCKPTLoader: }, } - def load_primere_ckpt(self, ckpt_name, use_yaml, strength_lcm_model, strength_lcm_clip, model_concept = "Normal", concept_data = None, lightning_selector = 'SAFETENSOR', lightning_model_step = 8, loaded_model = None, loaded_clip = None, loaded_vae = None): + def load_primere_ckpt(self, ckpt_name, use_yaml, strength_lcm_model, strength_lcm_clip, + model_concept = "Normal", concept_data = None, + lightning_selector = 'SAFETENSOR', lightning_model_step = 8, + cascade_stage_a = None, cascade_stage_b = None, cascade_stage_c = None, cascade_clip = None, + loaded_model = None, loaded_clip = None, loaded_vae = None): + if concept_data is not None: if 'lightning_selector' in concept_data: lightning_selector = concept_data['lightning_selector'] if 'lightning_model_step' in concept_data: lightning_model_step = concept_data['lightning_model_step'] + if 'cascade_stage_a' in concept_data: + cascade_stage_a = concept_data['cascade_stage_a'] + if 'cascade_stage_b' in concept_data: + cascade_stage_b = concept_data['cascade_stage_b'] + if 'cascade_stage_c' in concept_data: + cascade_stage_c = concept_data['cascade_stage_c'] + if 'cascade_clip' in concept_data: + cascade_clip = concept_data['cascade_clip'] + + if model_concept == "Cascade" and cascade_stage_a is not None and cascade_stage_b is not None and cascade_stage_c is not None and cascade_clip is not None: + MODEL_VERSION = 'SDXL_2048' + is_sdxl = 1 + + OUTPUT_VAE = nodes.VAELoader.load_vae(self, cascade_stage_a)[0] + + MODEL_B = nodes.UNETLoader.load_unet(self, cascade_stage_b)[0] + MODEL_C = nodes.UNETLoader.load_unet(self, cascade_stage_c)[0] + + OUTPUT_CLIP = nodes.CLIPLoader.load_clip(self, cascade_clip, 'stable_cascade')[0] + OUTPUT_MODEL = [MODEL_B, MODEL_C] + + return (OUTPUT_MODEL,) + (OUTPUT_CLIP,) + (OUTPUT_VAE,) + (MODEL_VERSION,) + ModelConceptChanges = utility.ModelConceptNames(ckpt_name, model_concept, lightning_selector, lightning_model_step) ckpt_name = ModelConceptChanges['ckpt_name'] lora_name = ModelConceptChanges['lora_name'] @@ -734,44 +779,60 @@ class PrimereCLIP: if (adv_encode == True): if (is_sdxl == 0): + + if model_concept == 'Cascade': + positive_text = utility.clear_cascade(positive_text) + negative_text = utility.clear_cascade(negative_text) + 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, "", "") else: # embeddings_final_pos, pooled_pos = advanced_encode_XL(clip, sdxl_positive_l, positive_text, token_normalization, weight_interpretation, w_max = 1.0, clip_balance = sdxl_balance_l, apply_to_pooled = True) # embeddings_final_neg, pooled_neg = advanced_encode_XL(clip, sdxl_negative_l, negative_text, token_normalization, weight_interpretation, w_max = 1.0, clip_balance = sdxl_balance_l, apply_to_pooled = True) # return ([[embeddings_final_pos, {"pooled_output": pooled_pos}]],[[embeddings_final_neg, {"pooled_output": pooled_neg}]], positive_text, negative_text, sdxl_positive_l, sdxl_negative_l) + if model_concept == 'Cascade': + positive_text = utility.clear_cascade(positive_text) + negative_text = utility.clear_cascade(negative_text) tokens_p = clip.tokenize(positive_text) - tokens_p["l"] = clip.tokenize(sdxl_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(sdxl_positive_l): + tokens_p["l"] = clip.tokenize(sdxl_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) - tokens_n["l"] = clip.tokenize(sdxl_negative_l)["l"] + if 'l' in clip.tokenize(sdxl_negative_l): + tokens_n["l"] = clip.tokenize(sdxl_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"] + 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, sdxl_positive_l, sdxl_negative_l) else: + if model_concept == 'Cascade': + positive_text = utility.clear_cascade(positive_text) + negative_text = utility.clear_cascade(negative_text) + tokens = clip.tokenize(positive_text) cond_pos, pooled_pos = clip.encode_from_tokens(tokens, return_pooled = True) tokens = clip.tokenize(negative_text) cond_neg, pooled_neg = clip.encode_from_tokens(tokens, return_pooled = True) + return ([[cond_pos, {"pooled_output": pooled_pos}]], [[cond_neg, {"pooled_output": pooled_neg}]], positive_text, negative_text, "", "") class PrimereResolution: @@ -1075,9 +1136,12 @@ class PrimereClearPrompt: "remove_lycoris": ("BOOLEAN", {"default": False}), "remove_hypernetwork": ("BOOLEAN", {"default": False}), }, + "optional": { + "model_concept": ("STRING", {"default": "Normal", "forceInput": True}), + } } - def clean_prompt(self, positive_prompt, negative_prompt, remove_comfy_embedding, remove_a1111_embedding, remove_lora, remove_lycoris, remove_hypernetwork, remove_only_if_sdxl, model_version = 'BaseModel_1024'): + def clean_prompt(self, positive_prompt, negative_prompt, remove_comfy_embedding, remove_a1111_embedding, remove_lora, remove_lycoris, remove_hypernetwork, remove_only_if_sdxl, model_version = 'BaseModel_1024', model_concept = "Normal"): NETWORK_START = [] is_sdxl = 0 @@ -1113,9 +1177,13 @@ class PrimereClearPrompt: negative_prompt = re.sub(r'(, )\1+', r', ', negative_prompt).strip(', ').replace(' ,', ',') if len(NETWORK_START) > 0: - NETWORK_END = ['\n', '>', ' ', ',', '}', ')', '|'] + NETWORK_START - positive_prompt = utility.clear_prompt(NETWORK_START, NETWORK_END, positive_prompt) - negative_prompt = utility.clear_prompt(NETWORK_START, NETWORK_END, negative_prompt) + NETWORK_END = ['\n', '>', ' ', ',', '}', ')', '|'] + NETWORK_START + positive_prompt = utility.clear_prompt(NETWORK_START, NETWORK_END, positive_prompt) + negative_prompt = utility.clear_prompt(NETWORK_START, NETWORK_END, negative_prompt) + + if model_concept == 'Cascade': + positive_prompt = utility.clear_cascade(positive_prompt) + negative_prompt = utility.clear_cascade(negative_prompt) return (positive_prompt, negative_prompt,) @@ -1359,11 +1427,16 @@ class PrimereConceptDataTuple: CATEGORY = TREE_DASHBOARD @classmethod - def INPUT_TYPES(s): + def INPUT_TYPES(cls): return { "required": { "lightning_selector": ("STRING", {"default": "SAFETENSOR", "forceInput": True}), "lightning_model_step": ("INT", {"default": 8, "forceInput": True}), + + "cascade_stage_a": ("STRING", {"forceInput": True}), + "cascade_stage_b": ("STRING", {"forceInput": True}), + "cascade_stage_c": ("STRING", {"forceInput": True}), + "cascade_clip": ("STRING", {"forceInput": True}), }, } diff --git a/Nodes/Inputs.py b/Nodes/Inputs.py index d0bc39f..37e9073 100644 --- a/Nodes/Inputs.py +++ b/Nodes/Inputs.py @@ -358,15 +358,20 @@ class PrimereVAESelector: "required": { "vae_sd": ("VAE",), "vae_sdxl": ("VAE",), + "vae_cascade": ("VAE",), "model_version": ("STRING", {"default": 'BaseModel_1024', "forceInput": True}), + "model_concept": ("STRING", {"default": "Normal", "forceInput": True}), } } - def primere_vae_selector(self, vae_sd, vae_sdxl, model_version = "BaseModel_1024"): - if model_version == 'SDXL_2048': - return (vae_sdxl, ) - else: - return (vae_sd, ) + def primere_vae_selector(self, vae_sd, vae_sdxl, vae_cascade, model_version = "BaseModel_1024", model_concept = 'Normal'): + match model_concept: + case 'Cascade': + return (vae_cascade,) + match model_version: + case 'SDXL_2048': + return (vae_sdxl,) + return (vae_sd,) class PrimereMetaRead: CATEGORY = TREE_INPUTS diff --git a/Nodes/Outputs.py b/Nodes/Outputs.py index 0ed142c..6dc22f2 100644 --- a/Nodes/Outputs.py +++ b/Nodes/Outputs.py @@ -14,6 +14,9 @@ import comfy.samplers from .modules import exif_data_checker from nodes import common_ksampler import comfy_extras.nodes_custom_sampler as nodes_custom_sampler +import comfy_extras.nodes_stable_cascade as nodes_stable_cascade +import torch +from ..components import utility ALLOWED_EXT = ('.jpeg', '.jpg', '.png', '.tiff', '.gif', '.bmp', '.webp') @@ -467,6 +470,30 @@ class PrimereKSampler: sampler = comfy.samplers.sampler_object(sampler_name) turbo_samples = nodes_custom_sampler.SamplerCustom().sample(model, True, seed, cfg, positive, negative, sampler, sigmas[0], latent_image) samples = (turbo_samples[0],) + if (model_concept == "Cascade"): + if type(model).__name__ == 'list': + latent_size = utility.getLatentSize(latent_image) + if (latent_size[0] < latent_size[1]): + orientation = 'Vertical' + else: + orientation = 'Horizontal' + + cascade_standards = [42, 84, 126, 168, 210, 252, 294, 336, 378, 420, 462, 504, 546, 588, 630, 672, 714, 756, 798, 840, 882, 924, 966, 1008, 1050, 1092, 1134, 1176, 1218, 1260, 1302, 1344, 1386, 1428, 1470, 1512, 1554, 1596, 1638, 1680] + dimensions = utility.calculate_dimensions(self, 'Square [1:1]', orientation, True, 'SDXL_2048', True, latent_size[0], latent_size[1], cascade_standards) + dimension_x = dimensions[0] + dimension_y = dimensions[1] + + height = dimension_y + width = dimension_x + compression = 42 + if type(model[0]).__name__ == 'ModelPatcher' and type(model[1]).__name__ == 'ModelPatcher': + c_latent = {"samples": torch.zeros([1, 16, height // compression, width // compression])} + b_latent = {"samples": torch.zeros([1, 4, height // 4, width // 4])} + samples_c = common_ksampler(model[1], seed, steps, cfg, sampler_name, scheduler_name, positive, negative, c_latent, denoise=denoise)[0] + conditining_c = nodes_stable_cascade.StableCascade_StageB_Conditioning.set_prior(self, positive, samples_c)[0] + samples = common_ksampler(model[0], seed, 10, 1.00, sampler_name, scheduler_name, conditining_c, negative, b_latent, denoise=denoise) + else: + samples = latent_image else: samples = common_ksampler(model, seed, steps, cfg, sampler_name, scheduler_name, positive, negative, latent_image, denoise=denoise) diff --git a/components/utility.py b/components/utility.py index 7fe7be9..c15e217 100644 --- a/components/utility.py +++ b/components/utility.py @@ -44,7 +44,10 @@ def remove_quotes(string): def add_quotes(string): return '"' + str(string) + '"' -def calculate_dimensions(self, ratio: str, orientation: str, round_to_standard: bool, model_version: str, calculate_by_custom: bool, custom_side_a: float, custom_side_b: float): +def calculate_dimensions(self, ratio: str, orientation: str, round_to_standard: bool, model_version: str, calculate_by_custom: bool, custom_side_a: float, custom_side_b: float, custom_standards=None): + if custom_standards is None: + custom_standards = [] + DEFAULT_RES = 768 match model_version: @@ -81,7 +84,11 @@ def calculate_dimensions(self, ratio: str, orientation: str, round_to_standard: side_base = round(math.sqrt(result_y)) side_a = round(ratio_1 * side_base) if round_to_standard == True: - side_a = min(STANDARD_SIDES, key=lambda x: abs(side_a - x)) + if len(custom_standards) > 1 and custom_standards is not None: + side_a = min(custom_standards, key=lambda x: abs(side_a - x)) + else: + side_a = min(STANDARD_SIDES, key=lambda x: abs(side_a - x)) + side_b = round(FullPixels / side_a) return sorted([side_a, side_b], reverse=True) @@ -142,6 +149,9 @@ def clear_prompt(NETWORK_START, NETWORK_END, promptstring): return promptstring_temp.replace('()', '').replace(' , ,', ',').replace('||', '').replace('{,', '').replace(' ', ' ').replace(', ,', ',').strip(', ') +def clear_cascade(prompt): + return re.sub("(\.\d+)|(:\d+)|[()]|BREAK|break", "", prompt).replace(' ', ' ') + def get_networks_prompt(NETWORK_START, NETWORK_END, promptstring): valid_networks = [] @@ -657,3 +667,15 @@ def LightningConceptModel(self, model_concept, lightningModeValid, lightning_sel OUTPUT_MODEL = nodes_model_advanced.ModelSamplingDiscrete.patch(self, OUTPUT_MODEL, "x0", False)[0] return OUTPUT_MODEL + +def getLatentSize(samples): + for tensor in samples['samples'][0]: + if isinstance(tensor, torch.Tensor): + shape = tensor.shape + tensor_height = shape[-2] + tensor_width = shape[-1] + return (tensor_width, tensor_height) + else: + return (None, None) + + return (None, None)