RioShiina commited on
Commit
e9e54d7
·
verified ·
1 Parent(s): 9331538

Upload krea2_style_reference_injector.py

Browse files
chain_injectors/krea2_style_reference_injector.py CHANGED
@@ -156,8 +156,41 @@ def inject(assembler, chain_definition, chain_items):
156
  neg_ref_node['inputs']['conditioning'] = [neg_encode_id, 0]
157
  assembler.workflow[neg_ref_id] = neg_ref_node
158
 
159
- assembler.workflow[ksampler_id]['inputs']['positive'] = [pos_ref_id, 0]
160
- assembler.workflow[ksampler_id]['inputs']['negative'] = [neg_ref_id, 0]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
161
 
162
  if pos_prompt_id and pos_prompt_id in assembler.workflow:
163
  del assembler.workflow[pos_prompt_id]
 
156
  neg_ref_node['inputs']['conditioning'] = [neg_encode_id, 0]
157
  assembler.workflow[neg_ref_id] = neg_ref_node
158
 
159
+ existing_pos = assembler.workflow[ksampler_id]['inputs'].get('positive')
160
+ existing_neg = assembler.workflow[ksampler_id]['inputs'].get('negative')
161
+
162
+ has_krea2_edit = False
163
+ if existing_pos and isinstance(existing_pos, (list, tuple)) and len(existing_pos) > 0:
164
+ pos_node_id = existing_pos[0]
165
+ if pos_node_id in assembler.workflow:
166
+ pos_node = assembler.workflow[pos_node_id]
167
+ if isinstance(pos_node, dict) and pos_node.get('class_type') == 'Krea2EditGroundedEncode':
168
+ has_krea2_edit = True
169
+
170
+ if not has_krea2_edit:
171
+ for node in assembler.workflow.values():
172
+ if isinstance(node, dict) and node.get('class_type') in ['Krea2EditModelPatch', 'Krea2EditGroundedEncode']:
173
+ has_krea2_edit = True
174
+ break
175
+
176
+ if has_krea2_edit and existing_pos and existing_neg:
177
+ combine_pos_id = assembler._get_unique_id()
178
+ combine_pos_node = create_node(assembler, "ConditioningCombine", "Conditioning (Combine)")
179
+ combine_pos_node['inputs']['conditioning_1'] = existing_pos
180
+ combine_pos_node['inputs']['conditioning_2'] = [pos_ref_id, 0]
181
+ assembler.workflow[combine_pos_id] = combine_pos_node
182
+
183
+ combine_neg_id = assembler._get_unique_id()
184
+ combine_neg_node = create_node(assembler, "ConditioningCombine", "Conditioning (Combine)")
185
+ combine_neg_node['inputs']['conditioning_1'] = existing_neg
186
+ combine_neg_node['inputs']['conditioning_2'] = [neg_ref_id, 0]
187
+ assembler.workflow[combine_neg_id] = combine_neg_node
188
+
189
+ assembler.workflow[ksampler_id]['inputs']['positive'] = [combine_pos_id, 0]
190
+ assembler.workflow[ksampler_id]['inputs']['negative'] = [combine_neg_id, 0]
191
+ else:
192
+ assembler.workflow[ksampler_id]['inputs']['positive'] = [pos_ref_id, 0]
193
+ assembler.workflow[ksampler_id]['inputs']['negative'] = [neg_ref_id, 0]
194
 
195
  if pos_prompt_id and pos_prompt_id in assembler.workflow:
196
  del assembler.workflow[pos_prompt_id]