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