Added Gemini Adv image gen node + image batch node

This commit is contained in:
Fill
2025-05-14 21:17:51 +09:00
parent d6a342ec17
commit c3e85fffd8
7 changed files with 604 additions and 3 deletions
+6
View File
@@ -122,6 +122,8 @@ from .nodes.FL_UnloadModel import FL_UnloadModel, FL_UnloadAllModels
from .nodes.FL_VideoTrim import FL_VideoTrim
from .nodes.FL_VideoCadence import FL_VideoCadence
from .nodes.FL_VideoCadenceCompile import FL_VideoCadenceCompile
from .nodes.FL_GeminiImageGenADV import FL_GeminiImageGenADV
from .nodes.FL_ImageBatch import FL_ImageBatch
NODE_CLASS_MAPPINGS = {
@@ -252,6 +254,8 @@ NODE_CLASS_MAPPINGS = {
"FL_VideoTrim": FL_VideoTrim,
"FL_VideoCadence": FL_VideoCadence,
"FL_VideoCadenceCompile": FL_VideoCadenceCompile,
"FL_GeminiImageGenADV": FL_GeminiImageGenADV,
"FL_ImageBatch": FL_ImageBatch,
}
@@ -383,6 +387,8 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"FL_VideoTrim": "FL Video Trim",
"FL_VideoCadence": "FL Video Cadence",
"FL_VideoCadenceCompile": "FL Video Cadence Compile",
"FL_GeminiImageGenADV": "FL Gemini Image Gen ADV",
"FL_ImageBatch": "FL Image Batch",
}
+399
View File
@@ -0,0 +1,399 @@
import os
import base64
import io
import json
import torch
import numpy as np
from PIL import Image, ImageDraw, ImageFont
import requests
import tempfile
from io import BytesIO
from google import genai
from google.genai import types
import time
import traceback
import asyncio
import concurrent.futures
import random
from typing import List, Tuple, Optional
# Assuming ImageBatch is still needed if we are batching results, or can be removed if Gemini returns a batch
# from nodes import ImageBatch
class FL_GeminiImageGenADV:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"inputcount": ("INT", {"default": 1, "min": 1, "max": 10, "step": 1}),
"api_key": ("STRING", {"default": os.getenv("GEMINI_API_KEY", ""), "multiline": False}),
"model": (["models/gemini-2.0-flash-exp", "models/gemini-2.0-flash-preview-image-generation"], {"default": "models/gemini-2.0-flash-preview-image-generation"}),
"temperature": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.05}),
"max_retries": ("INT", {"default": 3, "min": 1, "max": 5, "step": 1}),
"prompt_1": ("STRING", {"multiline": False, "default": "Describe image 1", "forceInput": True}),
},
"optional": {
"image_1": ("IMAGE", {}), # Moved image_1 to optional. Default will be None if not connected.
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), # Restored full seed range
# Subsequent image_i and prompt_i will be handled by **kwargs based on inputcount
}
}
RETURN_TYPES = ("IMAGE", "STRING")
RETURN_NAMES = ("images", "API_responses")
FUNCTION = "generate_images_advanced"
CATEGORY = "🏵️Fill Nodes/AI"
DESCRIPTION = """
Generates images using Gemini based on multiple image/prompt pairs.
Each pair triggers an asynchronous API call. Results are batched.
"""
def __init__(self):
self.log_messages = []
self.min_size = 1024 # Minimum size from Editor
try:
import importlib.metadata
genai_version = importlib.metadata.version('google-genai')
self._log(f"Current google-genai version: {genai_version}")
from packaging import version # Ensure packaging is imported
if version.parse(genai_version) < version.parse('0.8.0'): # Example version, check Gemini docs
self._log("Warning: google-genai version is too low, recommend upgrading to the latest version.")
self._log("Suggested: pip install -q -U google-genai")
except Exception as e:
self._log(f"Unable to check google-genai version: {e}")
def _log(self, message):
timestamp = time.strftime("%Y-%m-%d %H:%M:%S")
formatted_message = f"[FL_GeminiImageGenADV] {timestamp}: {message}"
print(formatted_message)
if hasattr(self, 'log_messages'):
self.log_messages.append(message)
return message
def _pad_image_to_minimum_size(self, pil_image):
width, height = pil_image.size
if width >= self.min_size and height >= self.min_size:
return pil_image
new_width = max(width, self.min_size)
new_height = max(height, self.min_size)
new_image = Image.new('RGB', (new_width, new_height), color=(255, 255, 255)) # White padding from Editor
paste_x = (new_width - width) // 2
paste_y = (new_height - height) // 2
new_image.paste(pil_image, (paste_x, paste_y))
self._log(f"Padded image from {width}x{height} to {new_width}x{new_height} with white borders")
return new_image
def _create_error_image(self, error_message="API Error", width=1024, height=1024): # Default size from Editor
image = Image.new('RGB', (width, height), color=(0, 0, 0)) # Black error image from Editor
draw = ImageDraw.Draw(image)
font = None
try:
# Try to find a font that exists on most systems
font_options = ['arial.ttf', 'DejaVuSans.ttf', 'FreeSans.ttf', 'NotoSans-Regular.ttf']
for font_name in font_options:
try:
font = ImageFont.truetype(font_name, 24) # Font size from Editor
break
except IOError:
continue
if font is None:
font = ImageFont.load_default()
except Exception:
font = ImageFont.load_default()
# Calculate text position (centered)
try: # Newer PIL versions
text_bbox = draw.textbbox((0,0), error_message, font=font)
text_width = text_bbox[2] - text_bbox[0]
text_height = text_bbox[3] - text_bbox[1]
except AttributeError: # Older PIL versions
text_width, text_height = draw.textsize(error_message, font=font)
text_x = (width - text_width) / 2
text_y = (height - text_height) / 2
draw.text((text_x, text_y), error_message, fill=(255, 0, 0), font=font) # Red text from Editor
img_array = np.array(image).astype(np.float32) / 255.0
img_tensor = torch.from_numpy(img_array).unsqueeze(0)
self._log(f"Created error image: '{error_message}'")
return img_tensor
def _process_tensor_to_pil(self, tensor_image: Optional[torch.Tensor], image_name: str = "Image") -> Optional[Image.Image]:
if tensor_image is None:
self._log(f"{image_name} is None, skipping PIL conversion.")
return None
if not isinstance(tensor_image, torch.Tensor):
self._log(f"{image_name} is not a tensor, skipping. Type: {type(tensor_image)}")
return None
if tensor_image.ndim == 4 and tensor_image.shape[0] > 0:
img_np = tensor_image[0].cpu().numpy() # Use the first image in the batch
img_np = (img_np * 255).astype(np.uint8)
pil_image = Image.fromarray(img_np)
self._log(f"Converted {image_name} (shape: {tensor_image.shape}) to PIL Image (size: {pil_image.size}).")
return pil_image
elif tensor_image.ndim == 3: # Assume (H, W, C)
img_np = tensor_image.cpu().numpy()
img_np = (img_np * 255).astype(np.uint8)
pil_image = Image.fromarray(img_np)
self._log(f"Converted {image_name} (shape: {tensor_image.shape}) to PIL Image (size: {pil_image.size}).")
return pil_image
else:
self._log(f"Cannot convert {image_name} with shape {tensor_image.shape} to PIL Image.")
return None
def _call_gemini_api(self, client_instance, model_name_full, contents, gen_config_obj, retry_count=0, max_retries=3, call_id="0"):
try:
self._log(f"[Call {call_id}] API call attempt #{retry_count + 1} to {model_name_full}")
# Using client.models.generate_content like in FL_GeminiImageEditor
response = client_instance.models.generate_content(
model=model_name_full, # FL_GeminiImageEditor passes the full model string here
contents=contents,
config=gen_config_obj # FL_GeminiImageEditor uses 'config' for GenerateContentConfig
)
# Validate response structure (adapted from FL_GeminiImageEditor)
if not hasattr(response, 'candidates') or not response.candidates:
self._log(f"[Call {call_id}] Empty response: No candidates found")
if retry_count < max_retries - 1:
self._log(f"[Call {call_id}] Retrying in 2 seconds... (Attempt {retry_count + 2}/{max_retries})")
time.sleep(2)
return self._call_gemini_api(client_instance, model_name_full, contents, gen_config_obj, retry_count + 1, max_retries, call_id)
else:
self._log(f"[Call {call_id}] Maximum retries ({max_retries}) reached. Returning empty response.")
return None
if not hasattr(response.candidates[0], 'content') or response.candidates[0].content is None:
self._log(f"[Call {call_id}] Invalid response: candidates[0].content is missing")
if retry_count < max_retries - 1:
self._log(f"[Call {call_id}] Retrying in 2 seconds... (Attempt {retry_count + 2}/{max_retries})")
time.sleep(2)
return self._call_gemini_api(client_instance, model_name_full, contents, gen_config_obj, retry_count + 1, max_retries, call_id)
else:
self._log(f"[Call {call_id}] Maximum retries ({max_retries}) reached. Returning empty response.")
return None
if not hasattr(response.candidates[0].content, 'parts') or response.candidates[0].content.parts is None:
self._log(f"[Call {call_id}] Invalid response: candidates[0].content.parts is missing")
if retry_count < max_retries - 1:
self._log(f"[Call {call_id}] Retrying in 2 seconds... (Attempt {retry_count + 2}/{max_retries})")
time.sleep(2)
return self._call_gemini_api(client_instance, model_name_full, contents, gen_config_obj, retry_count + 1, max_retries, call_id)
else:
self._log(f"[Call {call_id}] Maximum retries ({max_retries}) reached. Returning empty response.")
return None
self._log(f"[Call {call_id}] Valid API response received.")
return response
except Exception as e:
self._log(f"[Call {call_id}] API call error: {str(e)}")
if retry_count < max_retries - 1:
wait_time = 2 * (retry_count + 1) # Progressive backoff
self._log(f"[Call {call_id}] Retrying in {wait_time}s... (Attempt {retry_count + 2}/{max_retries})")
time.sleep(wait_time)
return self._call_gemini_api(client_instance, model_name_full, contents, gen_config_obj, retry_count + 1, max_retries, call_id)
else:
self._log(f"[Call {call_id}] Max retries ({max_retries}) reached. Giving up.")
return None
def _process_api_response(self, response, call_id="0"):
if response is None: # Simplified check from Editor
self._log(f"[Call {call_id}] No valid response to process.")
error_msg = "API Error: No content in response"
return self._create_error_image(error_msg), error_msg
response_text_parts = [] # Changed from response_text to response_text_parts to match ADV logic initially
image_tensor = None
if not hasattr(response, 'candidates') or not response.candidates: # Check from Editor
self._log(f"[Call {call_id}] No candidates in API response")
error_msg = "API returned an empty response"
return self._create_error_image(error_msg), error_msg
# Iterate through response parts (similar to Editor, but adapted for ADV's single image focus per call)
for part in response.candidates[0].content.parts:
if hasattr(part, 'text') and part.text is not None:
text_content = part.text
response_text_parts.append(text_content)
self._log(
f"[Call {call_id}] API returned text: {text_content[:100]}..." if len(
text_content) > 100 else text_content)
elif hasattr(part, 'inline_data') and part.inline_data is not None:
self._log(f"[Call {call_id}] API returned image data")
try:
image_data = part.inline_data.data
if not image_data or len(image_data) < 100: # Check from Editor
self._log(f"[Call {call_id}] Warning: Image data is empty or too small")
continue
pil_image = None
try:
pil_image = Image.open(BytesIO(image_data))
self._log(
f"[Call {call_id}] Direct PIL open successful, size: {pil_image.width}x{pil_image.height}")
except Exception as e1:
self._log(f"[Call {call_id}] Direct PIL open failed: {str(e1)}")
# Optional: Add temp file saving from Editor if direct open fails often
try:
temp_dir = tempfile.gettempdir()
# Ensure temp_dir is writable, or fall back
if not os.access(temp_dir, os.W_OK):
temp_dir = "." # Current directory as fallback
self._log(f"[Call {call_id}] Temp directory {tempfile.gettempdir()} not writable, using current directory.")
temp_file_path = os.path.join(temp_dir, f"gemini_image_adv_{call_id}_{int(time.time())}.png")
with open(temp_file_path, "wb") as f:
f.write(image_data)
pil_image = Image.open(temp_file_path)
self._log(f"[Call {call_id}] Opening via temp file {temp_file_path} successful")
try:
os.remove(temp_file_path) # Clean up temp file
except Exception as e_remove:
self._log(f"[Call {call_id}] Could not remove temp file {temp_file_path}: {e_remove}")
except Exception as e2:
self._log(f"[Call {call_id}] Opening via temp file failed: {str(e2)}")
if pil_image is None:
self._log(f"[Call {call_id}] Cannot open image, skipping")
continue
if pil_image.mode != 'RGB':
pil_image = pil_image.convert('RGB')
self._log(f"[Call {call_id}] Image converted to RGB mode")
# Optional padding, if self.min_size is set appropriately (currently 1024 from Editor)
# Consider if each image from ADV should be padded or if it's context-dependent
# For now, let's assume padding is desired if min_size is met.
if pil_image.width < self.min_size or pil_image.height < self.min_size:
self._log(
f"[Call {call_id}] Image size {pil_image.width}x{pil_image.height} is smaller than minimum {self.min_size}x{self.min_size}, padding needed")
pil_image = self._pad_image_to_minimum_size(pil_image)
img_array = np.array(pil_image).astype(np.float32) / 255.0
image_tensor = torch.from_numpy(img_array).unsqueeze(0) # Batch dimension
self._log(f"[Call {call_id}] Image processed from API response. Shape: {image_tensor.shape}")
break # Assuming one image per response for ADV node
except Exception as e:
self._log(f"[Call {call_id}] Error processing image from API response: {e}")
traceback.print_exc()
final_response_text = "\n".join(response_text_parts)
if image_tensor is None:
self._log(f"[Call {call_id}] No image found in API response parts.")
error_msg = "API Error: No image data in response" # More specific than Editor's default
image_tensor = self._create_error_image(error_msg) # Use the updated _create_error_image
if not final_response_text: final_response_text = error_msg # Keep this logic
return image_tensor, final_response_text
async def _generate_single_image_async(self, api_key, model_name_full, prompt_text, input_pil_image, temperature, max_retries, seed_val, call_id):
try:
# Use genai.Client for initialization, similar to FL_GeminiImageEditor
# Ensure API key is passed if required by genai.Client constructor
# Some SDK versions might use genai.configure(api_key=...) globally first,
# then genai.Client() without api_key arg.
# Given the error, direct client instantiation is safer.
try:
client_instance = genai.Client(api_key=api_key)
except TypeError: # Fallback if genai.Client() doesn't take api_key (older versions might not)
genai.configure(api_key=api_key) # Try global configure
client_instance = genai.Client() # Then instantiate
except AttributeError: # If genai.Client itself is not found, this is a deeper SDK issue.
self._log(f"[Call {call_id}] CRITICAL: genai.Client not found. Please check google-genai SDK installation and version.")
error_msg = f"Call {call_id} Error: genai.Client not found."
return self._create_error_image(error_msg), error_msg, call_id
actual_seed = seed_val if seed_val != 0 else random.randint(1, 0xffffffffffffffff)
self._log(f"[Call {call_id}] Using seed: {actual_seed} for prompt: '{prompt_text[:50]}...'")
# Use types.GenerateContentConfig as in FL_GeminiImageEditor
gen_config_params = {
"temperature": temperature,
"response_modalities": ['Text', 'Image'] # From FL_GeminiImageEditor
}
if actual_seed != 0:
gen_config_params["seed"] = actual_seed
gen_config_obj = types.GenerateContentConfig(**gen_config_params)
if actual_seed != 0:
current_seed_in_config = getattr(gen_config_obj, 'seed', None)
if current_seed_in_config != actual_seed:
self._log(f"[Call {call_id}] Warning: Seed {actual_seed} was specified. GenerateContentConfig has seed: {current_seed_in_config}. Ensure model supports seed via this config.")
contents = [prompt_text]
if input_pil_image:
contents.append(input_pil_image)
loop = asyncio.get_event_loop()
response = await loop.run_in_executor(
None,
lambda: self._call_gemini_api(client_instance, model_name_full, contents, gen_config_obj, 0, max_retries, call_id)
)
img_tensor, response_text = self._process_api_response(response, call_id)
return img_tensor, response_text, call_id # Return call_id to map results
except Exception as e:
self._log(f"[Call {call_id}] Error in async generation: {str(e)}")
error_msg = f"Call {call_id} Error: {str(e)}"
return self._create_error_image(error_msg), error_msg, call_id
def generate_images_advanced(self, inputcount, api_key, model, temperature, max_retries, prompt_1, image_1=None, seed=0, **kwargs): # image_1 now has default None
self.log_messages = []
if not api_key:
error_msg = "API key not provided."
self._log(error_msg)
# Ensure error image is created with the correct (updated) default size if needed
error_img_instance = self._create_error_image(error_msg)
return ([error_img_instance] * inputcount, error_msg)
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
tasks = []
for i in range(1, inputcount + 1):
current_prompt = prompt_1 if i == 1 else kwargs.get(f"prompt_{i}", f"Default prompt for image {i}")
# Handle image_1 being optional, and subsequent image_i from kwargs
current_image_tensor = None
if i == 1:
current_image_tensor = image_1 # This could be None if not provided
else:
current_image_tensor = kwargs.get(f"image_{i}") # This could also be None
current_pil_image = self._process_tensor_to_pil(current_image_tensor, f"InputImage{i}") # This function already handles None
# Increment seed for each task if a non-zero base seed is provided
current_seed = seed + (i-1) if seed != 0 else 0
tasks.append(self._generate_single_image_async(
api_key, model, current_prompt, current_pil_image,
temperature, max_retries, current_seed, str(i)
))
try:
results_with_id = loop.run_until_complete(asyncio.gather(*tasks))
finally:
loop.close()
# Sort results by call_id to maintain order, though gather usually preserves submission order
results_with_id.sort(key=lambda x: int(x[2]))
output_images = []
output_texts = []
for img_tensor, response_text, call_id_res in results_with_id:
output_images.append(img_tensor)
output_texts.append(f"Response for Input {call_id_res}:\n{response_text}")
batched_images = torch.cat(output_images, dim=0) if output_images else self._create_error_image("No images generated")
combined_responses = "\n\n".join(output_texts)
# Prepend logs
final_log_output = "Processing Logs:\n" + "\n".join(self.log_messages) + "\n\n" + combined_responses
return (batched_images, final_log_output)
+48
View File
@@ -0,0 +1,48 @@
import torch
from nodes import ImageBatch # Assuming ImageBatch is a ComfyUI core node or accessible
class FL_ImageBatch: # Renamed class
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"inputcount": ("INT", {"default": 2, "min": 2, "max": 100, "step": 1}),
"image_1": ("IMAGE",),
"image_2": ("IMAGE",),
},
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("images",)
FUNCTION = "combine_images"
CATEGORY = "🏵️Fill Nodes/Image" # Kept category, can be changed if needed
DESCRIPTION = """
FL Image Batch allows combining multiple images into a batch.
The number of image inputs can be dynamically adjusted.
""" # Updated description
def combine_images(self, inputcount, **kwargs):
image_batch_node = ImageBatch()
if "image_1" not in kwargs:
return (torch.empty(0),)
image = kwargs["image_1"]
if inputcount == 1:
return (image,)
for c in range(1, inputcount):
image_key = f"image_{c + 1}"
if image_key in kwargs:
new_image = kwargs[image_key]
if image.ndim == 3:
image = image.unsqueeze(0)
if new_image.ndim == 3:
new_image = new_image.unsqueeze(0)
image, = image_batch_node.batch(image, new_image)
else:
print(f"Warning: {image_key} not found in inputs for FL_ImageBatch.")
pass
return (image,)
+2 -2
View File
@@ -582,7 +582,7 @@ class FL_PixVerseAPI:
# Convert PIL image to bytes
img_byte_arr = io.BytesIO()
pil_image.save(img_byte_arr, format='PNG')
pil_image.save(img_byte_arr, format='JPEG', quality=80) # Changed to JPEG with quality 80
img_byte_arr.seek(0) # Reset pointer to beginning of buffer
upload_url = "https://app-api.pixverse.ai/openapi/v2/image/upload"
@@ -592,7 +592,7 @@ class FL_PixVerseAPI:
}
# Send bytes directly without saving to disk
files = {'image': (f'{image_type}.png', img_byte_arr, 'image/png')}
files = {'image': (f'{image_type}.jpg', img_byte_arr, 'image/jpeg')} # Changed to jpg and image/jpeg
upload_response = requests.post(upload_url, headers=upload_headers, files=files)
if upload_response.status_code != 200:
+1 -1
View File
@@ -7,7 +7,7 @@ glfw
scipy >=1.13.1
requests
aiohttp
moviepy
moviepy==1.0.3
matplotlib
reportlab
openai
+91
View File
@@ -0,0 +1,91 @@
import { app } from "../../../scripts/app.js";
app.registerExtension({
name: "FillNodes.GeminiImageGenADV", // Unique name for the extension
async beforeRegisterNodeDef(nodeType, nodeData, app) {
// Check if this is the correct node we want to modify
if (nodeData.name === "FL_GeminiImageGenADV") {
// This function is called when a new node of this type is created
nodeType.prototype.onNodeCreated = function () {
this._image_type = "IMAGE";
this._prompt_type = "STRING";
// Add the "Update inputs" button to this node's widget list
this.addWidget("button", "Update inputs", null, () => {
if (!this.inputs) {
this.inputs = [];
}
const inputCountWidget = this.widgets.find(w => w.name === "inputcount");
if (!inputCountWidget) {
console.error("FL_GeminiImageGenADV: 'inputcount' widget not found on this node!");
return;
}
const target_pairs = parseInt(inputCountWidget.value);
// Current number of *pairs* (image_1/prompt_1 is one pair)
// image_1 and prompt_1 are required and always present.
// So, count additional pairs starting from image_2/prompt_2
let current_pairs = 1; // Start with 1 for the required image_1/prompt_1
for(let i = 0; i < this.inputs.length; i++) {
if (this.inputs[i].name === `image_${current_pairs + 1}`) {
current_pairs++;
}
}
if (target_pairs === current_pairs) {
return; // No change needed
}
if (target_pairs < current_pairs) {
// Reduce the number of pairs
const pairs_to_remove = current_pairs - target_pairs;
for (let i = 0; i < pairs_to_remove; i++) {
// Remove the last prompt and then the last image input of the highest index pair
let last_prompt_index = -1;
let last_image_index = -1;
const pair_num_to_remove = current_pairs - i;
for (let j = this.inputs.length - 1; j >= 0; j--) {
if (this.inputs[j].name === `prompt_${pair_num_to_remove}`) {
last_prompt_index = j;
} else if (this.inputs[j].name === `image_${pair_num_to_remove}`) {
last_image_index = j;
}
}
if (last_prompt_index !== -1) this.removeInput(last_prompt_index);
if (last_image_index !== -1 && last_image_index < this.inputs.length) { // Check if index is still valid after prompt removal
// Need to re-check index if prompt was before image and got removed
let current_last_image_index = -1;
for (let k = this.inputs.length - 1; k >=0; k--) {
if (this.inputs[k].name === `image_${pair_num_to_remove}`) {
current_last_image_index = k;
break;
}
}
if(current_last_image_index !== -1) this.removeInput(current_last_image_index);
} else if (last_image_index !== -1) { // If prompt was after image or not found
this.removeInput(last_image_index);
}
}
} else {
// Increase the number of pairs
// Start from current_pairs + 1 because image_1/prompt_1 up to image_{current_pairs}/prompt_{current_pairs} exist
for (let i = current_pairs + 1; i <= target_pairs; ++i) {
this.addInput(`image_${i}`, this._image_type);
this.addInput(`prompt_${i}`, this._prompt_type, { multiline: false, default: `prompt for image ${i}` });
}
}
// Refresh the node's appearance
this.setDirtyCanvas(true, true);
});
// Initial call to sync inputs if loaded from workflow with different inputcount
// Ensure widgets are available before calling
if (this.widgets && this.widgets.find(w => w.name === "inputcount")) {
this.widgets.find(w => w.name === "Update inputs").callback();
}
};
}
},
});
+57
View File
@@ -0,0 +1,57 @@
import { app } from "../../../scripts/app.js";
app.registerExtension({
name: "FillNodes.ImageBatch", // Unique name for this new node's extension
async beforeRegisterNodeDef(nodeType, nodeData, app) {
// Check if this is the "FL_ImageBatch" node
if (nodeData.name === "FL_ImageBatch") {
// This function is called when a new "FL_ImageBatch" node is created
nodeType.prototype.onNodeCreated = function () {
this._type = "IMAGE"; // Define the type of input we are managing
// Add the "Update inputs" button to this node's widget list
this.addWidget("button", "Update inputs", null, () => {
if (!this.inputs) {
this.inputs = [];
}
// Find the 'inputcount' widget (defined in Python) to get the target number
const inputCountWidget = this.widgets.find(w => w.name === "inputcount");
if (!inputCountWidget) {
console.error("FL_ImageBatch: 'inputcount' widget not found on this node!");
return;
}
const target_number_of_inputs = parseInt(inputCountWidget.value); // Ensure it's a number
// Count current IMAGE inputs that match our naming convention
const num_inputs = this.inputs.filter(input => input.type === this._type && input.name.startsWith("image_")).length;
if (target_number_of_inputs === num_inputs) {
return; // No change needed
}
if (target_number_of_inputs < num_inputs) {
// Reduce the number of inputs
const inputs_to_remove = num_inputs - target_number_of_inputs;
for (let i = 0; i < inputs_to_remove; i++) {
// Remove the last IMAGE input matching our convention
for (let j = this.inputs.length - 1; j >= 0; j--) {
if (this.inputs[j].type === this._type && this.inputs[j].name.startsWith("image_")) {
this.removeInput(j);
break; // Exit inner loop once an input is removed
}
}
}
} else {
// Increase the number of inputs
for (let i = num_inputs + 1; i <= target_number_of_inputs; ++i) {
this.addInput(`image_${i}`, this._type);
}
}
// Refresh the node's appearance
this.setDirtyCanvas(true, true);
});
};
}
},
});