RioShiina commited on
Commit
e372da2
·
verified ·
1 Parent(s): e4288e9

Krea-2 Identity Edit injector: Add image scaling before VAE encode.

Browse files
chain_injectors/krea2_identity_edit_injector.py CHANGED
@@ -1,6 +1,18 @@
1
  import os
2
  from utils.app_utils import ensure_file_downloaded
3
 
 
 
 
 
 
 
 
 
 
 
 
 
4
  def inject(assembler, chain_definition, chain_items):
5
  if not chain_items:
6
  return
@@ -52,6 +64,11 @@ def inject(assembler, chain_definition, chain_items):
52
  vae_connection = None
53
  if vae_loader_name in assembler.node_map:
54
  vae_connection = [assembler.node_map[vae_loader_name], 0]
 
 
 
 
 
55
 
56
  clip_connection = None
57
  if clip_loader_name in assembler.node_map:
@@ -61,11 +78,10 @@ def inject(assembler, chain_definition, chain_items):
61
  clip_connection = assembler.workflow[pos_id]['inputs'].get('clip')
62
 
63
  lora_loader_id = assembler._get_unique_id()
64
- lora_loader_node = assembler._get_node_template("LoraLoaderModelOnly")
65
  lora_loader_node['inputs']['lora_name'] = lora_filename
66
  lora_loader_node['inputs']['strength_model'] = 1.0
67
  lora_loader_node['inputs']['model'] = current_model_connection
68
- lora_loader_node['_meta']['title'] = "Load LoRA (Krea2 Identity Edit)"
69
  assembler.workflow[lora_loader_id] = lora_loader_node
70
 
71
  image_ids = []
@@ -73,23 +89,29 @@ def inject(assembler, chain_definition, chain_items):
73
 
74
  for i, img_filename in enumerate(valid_images):
75
  load_id = assembler._get_unique_id()
76
- load_node = assembler._get_node_template("LoadImage")
77
  load_node['inputs']['image'] = img_filename
78
- load_node['_meta']['title'] = f"Load Image (Ref {i+1})"
79
  assembler.workflow[load_id] = load_node
80
- image_ids.append(load_id)
 
 
 
 
 
 
 
 
81
 
82
  vae_enc_id = assembler._get_unique_id()
83
- vae_enc_node = assembler._get_node_template("VAEEncode")
84
- vae_enc_node['inputs']['pixels'] = [load_id, 0]
85
  if vae_connection:
86
  vae_enc_node['inputs']['vae'] = vae_connection
87
- vae_enc_node['_meta']['title'] = f"VAE Encode (Ref {i+1})"
88
  assembler.workflow[vae_enc_id] = vae_enc_node
89
  vae_encode_ids.append(vae_enc_id)
90
 
91
  patch_id = assembler._get_unique_id()
92
- patch_node = assembler._get_node_template("Krea2EditModelPatch")
93
  patch_node['inputs']['ref_boost'] = 4
94
  patch_node['inputs']['ref_boost_a'] = 1
95
  patch_node['inputs']['fit_mode'] = "fit"
@@ -104,7 +126,6 @@ def inject(assembler, chain_definition, chain_items):
104
  patch_node['inputs']['source_latent_b'] = [vae_encode_ids[1], 0]
105
  patch_node['inputs']['source_image_b'] = [image_ids[1], 0]
106
 
107
- patch_node['_meta']['title'] = "Krea2 Edit (source patch)"
108
  assembler.workflow[patch_id] = patch_node
109
 
110
  assembler.workflow[ksampler_id]['inputs']['model'] = [patch_id, 0]
 
1
  import os
2
  from utils.app_utils import ensure_file_downloaded
3
 
4
+ def create_node(assembler, class_type, title):
5
+ try:
6
+ node = assembler._get_node_template(class_type)
7
+ except Exception:
8
+ node = {
9
+ "inputs": {},
10
+ "class_type": class_type,
11
+ "_meta": {"title": title}
12
+ }
13
+ node['_meta']['title'] = title
14
+ return node
15
+
16
  def inject(assembler, chain_definition, chain_items):
