File size: 7,454 Bytes
b4cd781 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 | import os
from utils.app_utils import ensure_file_downloaded
def create_node(assembler, class_type, title):
try:
node = assembler._get_node_template(class_type)
except Exception:
node = {
"inputs": {},
"class_type": class_type,
"_meta": {"title": title}
}
node['_meta']['title'] = title
return node
def inject(assembler, chain_definition, chain_items):
if not chain_items:
return
valid_images = []
for item in chain_items:
if not item:
continue
img_path = item
if isinstance(item, dict):
img_path = item.get('image') or item.get('filename') or item.get('path')
if img_path:
valid_images.append(img_path)
if not valid_images:
return
valid_images = valid_images[:3]
lora_filename = "krea2_style_reference.safetensors"
try:
ensure_file_downloaded(lora_filename)
except Exception as e:
print(f"Warning: Failed to ensure '{lora_filename}' downloaded: {e}")
ksampler_name = chain_definition.get('ksampler_node', 'ksampler')
pos_prompt_name = chain_definition.get('pos_prompt_node', 'pos_prompt')
neg_prompt_name = chain_definition.get('neg_prompt_node', 'neg_prompt')
clip_loader_name = chain_definition.get('clip_loader_node', 'clip_loader')
vae_loader_name = chain_definition.get('vae_loader_node', 'vae_loader')
if ksampler_name not in assembler.node_map:
print(f"Warning: Target node '{ksampler_name}' for Krea2 Style Reference Edit chain not found. Skipping.")
return
ksampler_id = assembler.node_map[ksampler_name]
if 'model' not in assembler.workflow[ksampler_id]['inputs']:
print(f"Warning: KSampler node '{ksampler_name}' is missing 'model' input. Skipping.")
return
current_model_connection = assembler.workflow[ksampler_id]['inputs']['model']
vae_connection = None
if vae_loader_name in assembler.node_map:
vae_connection = [assembler.node_map[vae_loader_name], 0]
else:
for node_id, node in assembler.workflow.items():
if isinstance(node, dict) and node.get('class_type') == 'VAELoader':
vae_connection = [node_id, 0]
break
clip_connection = None
if clip_loader_name in assembler.node_map:
clip_connection = [assembler.node_map[clip_loader_name], 0]
elif pos_prompt_name in assembler.node_map:
pos_id = assembler.node_map[pos_prompt_name]
clip_connection = assembler.workflow[pos_id]['inputs'].get('clip')
scaled_image_ids = []
for i, img_filename in enumerate(valid_images):
load_id = assembler._get_unique_id()
load_node = create_node(assembler, "LoadImage", f"Load Reference Image {i+1}")
load_node['inputs']['image'] = img_filename
assembler.workflow[load_id] = load_node
scale_id = assembler._get_unique_id()
scale_node = create_node(assembler, "ImageScaleToTotalPixels", f"Scale Reference {i+1}")
scale_node['inputs']['upscale_method'] = "nearest-exact"
scale_node['inputs']['megapixels'] = 1
scale_node['inputs']['resolution_steps'] = 1
scale_node['inputs']['image'] = [load_id, 0]
assembler.workflow[scale_id] = scale_node
scaled_image_ids.append(scale_id)
lora_loader_id = assembler._get_unique_id()
lora_loader_node = create_node(assembler, "LoraLoaderModelOnly", "Load LoRA (Krea2 Style Reference)")
lora_loader_node['inputs']['lora_name'] = lora_filename
lora_loader_node['inputs']['strength_model'] = 1.0
lora_loader_node['inputs']['model'] = current_model_connection
assembler.workflow[lora_loader_id] = lora_loader_node
assembler.workflow[ksampler_id]['inputs']['model'] = [lora_loader_id, 0]
pos_prompt_id = assembler.node_map.get(pos_prompt_name)
neg_prompt_id = assembler.node_map.get(neg_prompt_name)
pos_text = ""
if pos_prompt_id and pos_prompt_id in assembler.workflow:
pos_text = assembler.workflow[pos_prompt_id]['inputs'].get('text', '')
elif hasattr(assembler, 'ui_values') and isinstance(assembler.ui_values, dict):
pos_text = assembler.ui_values.get('positive_prompt') or assembler.ui_values.get('prompt') or ''
if not pos_text:
for node_id, node in assembler.workflow.items():
if isinstance(node, dict):
cls = node.get('class_type', '')
if cls in ['Krea2EditGroundedEncode', 'TextEncodeQwenImageEditPlus', 'CLIPTextEncode']:
t = node.get('inputs', {}).get('prompt') or node.get('inputs', {}).get('text')
if t:
pos_text = t
break
neg_text = ""
if neg_prompt_id and neg_prompt_id in assembler.workflow:
neg_text = assembler.workflow[neg_prompt_id]['inputs'].get('text', '')
elif hasattr(assembler, 'ui_values') and isinstance(assembler.ui_values, dict):
neg_text = assembler.ui_values.get('negative_prompt') or assembler.ui_values.get('neg_prompt') or ''
pos_encode_id = assembler._get_unique_id()
pos_encode_node = create_node(assembler, "TextEncodeQwenImageEditPlus", "TextEncodeQwenImageEditPlus (Positive)")
pos_encode_node['inputs']['prompt'] = pos_text
if clip_connection:
pos_encode_node['inputs']['clip'] = clip_connection
if vae_connection:
pos_encode_node['inputs']['vae'] = vae_connection
for idx, s_id in enumerate(scaled_image_ids):
pos_encode_node['inputs'][f"image{idx+1}"] = [s_id, 0]
assembler.workflow[pos_encode_id] = pos_encode_node
neg_encode_id = assembler._get_unique_id()
neg_encode_node = create_node(assembler, "TextEncodeQwenImageEditPlus", "TextEncodeQwenImageEditPlus (Negative)")
neg_encode_node['inputs']['prompt'] = neg_text
if clip_connection:
neg_encode_node['inputs']['clip'] = clip_connection
if vae_connection:
neg_encode_node['inputs']['vae'] = vae_connection
for idx, s_id in enumerate(scaled_image_ids):
neg_encode_node['inputs'][f"image{idx+1}"] = [s_id, 0]
assembler.workflow[neg_encode_id] = neg_encode_node
pos_ref_id = assembler._get_unique_id()
pos_ref_node = create_node(assembler, "FluxKontextMultiReferenceLatentMethod", "Edit Model Reference Method")
pos_ref_node['inputs']['reference_latents_method'] = "index_timestep_zero"
pos_ref_node['inputs']['conditioning'] = [pos_encode_id, 0]
assembler.workflow[pos_ref_id] = pos_ref_node
neg_ref_id = assembler._get_unique_id()
neg_ref_node = create_node(assembler, "FluxKontextMultiReferenceLatentMethod", "Edit Model Reference Method")
neg_ref_node['inputs']['reference_latents_method'] = "index_timestep_zero"
neg_ref_node['inputs']['conditioning'] = [neg_encode_id, 0]
assembler.workflow[neg_ref_id] = neg_ref_node
assembler.workflow[ksampler_id]['inputs']['positive'] = [pos_ref_id, 0]
assembler.workflow[ksampler_id]['inputs']['negative'] = [neg_ref_id, 0]
if pos_prompt_id and pos_prompt_id in assembler.workflow:
del assembler.workflow[pos_prompt_id]
if neg_prompt_id and neg_prompt_id in assembler.workflow:
del assembler.workflow[neg_prompt_id]
print(f"Krea2 Style Reference Edit injector applied with {len(valid_images)} reference image(s). Original CLIPTextEncode nodes replaced.")
|