| 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.") |
|
|