V 2.0.0 - Auto config #24 - refiners for all
This commit is contained in:
@@ -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
|
||||
|
||||
+79
-33
@@ -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
|
||||
|
||||
+38
-27
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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,)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user