17
  if not chain_items:
18
  return
 
64
  vae_connection = None
65
  if vae_loader_name in assembler.node_map:
66
  vae_connection = [assembler.node_map[vae_loader_name], 0]
67
+ else:
68
+ for node_id, node in assembler.workflow.items():
69
+ if isinstance(node, dict) and node.get('class_type') == 'VAELoader':
70
+ vae_connection = [node_id, 0]
71
+ break
72
 
73
  clip_connection = None
74
  if clip_loader_name in assembler.node_map:
 
78
  clip_connection = assembler.workflow[pos_id]['inputs'].get('clip')
79
 
80
  lora_loader_id = assembler._get_unique_id()
81
+ lora_loader_node = create_node(assembler, "LoraLoaderModelOnly", "Load LoRA (Krea2 Identity Edit)")
82
  lora_loader_node['inputs']['lora_name'] = lora_filename
83
  lora_loader_node['inputs']['strength_model'] = 1.0
84
  lora_loader_node['inputs']['model'] = current_model_connection
 
85
  assembler.workflow[lora_loader_id] = lora_loader_node
86
 
87
  image_ids = []
 
89
 
90
  for i, img_filename in enumerate(valid_images):
91
  load_id = assembler._get_unique_id()
92
+ load_node = create_node(assembler, "LoadImage", f"Load Image (Ref {i+1})")
93
  load_node['inputs']['image'] = img_filename
 
94
  assembler.workflow[load_id] = load_node
95
+
96
+ scale_id = assembler._get_unique_id()
97
+ scale_node = create_node(assembler, "ImageScaleToTotalPixels", f"Scale Reference {i+1}")
98
+ scale_node['inputs']['upscale_method'] = "lanczos"
99
+ scale_node['inputs']['megapixels'] = 1.0
100
+ scale_node['inputs']['resolution_steps'] = 1
101
+ scale_node['inputs']['image'] = [load_id, 0]
102
+ assembler.workflow[scale_id] = scale_node
103
+ image_ids.append(scale_id)
104
 
105
  vae_enc_id = assembler._get_unique_id()
106
+ vae_enc_node = create_node(assembler, "VAEEncode", f"VAE Encode (Ref {i+1})")
107
+ vae_enc_node['inputs']['pixels'] = [scale_id, 0]
108
  if vae_connection:
109
  vae_enc_node['inputs']['vae'] = vae_connection
 
110
  assembler.workflow[vae_enc_id] = vae_enc_node
111
  vae_encode_ids.append(vae_enc_id)
112
 
113
  patch_id = assembler._get_unique_id()
114
+ patch_node = create_node(assembler, "Krea2EditModelPatch", "Krea2 Edit (source patch)")
115
  patch_node['inputs']['ref_boost'] = 4
116
  patch_node['inputs']['ref_boost_a'] = 1
117
  patch_node['inputs']['fit_mode'] = "fit"
 
126
  patch_node['inputs']['source_latent_b'] = [vae_encode_ids[1], 0]
127
  patch_node['inputs']['source_image_b'] = [image_ids[1], 0]
128
 
 
129
  assembler.workflow[patch_id] = patch_node
130
 
131
  assembler.workflow[ksampler_id]['inputs']['model'] = [patch_id, 0]
requirements.txt CHANGED
@@ -1,5 +1,5 @@
1
  comfyui-frontend-package==1.51.9
2
- comfyui-workflow-templates==0.11.50
3
  comfyui-embedded-docs==0.5.10
4
  torch
5
  torchsde
 
1
  comfyui-frontend-package==1.51.9
2
+ comfyui-workflow-templates==0.11.54
3
  comfyui-embedded-docs==0.5.10
4
  torch
5
  torchsde