From 16d6d5d97e06ad9e984865489018da4073fdcc60 Mon Sep 17 00:00:00 2001 From: "DESKTOP-TVBJISQ\\Primere" Date: Mon, 16 Mar 2026 09:23:35 +0100 Subject: [PATCH] V 2.0.0 - Auto config #24 - refiners for all --- Nodes/Outputs.py | 16 +++++ components/clipping.py | 112 ++++++++++++++++++++++++---------- components/models.py | 65 ++++++++++++-------- components/primeresamplers.py | 63 ++++++++++--------- 4 files changed, 164 insertions(+), 92 deletions(-) diff --git a/Nodes/Outputs.py b/Nodes/Outputs.py index c375eac..a6e72fd 100644 --- a/Nodes/Outputs.py +++ b/Nodes/Outputs.py @@ -546,6 +546,19 @@ class PrimereKSampler: noise_constant = noise_extender_ksampler WORKFLOWDATA = extra_pnginfo['workflow']['nodes'] + refiner_model_data = None + if isinstance(model, dict) and 'main' in model: + refiner_model_data = model.get('refiner') + model = model['main'] + refiner_cond_pos = None + refiner_cond_neg = None + if isinstance(positive, dict) and 'main' in positive: + refiner_cond_pos = positive.get('refiner') + positive = positive['main'] + if isinstance(negative, dict) and 'main' in negative: + refiner_cond_neg = negative.get('refiner') + negative = negative['main'] + match model_concept: case 'SANA1024' | 'SANA512': if scheduler_name == 'flow_dpm-solver': @@ -670,6 +683,9 @@ class PrimereKSampler: variation_extender, variation_batch_step_original, batch_counter, variation_extender_original, variation_batch_step, variation_level, variation_limit, align_your_steps, noise_extender_ksampler, None, control_data)[0] + if refiner_model_data is not None: + samples_out = primeresamplers._run_refiner_pass(self, refiner_model_data, refiner_cond_pos, refiner_cond_neg, samples_out, control_data, seed) + if control_data is not None: control_data['sampler_settings'] = {} control_data['sampler_settings']['denoise'] = denoise diff --git a/components/clipping.py b/components/clipping.py index 2b37401..76f4e3e 100644 --- a/components/clipping.py +++ b/components/clipping.py @@ -600,6 +600,10 @@ def detect_attn_preset(model_name, default='Off'): def apply_clip_overrides(loader_self, clip, control_data): if not control_data: return clip + refiner_clip = None + if isinstance(clip, dict): + refiner_clip = clip.get('refiner') + clip = clip['main'] encoder_1 = control_data.get('encoder_1', None) last_layer = int(control_data.get('last_layer', 0)) baked_clip = clip @@ -623,6 +627,8 @@ def apply_clip_overrides(loader_self, clip, control_data): if last_layer < 0: clip = nodes.CLIPSetLastLayer.set_last_layer(loader_self, clip, last_layer)[0] + if refiner_clip is not None: + return {'main': clip, 'refiner': refiner_clip} return clip @@ -635,20 +641,45 @@ def apply_clip_attention_multiply(clip, control_data): out = float(control_data.get('clip_attn_out', 1.0)) if q == 1.0 and k == 1.0 and v == 1.0 and out == 1.0: return clip + if isinstance(clip, dict): + try: + clip['main'] = nodes_attention_multiply.CLIPAttentionMultiply.execute(clip['main'], q, k, v, out)[0] + except Exception: + pass + return clip try: return nodes_attention_multiply.CLIPAttentionMultiply.execute(clip, q, k, v, out)[0] except Exception: return clip +def _maybe_wrap_cond(pos_cond, neg_cond, refiner_clip, positive_text, negative_text): + if refiner_clip is None: + return pos_cond, neg_cond + tokens_pos = refiner_clip.tokenize(positive_text) + tokens_neg = refiner_clip.tokenize(negative_text) + out_pos = refiner_clip.encode_from_tokens(tokens_pos, return_pooled=True, return_dict=True) + out_neg = refiner_clip.encode_from_tokens(tokens_neg, return_pooled=True, return_dict=True) + cond_pos_ref = out_pos.pop("cond") + cond_neg_ref = out_neg.pop("cond") + return {'main': pos_cond, 'refiner': [[cond_pos_ref, out_pos]]}, {'main': neg_cond, 'refiner': [[cond_neg_ref, out_neg]]} + + def encode_standard(clip, positive_text, negative_text, t5xxl_prompt, adv_encode, token_normalization, weight_interpretation, positive_l, negative_l, width, height, control_data, advanced_encode_fn): + refiner_clip = None + if isinstance(clip, dict): + refiner_clip = clip.get('refiner') + clip = clip['main'] 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, "", "", control_data) + pos_cond = [[embeddings_final_pos, {"pooled_output": pooled_pos}]] + neg_cond = [[embeddings_final_neg, {"pooled_output": pooled_neg}]] + pos_cond, neg_cond = _maybe_wrap_cond(pos_cond, neg_cond, refiner_clip, positive_text, negative_text) + return (pos_cond, neg_cond, positive_text, negative_text, t5xxl_prompt, "", "", control_data) else: if 'l' in clip.tokenize(positive_l): tokens_p["l"] = clip.tokenize(positive_l)["l"] @@ -668,7 +699,10 @@ def encode_standard(clip, positive_text, negative_text, t5xxl_prompt, adv_encode 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, control_data) + pos_cond = [[cond_p, {"pooled_output": pooled_p, "width": width, "height": height, "crop_w": 0, "crop_h": 0, "target_width": width, "target_height": height}]] + neg_cond = [[cond_n, {"pooled_output": pooled_n, "width": width, "height": height, "crop_w": 0, "crop_h": 0, "target_width": width, "target_height": height}]] + pos_cond, neg_cond = _maybe_wrap_cond(pos_cond, neg_cond, refiner_clip, positive_text, negative_text) + return (pos_cond, neg_cond, positive_text, negative_text, "", positive_l, negative_l, control_data) else: tokens_pos = clip.tokenize(positive_text) tokens_neg = clip.tokenize(negative_text) @@ -680,16 +714,22 @@ def encode_standard(clip, positive_text, negative_text, t5xxl_prompt, adv_encode 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, "", "", control_data) + pos_cond, neg_cond = _maybe_wrap_cond([[cond_pos, out_pos]], [[cond_neg, out_neg]], refiner_clip, positive_text, negative_text) + return (pos_cond, neg_cond, positive_text, negative_text, t5xxl_prompt, "", "", control_data) def encode_sd3(clip, positive_text, negative_text, t5xxl_prompt, control_data): + refiner_clip = None + if isinstance(clip, dict): + refiner_clip = clip.get('refiner') + clip = clip['main'] 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, "", "", control_data) + pos_cond, neg_cond = _maybe_wrap_cond(pos_out[0], [[cond_neg, out_neg]], refiner_clip, positive_text, negative_text) + return (pos_cond, neg_cond, positive_text, negative_text, t5xxl_prompt, "", "", control_data) else: tokens_pos = clip.tokenize(positive_text) tokens_neg = clip.tokenize(negative_text) @@ -697,68 +737,74 @@ def encode_sd3(clip, positive_text, negative_text, t5xxl_prompt, control_data): 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, "", "", "", control_data) + pos_cond, neg_cond = _maybe_wrap_cond([[cond_pos, out_pos]], [[cond_neg, out_neg]], refiner_clip, positive_text, negative_text) + return (pos_cond, neg_cond, positive_text, negative_text, "", "", "", control_data) def encode_stable_cascade(clip, positive_text, negative_text, control_data): + refiner_clip = None + if isinstance(clip, dict): + refiner_clip = clip.get('refiner') + clip = clip['main'] 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, "", "", "", control_data) + pos_cond, neg_cond = _maybe_wrap_cond([[cond_pos, {"pooled_output": pooled_pos}]], [[cond_neg, {"pooled_output": pooled_neg}]], refiner_clip, positive_text, negative_text) + return (pos_cond, neg_cond, positive_text, negative_text, "", "", "", control_data) def encode_pixart_sigma(clip, positive_text, negative_text, control_data): + refiner_clip = clip.get('refiner') if isinstance(clip, dict) else None + clip = clip['main'] if isinstance(clip, dict) else clip positive_text = utility.DiT_cleaner(positive_text) negative_text = utility.DiT_cleaner(negative_text) - - cond_pos_ref = cond_neg_ref = out_pos_ref = out_neg_ref = None - - if clip['refiner'] is not None: - clipRef = clip['refiner'] - tokens_pos_ref = clipRef.tokenize(positive_text) - tokens_neg_ref = clipRef.tokenize(negative_text) - out_pos_ref = clipRef.encode_from_tokens(tokens_pos_ref, return_pooled=True, return_dict=True) - out_neg_ref = clipRef.encode_from_tokens(tokens_neg_ref, return_pooled=True, return_dict=True) - cond_pos_ref = out_pos_ref.pop("cond") - cond_neg_ref = out_neg_ref.pop("cond") - - clipMain = clip['main'] - tokens_pos_main = clipMain.tokenize(positive_text) - tokens_neg_main = clipMain.tokenize(negative_text) - out_pos_main = clipMain.encode_from_tokens(tokens_pos_main, return_pooled=True, return_dict=True) - out_neg_main = clipMain.encode_from_tokens(tokens_neg_main, return_pooled=True, return_dict=True) - cond_pos_main = out_pos_main.pop("cond") - cond_neg_main = out_neg_main.pop("cond") - - return ({'refiner': [[cond_pos_ref, out_pos_ref]], 'main': [[cond_pos_main, out_pos_main]]}, {'refiner': [[cond_neg_ref, out_neg_ref]], 'main': [[cond_neg_main, out_neg_main]]}, positive_text, negative_text, "", "", "", control_data) - - -def encode_chroma(clip, positive_text, negative_text, control_data): 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, "", "", "", control_data) + pos_cond, neg_cond = _maybe_wrap_cond([[cond_pos, out_pos]], [[cond_neg, out_neg]], refiner_clip, positive_text, negative_text) + return (pos_cond, neg_cond, positive_text, negative_text, "", "", "", control_data) + + +def encode_chroma(clip, positive_text, negative_text, control_data): + refiner_clip = None + if isinstance(clip, dict): + refiner_clip = clip.get('refiner') + clip = clip['main'] + 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") + pos_cond, neg_cond = _maybe_wrap_cond([[cond_pos, out_pos]], [[cond_neg, out_neg]], refiner_clip, positive_text, negative_text) + return (pos_cond, neg_cond, positive_text, negative_text, "", "", "", control_data) def encode_flux(clip, positive_text, negative_text, t5xxl_prompt, control_data): + refiner_clip = None + if isinstance(clip, dict): + refiner_clip = clip.get('refiner') + clip = clip['main'] FLUX_SAMPLER = control_data.get('sampler', 'ksampler') FLUX_GUIDANCE = control_data.get('guidance', 2) if FLUX_SAMPLER == 'custom_advanced' and len(t5xxl_prompt) > 5: 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, "", "", control_data) + pos_cond, neg_cond = _maybe_wrap_cond(CONDITIONING_POS, CONDITIONING_POS, refiner_clip, positive_text, negative_text) + return (pos_cond, neg_cond, positive_text, negative_text, t5xxl_prompt, "", "", control_data) 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, t5xxl_prompt, "", "", control_data) + pos_cond, neg_cond = _maybe_wrap_cond([[cond_pos, out_pos]], [[cond_neg, out_neg]], refiner_clip, positive_text, negative_text) + return (pos_cond, neg_cond, positive_text, negative_text, t5xxl_prompt, "", "", control_data) _SANA_MAX_TOKENS = 300 diff --git a/components/models.py b/components/models.py index 84e88fe..ac1d940 100644 --- a/components/models.py +++ b/components/models.py @@ -132,6 +132,23 @@ def pick_lora(concept_data): return None, 0 +def _load_refiner(loader_self, concept_data): + if concept_data.get('refiner') != True: + return None, None, None + refiner_model_name = concept_data.get('refiner_model') + if not refiner_model_name or refiner_model_name == 'None': + return None, None, None + ckpt = nodes.CheckpointLoaderSimple.load_checkpoint(loader_self, refiner_model_name) + return ckpt[0], ckpt[1], ckpt[2] + + +def _wrap_refiner(output_model, output_clip, output_vae, loader_self, concept_data): + ref_model, ref_clip, ref_vae = _load_refiner(loader_self, concept_data) + if ref_model is None: + return output_model, output_clip, output_vae + return {'main': output_model, 'refiner': ref_model}, {'main': output_clip, 'refiner': ref_clip}, output_vae + + 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) @@ -154,7 +171,7 @@ def load_sd_model(loader_self, ckpt_name, use_yaml, model_config_full_path, conc else: OUTPUT_VAE = LOADED_CHECKPOINT[2] OUTPUT_MODEL = apply_generic_patches(loader_self, OUTPUT_MODEL, concept_data) - return OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE + return _wrap_refiner(OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE, loader_self, concept_data) def load_sd3_model(loader_self, ckpt_name, concept_data): @@ -184,7 +201,7 @@ def load_sd3_model(loader_self, ckpt_name, concept_data): if lora_path: OUTPUT_MODEL = apply_lora(loader_self, OUTPUT_MODEL, lora_path, lora_strength) OUTPUT_MODEL = apply_generic_patches(loader_self, OUTPUT_MODEL, concept_data) - return OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE + return _wrap_refiner(OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE, loader_self, concept_data) def load_stable_cascade_model(loader_self, ckpt_name, concept_data): @@ -199,7 +216,7 @@ def load_stable_cascade_model(loader_self, ckpt_name, concept_data): 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 + return _wrap_refiner([MODEL_B, MODEL_C], OUTPUT_CLIP, OUTPUT_VAE, loader_self, concept_data) def load_zimage_model(loader_self, ckpt_name, concept_data): @@ -229,7 +246,7 @@ def load_zimage_model(loader_self, ckpt_name, concept_data): 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 + return _wrap_refiner(OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE, loader_self, concept_data) def load_flux_model(loader_self, ckpt_name, concept_data): @@ -272,7 +289,7 @@ def load_flux_model(loader_self, ckpt_name, concept_data): except Exception as e: print(f"Primere: ModelSamplingFlux failed: {e}") OUTPUT_MODEL = apply_generic_patches(loader_self, OUTPUT_MODEL, concept_data) - return OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE + return _wrap_refiner(OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE, loader_self, concept_data) def load_auraflow_model(loader_self, ckpt_name, concept_data): @@ -287,7 +304,7 @@ def load_auraflow_model(loader_self, ckpt_name, concept_data): encoder_1 = concept_data.get('encoder_1', None) OUTPUT_CLIP = nodes.CLIPLoader.load_clip(loader_self, encoder_1, 'stable_diffusion')[0] OUTPUT_VAE = utility.vae_loader_class.load_vae(concept_data.get('vae', None))[0] - return OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE + return _wrap_refiner(OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE, loader_self, concept_data) def load_pixart_model(loader_self, ckpt_name, concept_data): @@ -309,19 +326,8 @@ def load_pixart_model(loader_self, ckpt_name, concept_data): encoder_1 = concept_data.get('encoder_1', None) OUTPUT_CLIP_MAIN = nodes.CLIPLoader.load_clip(loader_self, encoder_1, 'sd3')[0] - - refiner_model = concept_data.get('refiner_model', None) - OUTPUT_MODEL_REFINER = None - OUTPUT_CLIP_REFINER = None - if concept_data.get('refiner') == True and refiner_model and refiner_model != 'None': - REFINER_CHECKPOINT = nodes.CheckpointLoaderSimple.load_checkpoint(loader_self, refiner_model) - OUTPUT_MODEL_REFINER = REFINER_CHECKPOINT[0] - OUTPUT_CLIP_REFINER = REFINER_CHECKPOINT[1] - OUTPUT_VAE = REFINER_CHECKPOINT[2] - else: - OUTPUT_VAE = utility.vae_loader_class.load_vae(concept_data.get('vae', None))[0] - - return {'main': OUTPUT_MODEL_MAIN, 'refiner': OUTPUT_MODEL_REFINER}, {'main': OUTPUT_CLIP_MAIN, 'refiner': OUTPUT_CLIP_REFINER}, OUTPUT_VAE + OUTPUT_VAE = utility.vae_loader_class.load_vae(concept_data.get('vae', None))[0] + return _wrap_refiner(OUTPUT_MODEL_MAIN, OUTPUT_CLIP_MAIN, OUTPUT_VAE, loader_self, concept_data) def load_playground_model(loader_self, ckpt_name, use_yaml, model_config_full_path, concept_data): @@ -329,10 +335,15 @@ def load_playground_model(loader_self, ckpt_name, use_yaml, model_config_full_pa sigma_max = concept_data.get('sigma_max', 120) sigma_min = concept_data.get('sigma_min', 0.002) edm_sampling = concept_data.get('edm_sampling', 'edm_playground_v2.5') + main_model = OUTPUT_MODEL['main'] if isinstance(OUTPUT_MODEL, dict) else OUTPUT_MODEL try: - OUTPUT_MODEL = nodes_model_advanced.ModelSamplingContinuousEDM.patch(loader_self, OUTPUT_MODEL, edm_sampling, sigma_max, sigma_min)[0] + main_model = nodes_model_advanced.ModelSamplingContinuousEDM.patch(loader_self, main_model, edm_sampling, sigma_max, sigma_min)[0] except Exception as e: print(f"Primere: ModelSamplingContinuousEDM failed: {e}") + if isinstance(OUTPUT_MODEL, dict): + OUTPUT_MODEL['main'] = main_model + else: + OUTPUT_MODEL = main_model return OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE @@ -365,7 +376,7 @@ def load_lightning_hyper_model(loader_self, ckpt_name, concept_data): OUTPUT_MODEL = utility.BDanceConceptHelper(loader_self, model_concept, True, 'LORA', None, OUTPUT_MODEL, lora_path, None, None, lora_strength) OUTPUT_MODEL = apply_generic_patches(loader_self, OUTPUT_MODEL, concept_data) - return OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE + return _wrap_refiner(OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE, loader_self, concept_data) def load_lcm_model(loader_self, ckpt_name, concept_data): @@ -388,7 +399,7 @@ def load_lcm_model(loader_self, ckpt_name, concept_data): m.add_object_patch("model_sampling", ModelSamplingAdvanced()) OUTPUT_MODEL = m - return OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE + return _wrap_refiner(OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE, loader_self, concept_data) def load_sana_model(loader_self, ckpt_name, concept_data): @@ -513,7 +524,7 @@ def load_sana_model(loader_self, ckpt_name, concept_data): "text_encoder_model": text_encoder_model, } - return SANA_MODEL, SANA_CLIP, SANA_VAE + return _wrap_refiner(SANA_MODEL, SANA_CLIP, SANA_VAE, loader_self, concept_data) def load_kolors_model(loader_self, ckpt_name, concept_data): @@ -559,7 +570,7 @@ def load_kolors_model(loader_self, ckpt_name, concept_data): 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 + return _wrap_refiner(KOLORS_MODEL, CHATGLM3_MODEL, OUTPUT_VAE, loader_self, concept_data) def load_hunyuan_model(loader_self, ckpt_name, concept_data): @@ -593,7 +604,7 @@ def load_hunyuan_model(loader_self, ckpt_name, concept_data): if t5_path: T5 = load_hydit_t5(model_path=t5_path, device='GPU', dtype=dtype) - return HUNYUAN_MODEL, {'clip': CLIP, 't5': T5}, HUNYUAN_VAE + return _wrap_refiner(HUNYUAN_MODEL, {'clip': CLIP, 't5': T5}, HUNYUAN_VAE, loader_self, concept_data) def load_qwen_model(loader_self, ckpt_name, concept_data): @@ -635,7 +646,7 @@ def load_qwen_model(loader_self, ckpt_name, concept_data): OUTPUT_MODEL = nodes_model_advanced.ModelSamplingSD3.patch(loader_self, OUTPUT_MODEL, 3, 1.0)[0] OUTPUT_MODEL = nodes_cfg.CFGNorm.execute(OUTPUT_MODEL, 1)[0] - return OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE + return _wrap_refiner(OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE, loader_self, concept_data) def load_chroma_model(loader_self, ckpt_name, concept_data): @@ -671,6 +682,6 @@ def load_chroma_model(loader_self, ckpt_name, concept_data): rescale_cfg = concept_data.get('rescale_cfg', 1.0) if rescale_cfg != 1.0: OUTPUT_MODEL = nodes_model_advanced.RescaleCFG.patch(loader_self, OUTPUT_MODEL, rescale_cfg)[0] - return OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE + return _wrap_refiner(OUTPUT_MODEL, OUTPUT_CLIP, OUTPUT_VAE, loader_self, concept_data) diff --git a/components/primeresamplers.py b/components/primeresamplers.py index cf08132..08e7e34 100644 --- a/components/primeresamplers.py +++ b/components/primeresamplers.py @@ -189,6 +189,30 @@ def PSamplerSana(self, device, seed, model, return (latent_out,) +def _run_refiner_pass(self, refiner_model, refiner_cond_pos, refiner_cond_neg, samples_main, control_data, seed): + main_vae = utility.vae_loader_class.load_vae(control_data.get('vae'))[0] + raw_image = nodes.VAEDecode.decode(self, main_vae, samples_main)[0] + refiner_ckpt = nodes.CheckpointLoaderSimple.load_checkpoint(self, control_data.get('refiner_model')) + encoded_image = nodes.VAEEncode.encode(self, refiner_ckpt[2], raw_image)[0] + refiner_sampler = control_data.get('refiner_sampler', 'dpmpp_2m') + refiner_scheduler = control_data.get('refiner_scheduler', 'normal') + refiner_cfg = float(control_data.get('refiner_cfg', 2.0)) + refiner_steps = int(control_data.get('refiner_steps', 22)) + refiner_denoise = float(control_data.get('refiner_denoise', 0.9)) + refiner_start = int(control_data.get('refiner_start', 12)) + sigmas = nodes_custom_sampler.BasicScheduler.execute(refiner_model, refiner_scheduler, refiner_steps, refiner_denoise)[0] + low_sigmas = nodes_custom_sampler.SplitSigmas.execute(sigmas, refiner_start)[1] + sampler = comfy.samplers.sampler_object(refiner_sampler) + if control_data.get('refiner_ignore_prompt', True): + empty_cond = nodes.CLIPTextEncode.encode(self, refiner_ckpt[1], "")[0] + pos_cond = empty_cond + neg_cond = empty_cond + else: + pos_cond = refiner_cond_pos if refiner_cond_pos is not None else nodes.CLIPTextEncode.encode(self, refiner_ckpt[1], "")[0] + neg_cond = refiner_cond_neg if refiner_cond_neg is not None else nodes.CLIPTextEncode.encode(self, refiner_ckpt[1], "")[0] + return nodes_custom_sampler.SamplerCustom.execute(refiner_model, True, seed, refiner_cfg, pos_cond, neg_cond, sampler, low_sigmas, encoded_image)[0] + + def PSamplerPixart(self, device, seed, model, steps, cfg, sampler_name, scheduler_name, positive, negative, @@ -200,50 +224,25 @@ def PSamplerPixart(self, device, seed, model, scheduler_name = control_data.get('scheduler_name', scheduler_name) steps = control_data.get('steps', steps) cfg = control_data.get('cfg', cfg) - PIXART_DENOISE = float(control_data.get('refiner_sampling_denoise', denoise)) if control_data and control_data.get('refiner') == True else denoise + if control_data and control_data.get('refiner') == True: + denoise = float(control_data.get('refiner_sampling_denoise', denoise)) - sigmas_main = nodes_custom_sampler.BasicScheduler.execute(model['main'], scheduler_name, steps, denoise=PIXART_DENOISE)[0] + sigmas_main = nodes_custom_sampler.BasicScheduler.execute(model, scheduler_name, steps, denoise=denoise)[0] sampler = comfy.samplers.sampler_object(sampler_name) if variation_level == True: - samples_main = latentnoise.noisy_samples(model['main'], device, steps, cfg, sampler_name, scheduler_name, positive['main'], negative['main'], latent_image, PIXART_DENOISE, seed, noise_extender)[0] + samples_main = latentnoise.noisy_samples(model, device, steps, cfg, sampler_name, scheduler_name, positive, negative, latent_image, denoise, seed, noise_extender)[0] else: if variation_extender_original > 0 or device != 'DEFAULT' or variation_batch_step_original > 0: - samples_main = latentnoise.noisy_samples(model['main'], device, steps, cfg, sampler_name, scheduler_name, positive['main'], negative['main'], latent_image, PIXART_DENOISE, seed, noise_extender)[0] + samples_main = latentnoise.noisy_samples(model, device, steps, cfg, sampler_name, scheduler_name, positive, negative, latent_image, denoise, seed, noise_extender)[0] else: if align_your_steps == True: model_type = 'SDXL' sigmas = nodes_align_your_steps.AlignYourStepsScheduler.get_sigmas(self, model_type, steps, denoise) sampler = comfy.samplers.sampler_object(sampler_name) - AYS_samples = nodes_custom_sampler.SamplerCustom.execute(model['main'], True, seed, cfg, positive['main'], negative['main'], sampler, sigmas[0], latent_image) + AYS_samples = nodes_custom_sampler.SamplerCustom.execute(model, True, seed, cfg, positive, negative, sampler, sigmas[0], latent_image) samples_main = AYS_samples[0] else: - samples_main = nodes_custom_sampler.SamplerCustom.execute(model['main'], True, seed, cfg, positive['main'], negative['main'], sampler, sigmas_main, latent_image)[0] - - if 'refiner' in model and model['refiner'] is not None: - PIXART_VAE = utility.vae_loader_class.load_vae(control_data.get('vae'))[0] - RAW_IMAGE = nodes.VAEDecode.decode(self, PIXART_VAE, samples_main)[0] - PIXART_REFINER_CHECKPOINT = nodes.CheckpointLoaderSimple.load_checkpoint(self, control_data.get('refiner_model')) - RAW_IMAGE_ENCODED = nodes.VAEEncode.encode(self, PIXART_REFINER_CHECKPOINT[2], RAW_IMAGE)[0] - - REFINER_SAMPLER = control_data.get('refiner_sampler', 'dpmpp_2m') - REFINER_SCHEDULER = control_data.get('refiner_scheduler', 'normal') - REFINER_CFG = float(control_data.get('refiner_cfg', 2.0)) - REFINER_STEPS = int(control_data.get('refiner_steps', 22)) - PIXART_DENOISE_REFINER = float(control_data.get('refiner_denoise', 0.9)) - PIXART_REFINER_START = int(control_data.get('refiner_start', 12)) - - sigmas_refiner = nodes_custom_sampler.BasicScheduler.execute(model['refiner'], REFINER_SCHEDULER, REFINER_STEPS, PIXART_DENOISE_REFINER)[0] - splitted_low_sigma = nodes_custom_sampler.SplitSigmas.execute(sigmas_refiner, PIXART_REFINER_START)[1] - sampler_refiner = comfy.samplers.sampler_object(REFINER_SAMPLER) - REFINER_IGNORE_PROMPT = control_data.get('refiner_ignore_prompt', False) - if REFINER_IGNORE_PROMPT: - empty_cond = nodes.CLIPTextEncode.encode(self, PIXART_REFINER_CHECKPOINT[1], "")[0] - refiner_pos = empty_cond - refiner_neg = empty_cond - else: - refiner_pos = positive['refiner'] - refiner_neg = negative['refiner'] - samples_main = nodes_custom_sampler.SamplerCustom.execute(model['refiner'], True, seed, REFINER_CFG, refiner_pos, refiner_neg, sampler_refiner, splitted_low_sigma, RAW_IMAGE_ENCODED)[0] + samples_main = nodes_custom_sampler.SamplerCustom.execute(model, True, seed, cfg, positive, negative, sampler, sigmas_main, latent_image)[0] return (samples_main,)