Streamlined and cleaned up code

This commit is contained in:
Aryan185
2026-01-02 13:54:01 +00:00
parent a42bcfdbba
commit 0d907bf82a
6 changed files with 282 additions and 534 deletions
+41 -78
View File
@@ -1,16 +1,13 @@
import os
import io
import numpy as np
import torch
import wave
import torch
import numpy as np
from PIL import Image
from typing import Optional
from google import genai
from google.genai import types
class GeminiChatNode:
"""ComfyUI Node for Gemini API Chat with optional image and audio input"""
@classmethod
def INPUT_TYPES(cls):
return {
@@ -24,7 +21,7 @@ class GeminiChatNode:
},
"optional": {
"system_instruction": ("STRING", {"multiline": True, "default": ""}),
"thinking_budget": ("INT", {"default": -1, "min": -1, "max": 24576, "step": 1}),
"thinking_budget": ("INT", {"default": 0, "min": -1, "max": 24576, "step": 1}),
"image": ("IMAGE",),
"audio": ("AUDIO",),
}
@@ -35,94 +32,62 @@ class GeminiChatNode:
FUNCTION = "generate"
CATEGORY = "text/generation"
def audio_to_bytes(self, audio):
if isinstance(audio, dict):
audio_data = audio.get("waveform")
sr = audio.get("sample_rate", 44100)
elif isinstance(audio, (list, tuple)) and len(audio) >= 2:
audio_data, sr = audio[0], audio[1]
else:
raise ValueError(f"Invalid audio input format: {type(audio)}")
if audio_data is None:
raise ValueError("Missing audio data")
if isinstance(audio_data, torch.Tensor):
audio_data = audio_data.cpu().numpy()
# Convert to WAV bytes
audio_data = np.squeeze(audio_data)
if audio_data.dtype in [np.float32, np.float64]:
audio_data = np.clip(audio_data, -1.0, 1.0)
audio_data = (audio_data * 32767).astype(np.int16)
wav_buffer = io.BytesIO()
with wave.open(wav_buffer, 'wb') as wav_file:
wav_file.setnchannels(1)
wav_file.setsampwidth(2)
wav_file.setframerate(int(sr))
wav_file.writeframes(audio_data.tobytes())
return wav_buffer.getvalue()
def generate(self, prompt: str, model: str, temperature: float, thinking: bool, seed: int, api_key: str,
system_instruction: Optional[str] = None, thinking_budget: int = -1,
image: Optional[torch.Tensor] = None, audio: Optional[dict] = None) -> tuple:
def generate(self, prompt, model, temperature, thinking, seed, api_key,
system_instruction=None, thinking_budget=-1, image=None, audio=None):
key = api_key.strip() or os.environ.get("GEMINI_API_KEY")
if not key:
raise ValueError("Error: No API key provided.")
if not key: raise ValueError("Error: No API key provided.")
# Initialize client and build parts
client = genai.Client(api_key=key, http_options=types.HttpOptions(retry_options=types.HttpRetryOptions(attempts=3, jitter=10)))
client = genai.Client(api_key=key, http_options={'api_version': 'v1beta'})
parts = [types.Part.from_text(text=prompt)]
# Handle image input
if image is not None:
img_array = image.cpu().numpy() if isinstance(image, torch.Tensor) else image
if len(img_array.shape) == 4:
img_array = img_array[0]
if img_array.dtype in [np.float32, np.float64]:
img_array = (img_array * 255).astype(np.uint8)
buffered = io.BytesIO()
Image.fromarray(img_array).save(buffered, format="PNG")
parts.append(types.Part.from_bytes(mime_type="image/png", data=buffered.getvalue()))
arr = (image[0].cpu().numpy() * 255).astype(np.uint8)
buf = io.BytesIO()
Image.fromarray(arr).save(buf, format="PNG")
parts.append(types.Part.from_bytes(mime_type="image/png", data=buf.getvalue()))
# Handle audio input
if audio is not None:
audio_bytes = self.audio_to_bytes(audio)
parts.append(types.Part.from_bytes(mime_type="audio/wav", data=audio_bytes))
wf = audio.get("waveform") if isinstance(audio, dict) else audio[0]
sr = audio.get("sample_rate", 44100) if isinstance(audio, dict) else audio[1]
wf = wf.cpu().numpy() if isinstance(wf, torch.Tensor) else wf
if wf.ndim > 1: wf = wf.mean(axis=0) if wf.shape[0] > 1 else wf.squeeze()
wf_int16 = (np.clip(wf, -1, 1) * 32767).astype(np.int16)
buf = io.BytesIO()
with wave.open(buf, 'wb') as w:
w.setnchannels(1); w.setsampwidth(2); w.setframerate(sr)
w.writeframes(wf_int16.tobytes())
parts.append(types.Part.from_bytes(mime_type="audio/wav", data=buf.getvalue()))
model_lower = model.lower()
t_config = None
if "gemini-2.0" in model_lower:
print("Gemini-2.0 models do not support thinking - disabling thinking config")
final_thinking_budget = None
elif not thinking:
final_thinking_budget = 0
if "gemini-2.5-pro" in model_lower:
print("Gemini-2.5-Pro cannot have thinking turned off - defaulting thinking budget to -1")
final_thinking_budget = -1
else:
final_thinking_budget = thinking_budget
if "gemini-2.5-pro" in model_lower and final_thinking_budget == 0:
print("Gemini-2.5-Pro cannot have thinking turned off - defaulting thinking budget to -1")
final_thinking_budget = -1
final_budget = 0 # Default disabled
if not thinking:
if "gemini-2.5-pro" in model_lower or "gemini-3-pro-preview" in model_lower:
print("Pro models cannot have thinking turned off - defaulting thinking budget to -1")
final_budget = -1
else:
final_budget = thinking_budget
if ("gemini-2.5-pro" in model_lower or "gemini-3-pro-preview" in model_lower) and final_budget == 0:
print("Pro models cannot have thinking turned off - defaulting thinking budget to -1")
final_budget = -1
t_config = types.ThinkingConfig(thinking_budget=final_budget)
config = types.GenerateContentConfig(
temperature=temperature,
seed=seed,
response_mime_type="text/plain"
system_instruction=system_instruction.strip() if system_instruction else None,
thinking_config=t_config
)
if "gemini-2.0" not in model_lower:
config.thinking_config = types.ThinkingConfig(thinking_budget=final_thinking_budget)
if system_instruction and system_instruction.strip():
config.system_instruction = [types.Part.from_text(text=system_instruction.strip())]
response = client.models.generate_content(
model=model,
contents=[types.Content(role="user", parts=parts)],
@@ -130,8 +95,6 @@ class GeminiChatNode:
)
return (response.text,)
# Node mappings
NODE_CLASS_MAPPINGS = {"GeminiChatNode": GeminiChatNode}
NODE_DISPLAY_NAME_MAPPINGS = {"GeminiChatNode": "Gemini Chat"}
+65 -138
View File
@@ -1,16 +1,15 @@
import base64
import os
import io
import re
import json
import numpy as np
import base64
import torch
import numpy as np
from PIL import Image
from google import genai
from google.genai import types
class GeminiSegmentationNode:
"""ComfyUI Node for Gemini API Image Segmentation"""
@classmethod
def INPUT_TYPES(cls):
return {
@@ -33,153 +32,81 @@ class GeminiSegmentationNode:
FUNCTION = "generate_segmentation"
CATEGORY = "image/generation"
def generate_segmentation(self, image: torch.Tensor, segment_prompt: str, model: str,
temperature: float, thinking: bool, seed: int, api_key: str,
thinking_budget: int = 0) -> tuple:
def generate_segmentation(self, image, segment_prompt, model, temperature, thinking, seed, api_key, thinking_budget=0):
key = api_key.strip() or os.environ.get("GEMINI_API_KEY") or os.environ.get("GOOGLE_API_KEY")
if not key:
raise ValueError("Error: No API key provided. Set GEMINI_API_KEY or GOOGLE_API_KEY environment variable, or provide it in the node.")
if not key: raise ValueError("API Key missing")
client = genai.Client(api_key=key, http_options={'api_version': 'v1beta'})
img_array = image.cpu().numpy() if isinstance(image, torch.Tensor) else image
if len(img_array.shape) == 4:
img_array = img_array[0] # Remove batch dimension
if img_array.dtype in [np.float32, np.float64]:
img_array = (img_array * 255).astype(np.uint8)
img_np = (image[0].cpu().numpy() * 255).astype(np.uint8)
orig_img = Image.fromarray(img_np)
orig_w, orig_h = orig_img.size
original_image = Image.fromarray(img_array).convert('RGB')
original_width, original_height = original_image.size
scale = min(1024 / orig_w, 1024 / orig_h)
proc_img = orig_img.resize((int(orig_w * scale), int(orig_h * scale)), Image.Resampling.LANCZOS) if scale < 1 else orig_img
pw, ph = proc_img.size
max_size = 1024
scale = min(max_size / original_width, max_size / original_height)
if scale < 1:
new_width = int(original_width * scale)
new_height = int(original_height * scale)
processed_image = original_image.resize((new_width, new_height), Image.Resampling.LANCZOS)
else:
processed_image = original_image
buffer = io.BytesIO()
processed_image.save(buffer, format='PNG')
image_data = buffer.getvalue()
base_prompt = f"Give the segmentation masks for {segment_prompt}. Output a JSON list of segmentation masks where each entry contains the 2D bounding box in the key \"box_2d\", the segmentation mask in key \"mask\", and the text label in the key \"label\". Use descriptive labels."
client = genai.Client(api_key=key, http_options=types.HttpOptions(retry_options=types.HttpRetryOptions(attempts=3, jitter=10)))
parts = [
types.Part.from_bytes(mime_type="image/png", data=image_data),
types.Part.from_text(text=base_prompt)
]
model_lower = model.lower()
if "gemini-2.0" in model_lower:
print("Gemini-2.0 models do not support thinking - disabling thinking config")
final_thinking_budget = None
elif not thinking:
final_thinking_budget = 0
if "gemini-2.5-pro" in model_lower:
print("Gemini-2.5-Pro cannot have thinking turned off - defaulting thinking budget to -1")
final_thinking_budget = -1
else:
final_thinking_budget = thinking_budget
if "gemini-2.5-pro" in model_lower and final_thinking_budget == 0:
print("Gemini-2.5-Pro cannot have thinking turned off - defaulting thinking budget to -1")
final_thinking_budget = -1
config = types.GenerateContentConfig(
temperature=temperature,
seed=seed,
response_mime_type="text/plain"
img_buf = io.BytesIO()
proc_img.save(img_buf, format='PNG')
t_config = None
if "gemini-2.0" not in model.lower():
budget = thinking_budget if thinking else 0
if "gemini-2.5-pro" in model.lower() and budget <= 0:
print("Gemini-2.5-Pro enforces thinking. Defaulting to auto (-1).")
budget = -1
t_config = types.ThinkingConfig(thinking_budget=budget)
prompt = f'Give the segmentation masks for {segment_prompt}. Output a JSON list of segmentation masks where each entry contains the 2D bounding box in the key "box_2d", the segmentation mask in key "mask", and the text label in the key "label".'
response = client.models.generate_content(
model=model,
contents=[types.Content(role="user", parts=[
types.Part.from_bytes(mime_type="image/png", data=img_buf.getvalue()),
types.Part.from_text(text=prompt)
])],
config=types.GenerateContentConfig(temperature=temperature, seed=seed, thinking_config=t_config)
)
if "gemini-2.0" not in model_lower:
config.thinking_config = types.ThinkingConfig(thinking_budget=final_thinking_budget)
# Generate content
try:
response = client.models.generate_content(
model=model,
contents=[types.Content(role="user", parts=parts)],
config=config
)
response_text = response.text
if '```json' in response_text:
response_text = response_text.split('```json')[1].split('```')[0]
segments = json.loads(response_text)
txt = response.text
if "```json" in txt: txt = re.search(r"```json\n(.*)\n```", txt, re.DOTALL).group(1)
segments = json.loads(txt)
except Exception as e:
raise RuntimeError(f"Error calling Gemini API: {str(e)}")
# Create mask from segments
proc_width, proc_height = processed_image.size
mask_image = Image.new('L', (proc_width, proc_height), 0)
# Sort segments by size (largest first)
segments_with_size = []
for segment in segments:
box_2d = segment['box_2d']
ymin, xmin, ymax, xmax = box_2d
w = (xmax - xmin) / 1000
h = (ymax - ymin) / 1000
segments_with_size.append((segment, w * h))
raise RuntimeError(f"Gemini API Error: {e}")
final_mask = np.zeros((ph, pw), dtype=np.uint8)
segments_with_size.sort(key=lambda x: x[1], reverse=True)
# Process each segment
for i, (segment, _) in enumerate(segments_with_size):
for seg in segments:
try:
box_2d = segment['box_2d']
ymin, xmin, ymax, xmax = box_2d
# Calculate integer coords
ymin, xmin, ymax, xmax = seg['box_2d']
x1, y1 = int(xmin * pw / 1000), int(ymin * ph / 1000)
x2, y2 = int(xmax * pw / 1000), int(ymax * ph / 1000)
w, h = x2 - x1, y2 - y1
x = int(xmin / 1000 * proc_width)
y = int(ymin / 1000 * proc_height)
w = int((xmax - xmin) / 1000 * proc_width)
h = int((ymax - ymin) / 1000 * proc_height)
if w <= 0 or h <= 0: continue
# Decode & Resize Patch
mask_str = seg['mask'].split(",")[1] if "data:image" in seg['mask'] else seg['mask']
patch = Image.open(io.BytesIO(base64.b64decode(mask_str))).convert('L')
mask_data = segment['mask']
if patch.size != (w, h):
patch = patch.resize((w, h), Image.Resampling.NEAREST)
if isinstance(mask_data, str):
if mask_data.startswith('data:image'):
mask_data = mask_data.split(',')[1]
mask_bytes = base64.b64decode(mask_data)
mask_img = Image.open(io.BytesIO(mask_bytes)).convert('L')
else:
continue
patch_arr = np.array(patch)
patch_arr = np.where(patch_arr > 128, 255, 0).astype(np.uint8)
if mask_img.size != (w, h):
mask_img = mask_img.resize((w, h), Image.Resampling.LANCZOS)
mask_array = list(mask_img.getdata())
final_pixels = [255 if alpha > 128 else 0 for alpha in mask_array]
segment_mask = Image.new('L', (w, h))
segment_mask.putdata(final_pixels)
if x + w <= proc_width and y + h <= proc_height and x >= 0 and y >= 0:
region = mask_image.crop((x, y, x + w, y + h))
region_pixels = list(region.getdata())
segment_pixels = list(segment_mask.getdata())
combined_pixels = [max(r, s) for r, s in zip(region_pixels, segment_pixels)]
combined_region = Image.new('L', (w, h))
combined_region.putdata(combined_pixels)
mask_image.paste(combined_region, (x, y))
except Exception:
continue
if processed_image.size != original_image.size:
mask_image = mask_image.resize(original_image.size, Image.Resampling.LANCZOS)
# Convert PIL mask to ComfyUI mask format
mask_array = np.array(mask_image, dtype=np.float32) / 255.0
mask_tensor = torch.from_numpy(mask_array).unsqueeze(0) # Add batch dimension
return (mask_tensor,)
# Safe slicing to handle potential boundary issues
target_slice = final_mask[y1:y2, x1:x2]
if target_slice.shape == patch_arr.shape:
np.maximum(target_slice, patch_arr, out=target_slice)
except Exception: continue
if (pw, ph) != (orig_w, orig_h):
final_mask = np.array(Image.fromarray(final_mask).resize((orig_w, orig_h), Image.Resampling.NEAREST))
return (torch.from_numpy(final_mask.astype(np.float32) / 255.0).unsqueeze(0),)
# Node mappings
NODE_CLASS_MAPPINGS = {"GeminiSegmentationNode": GeminiSegmentationNode}
NODE_DISPLAY_NAME_MAPPINGS = {"GeminiSegmentationNode": "Gemini Segmentation"}
+32 -43
View File
@@ -30,64 +30,53 @@ class GeminiTTSNode:
if not text.strip():
raise ValueError("Text input cannot be empty.")
key = api_key.strip() or os.environ.get("GEMINI_TTS_API_KEY")
key = api_key.strip() or os.environ.get("GEMINI_API_KEY")
if not key:
raise ValueError("No API key provided.")
client = Client(
api_key=key,
http_options=types.HttpOptions(
retry_options=types.HttpRetryOptions(attempts=10, jitter=10)
client = Client(api_key=key)
final_prompt = text
if system_prompt.strip():
final_prompt = f"{system_prompt.strip()}\n\n{text}"
speech_config = types.SpeechConfig(
voice_config=types.VoiceConfig(
prebuilt_voice_config=types.PrebuiltVoiceConfig(voice_name=voice_id)
)
)
# Build prompt
prompt_text = text
if system_prompt.strip():
prompt_text = system_prompt.strip() + ":\n\n\"" + text + "\""
contents = [types.Content(role="user", parts=[types.Part.from_text(text=prompt_text)])]
config = types.GenerateContentConfig(
temperature=temperature,
seed=seed,
response_modalities=["audio"],
speech_config=types.SpeechConfig(
voice_config=types.VoiceConfig(
prebuilt_voice_config=types.PrebuiltVoiceConfig(voice_name=voice_id)
)
),
response_modalities=["AUDIO"],
speech_config=speech_config,
)
# Generate audio - collect raw PCM chunks
audio_data = b""
for chunk in client.models.generate_content_stream(
model=model,
contents=contents,
config=config
):
if (chunk.candidates and chunk.candidates[0].content and
chunk.candidates[0].content.parts and
chunk.candidates[0].content.parts[0].inline_data):
inline_data = chunk.candidates[0].content.parts[0].inline_data
audio_data += inline_data.data
if not audio_data:
raise ValueError("No audio data received from API.")
# Convert raw PCM to waveform tensor
waveform = torch.frombuffer(bytearray(audio_data), dtype=torch.int16)
try:
response = client.models.generate_content(
model=model,
contents=final_prompt,
config=config
)
except Exception as e:
raise RuntimeError(f"Gemini API Error: {str(e)}")
try:
inline_data = response.candidates[0].content.parts[0].inline_data
audio_bytes = inline_data.data
except (AttributeError, IndexError, TypeError):
raise ValueError("API returned a response, but it contained no audio data.")
waveform = torch.frombuffer(bytearray(audio_bytes), dtype=torch.int16)
waveform = waveform.to(torch.float32) / 32768.0
waveform = waveform.unsqueeze(0)
sample_rate = 24000
waveform = waveform.unsqueeze(0).unsqueeze(0)
return ({"waveform": waveform.unsqueeze(0), "sample_rate": sample_rate},)
return ({"waveform": waveform, "sample_rate": 24000},)
@classmethod
def IS_CHANGED(cls, **kwargs):
return f"{kwargs.get('text', '')}-{kwargs.get('voice_id', '')}-{kwargs.get('temperature', 1.0)}-{kwargs.get('model', '')}-{kwargs.get('seed', 69)}-{kwargs.get('system_prompt', '')}"
def IS_CHANGED(cls, seed, **kwargs):
return seed
NODE_CLASS_MAPPINGS = {"GeminiTTSNode": GeminiTTSNode}
NODE_DISPLAY_NAME_MAPPINGS = {"GeminiTTSNode": "Gemini TTS"}
+36 -79
View File
@@ -1,4 +1,5 @@
import os
import io
import torch
import numpy as np
from PIL import Image
@@ -6,12 +7,11 @@ from google import genai
from google.genai import types
class GoogleImagenNode:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"prompt": ("STRING", {"multiline": True, "default": "A majestic lion in the savanna"}),
"prompt": ("STRING", {"multiline": True}),
"api_key": ("STRING", {"multiline": False, "default": ""}),
"model": (["models/imagen-4.0-ultra-generate-001", "models/imagen-4.0-generate-001", "models/imagen-4.0-fast-generate-001", "models/imagen-3.0-generate-002"],),
"number_of_images": ("INT", {"default": 1, "min": 1, "max": 4, "step": 1}),
@@ -30,93 +30,50 @@ class GoogleImagenNode:
FUNCTION = "generate_images"
CATEGORY = "image/generation"
def pil_to_tensor(self, images):
if not isinstance(images, list):
images = [images]
tensors = []
for image in images:
if image.mode != 'RGB':
image = image.convert('RGB')
array = np.array(image).astype(np.float32) / 255.0
tensor = torch.from_numpy(array)
tensors.append(tensor)
return torch.stack(tensors)
def generate_images(self, prompt, api_key, model, number_of_images, aspect_ratio, image_size, seed, guidance_scale, negative_prompt=""):
key = api_key.strip() or os.environ.get("GEMINI_API_KEY")
if not key: raise ValueError("No API key provided.")
client = genai.Client(api_key=key)
config = types.GenerateImagesConfig(
number_of_images=number_of_images,
aspect_ratio=aspect_ratio,
guidance_scale=guidance_scale,
negative_prompt=negative_prompt.strip() if negative_prompt.strip() else None
)
if "imagen-4.0" in model and "fast" not in model:
config.image_size = image_size
try:
key = api_key.strip() or os.environ.get("GEMINI_API_KEY")
if not key:
raise ValueError("No API key provided.")
result = client.models.generate_images(model=model, prompt=prompt, config=config)
if not result.generated_images: raise ValueError("No images generated")
client = genai.Client(api_key=key)
config_params = {
"number_of_images": number_of_images,
"aspect_ratio": aspect_ratio,
"add_watermark": False,
"seed": seed,
"guidance_scale": guidance_scale,
}
if negative_prompt and negative_prompt.strip():
config_params["negative_prompt"] = negative_prompt
config = types.GenerateImagesConfig(**config_params)
if model in ["models/imagen-4.0-ultra-generate-001", "models/imagen-4.0-generate-001"]:
config.image_size = image_size
else:
print("2K resolution not supported for this model, using default image size")
result = client.models.generate_images(
model=model,
prompt=prompt,
config=config
)
if not result.generated_images:
raise ValueError("No images generated by the API")
pil_images = []
for generated_image in result.generated_images:
image_data = generated_image.image
tensors = []
for item in result.generated_images:
img_data = item.image
if hasattr(image_data, 'mode') and hasattr(image_data, 'size'):
pil_images.append(image_data)
elif hasattr(image_data, '_pil_image'):
pil_images.append(image_data._pil_image)
elif hasattr(image_data, 'show'):
try:
from io import BytesIO
buffer = BytesIO()
image_data.save(buffer, format='PNG')
buffer.seek(0)
pil_image = Image.open(buffer)
pil_images.append(pil_image)
except:
pil_images.append(Image.new('RGB', (512, 512), color='gray'))
elif hasattr(image_data, 'read') or isinstance(image_data, bytes):
from io import BytesIO
image_bytes = image_data.read() if hasattr(image_data, 'read') else image_data
pil_images.append(Image.open(BytesIO(image_bytes)))
if hasattr(img_data, "image_bytes"):
pil_img = Image.open(io.BytesIO(img_data.image_bytes))
elif hasattr(img_data, "convert"):
pil_img = img_data
else:
try:
pil_images.append(Image.open(image_data))
except:
pil_images.append(Image.new('RGB', (512, 512), color='gray'))
# Fallback for raw bytes
pil_img = Image.open(io.BytesIO(img_data))
pil_img = pil_img.convert("RGB")
tensors.append(torch.from_numpy(np.array(pil_img).astype(np.float32) / 255.0))
return (self.pil_to_tensor(pil_images),)
return (torch.stack(tensors),)
except Exception as e:
print(f"Google Imagen Error: {str(e)}")
error_image = Image.new('RGB', (512, 512), color='black')
return (self.pil_to_tensor([error_image]),)
print(f"Google Imagen Error: {e}")
raise RuntimeError(f"Google Imagen Error: {e}")
@classmethod
def IS_CHANGED(cls, **kwargs):
return f"{kwargs.get('prompt', '')}-{kwargs.get('model', '')}-{kwargs.get('number_of_images', 1)}-{kwargs.get('aspect_ratio', '1:1')}-{kwargs.get('image_size', '1K')}"
return float("nan")
NODE_CLASS_MAPPINGS = {"GoogleImagenNode": GoogleImagenNode}
NODE_DISPLAY_NAME_MAPPINGS = {"GoogleImagenNode": "Google Imagen Generator"}
+49 -118
View File
@@ -1,17 +1,15 @@
import os
import json
import io
import base64
import tempfile
import torch
import numpy as np
from PIL import Image
import base64
from io import BytesIO
from google import genai
from google.genai import types
from google.genai.types import RawReferenceImage, MaskReferenceImage
class GoogleImagenEditNode:
@classmethod
def INPUT_TYPES(cls):
return {
@@ -19,7 +17,6 @@ class GoogleImagenEditNode:
"image": ("IMAGE",),
"mask": ("MASK",),
"prompt": ("STRING", {"multiline": True, "default": "Edit this image"}),
"negative_prompt": ("STRING", {"multiline": True, "default": ""}),
"project_id": ("STRING", {"multiline": False, "default": ""}),
"location": (["global", "us-central1", "us-east1", "us-east4", "us-east5", "us-south1", "us-west1", "us-west2", "us-west3", "us-west4", "northamerica-northeast1", "northamerica-northeast2", "southamerica-east1", "southamerica-west1", "africa-south1", "europe-west1", "europe-north1", "europe-west2", "europe-west3", "europe-west4", "europe-west6", "europe-west8", "europe-west9", "europe-west12", "europe-southwest1", "europe-central2", "asia-east1", "asia-east2", "asia-northeast1", "asia-northeast2", "asia-northeast3", "asia-south1", "asia-south2", "asia-southeast1", "asia-southeast2", "australia-southeast1", "australia-southeast2", "me-central1", "me-central2", "me-west1"], {"default": "us-central1"}),
"service_account": ("STRING", {"multiline": True, "default": ""}),
@@ -29,7 +26,9 @@ class GoogleImagenEditNode:
"base_steps": ("INT", {"default": 50, "min": 10, "max": 100, "step": 1}),
"guidance_scale": ("FLOAT", {"default": 7.5, "min": 1.0, "max": 20.0, "step": 0.1}),
"mask_dilation": ("FLOAT", {"default": 0.03, "min": 0.0, "max": 1.0, "step": 0.01}),
},
"optional": {
"negative_prompt": ("STRING", {"multiline": True, "default": ""}),
}
}
@@ -38,139 +37,71 @@ class GoogleImagenEditNode:
FUNCTION = "edit_image"
CATEGORY = "image/edit"
def setup_client(self, service_account_json, project_id, location):
"""Setup Vertex AI client with service account JSON content"""
if not service_account_json.strip():
raise ValueError("Service account JSON content is required.")
def edit_image(self, image, mask, prompt, project_id, location, service_account,
edit_mode, number_of_images, seed, base_steps, guidance_scale, mask_dilation, negative_prompt=""):
if not project_id.strip():
raise ValueError("Project ID is required.")
creds_file = tempfile.NamedTemporaryFile(mode='w', suffix='.json', delete=False)
creds_file.write(service_account.strip())
creds_file.close()
os.environ['GOOGLE_APPLICATION_CREDENTIALS'] = creds_file.name
# Validate and write JSON content to temporary file
try:
json.loads(service_account_json) # Validate JSON format
except json.JSONDecodeError as e:
raise ValueError(f"Invalid JSON content: {str(e)}")
# Create temporary file with JSON content
temp_file = tempfile.NamedTemporaryFile(mode='w', suffix='.json', delete=False)
temp_file.write(service_account_json.strip())
temp_file.close()
# Set credentials path
os.environ['GOOGLE_APPLICATION_CREDENTIALS'] = temp_file.name
return genai.Client(vertexai=True, project=project_id.strip(), location=location.strip())
def tensor_to_pil(self, tensor):
array = (tensor.cpu().numpy() * 255).astype(np.uint8)
return Image.fromarray(array)
def mask_to_pil(self, mask):
if len(mask.shape) == 3 and mask.shape[0] == 1:
mask = mask.squeeze(0)
array = (mask.cpu().numpy() * 255).astype(np.uint8)
return Image.fromarray(array, mode='L')
def pil_to_tensor(self, images):
if not isinstance(images, list):
images = [images]
tensors = []
for image in images:
if image.mode != 'RGB':
image = image.convert('RGB')
array = np.array(image).astype(np.float32) / 255.0
tensors.append(torch.from_numpy(array))
return torch.stack(tensors)
def edit_image(self, image, mask, prompt, project_id, location, service_account, edit_mode, number_of_images, negative_prompt, seed, base_steps, guidance_scale, mask_dilation):
try:
# Initialize Vertex AI client with improved authentication
client = self.setup_client(service_account, project_id, location)
client = genai.Client(vertexai=True, project=project_id.strip(), location=location.strip())
input_image = self.tensor_to_pil(image[0])
input_mask = self.mask_to_pil(mask)
def to_b64(img):
b = io.BytesIO()
img.save(b, format='PNG')
return base64.b64encode(b.getvalue()).decode('utf-8')
img_pil = Image.fromarray((image[0].cpu().numpy() * 255).astype(np.uint8))
img_buffer = BytesIO()
input_image.save(img_buffer, format='PNG')
img_b64 = base64.b64encode(img_buffer.getvalue()).decode('utf-8')
mask_buffer = BytesIO()
input_mask.save(mask_buffer, format='PNG')
mask_b64 = base64.b64encode(mask_buffer.getvalue()).decode('utf-8')
raw_ref_image = RawReferenceImage(
reference_image={'image_bytes': img_b64},
reference_id=0
)
mask_ref_image = MaskReferenceImage(
reference_id=1,
reference_image={'image_bytes': mask_b64},
config=types.MaskReferenceConfig(
mask_mode="MASK_MODE_USER_PROVIDED",
mask_dilation=mask_dilation,
),
)
config_params = {
mask_np = mask.cpu().numpy()
if mask_np.ndim == 3: mask_np = mask_np[0]
mask_pil = Image.fromarray((mask_np * 255).astype(np.uint8), mode='L')
config_dict = {
"edit_mode": edit_mode,
"number_of_images": number_of_images,
"include_rai_reason": True,
"output_mime_type": "image/jpeg",
"base_steps": base_steps,
"seed": seed,
"guidance_scale": guidance_scale,
"output_mime_type": "image/jpeg",
"include_rai_reason": True,
}
if negative_prompt.strip():
config_params["negative_prompt"] = negative_prompt.strip()
config_dict["negative_prompt"] = negative_prompt.strip()
response = client.models.edit_image(
model="imagen-3.0-capability-001",
prompt=prompt,
reference_images=[raw_ref_image, mask_ref_image],
config=types.EditImageConfig(**config_params),
reference_images=[
types.RawReferenceImage(reference_id=0, reference_image={'image_bytes': to_b64(img_pil)}),
types.MaskReferenceImage(reference_id=1, reference_image={'image_bytes': to_b64(mask_pil)},
config=types.MaskReferenceConfig(mask_mode="MASK_MODE_USER_PROVIDED", mask_dilation=mask_dilation))
],
config=types.EditImageConfig(**config_dict)
)
if not response.generated_images: raise ValueError("No images generated")
output_tensors = []
for item in response.generated_images:
img_bytes = item.image.image_bytes
res_img = Image.open(io.BytesIO(img_bytes)).convert("RGB")
output_tensors.append(torch.from_numpy(np.array(res_img).astype(np.float32) / 255.0))
if not response.generated_images:
raise ValueError("No images generated by the API")
pil_images = []
for generated_image in response.generated_images:
image_data = generated_image.image
if hasattr(image_data, 'mode') and hasattr(image_data, 'size'):
pil_images.append(image_data)
elif hasattr(image_data, '_pil_image'):
pil_images.append(image_data._pil_image)
elif hasattr(image_data, 'show'):
try:
buffer = BytesIO()
image_data.save(buffer, format='PNG')
buffer.seek(0)
pil_images.append(Image.open(buffer))
except:
pil_images.append(Image.new('RGB', (512, 512), color='gray'))
elif hasattr(image_data, 'read') or isinstance(image_data, bytes):
image_bytes = image_data.read() if hasattr(image_data, 'read') else image_data
pil_images.append(Image.open(BytesIO(image_bytes)))
else:
try:
pil_images.append(Image.open(image_data))
except:
pil_images.append(Image.new('RGB', (512, 512), color='gray'))
return (self.pil_to_tensor(pil_images),)
return (torch.stack(output_tensors),)
except Exception as e:
print(f"Google Imagen Edit Error: {str(e)}")
error_image = Image.new('RGB', (512, 512), color='black')
return (self.pil_to_tensor([error_image]),)
print(f"Google Imagen Edit Error: {e}")
raise RuntimeError(f"Google Imagen Edit Error: {e}")
finally:
if os.path.exists(creds_file.name): os.remove(creds_file.name)
@classmethod
def IS_CHANGED(cls, **kwargs):
return f"{kwargs.get('prompt', '')}-{kwargs.get('negative_prompt', '')}-{kwargs.get('edit_mode', '')}-{kwargs.get('number_of_images', 1)}-{kwargs.get('seed', 12345)}-{kwargs.get('base_steps', 50)}"
return float("nan")
NODE_CLASS_MAPPINGS = {"GoogleImagenEditNode": GoogleImagenEditNode}
NODE_DISPLAY_NAME_MAPPINGS = {"GoogleImagenEditNode": "Google Imagen Edit (Vertex AI only)"}
+59 -78
View File
@@ -12,22 +12,22 @@ class NanoBananaNode:
def INPUT_TYPES(cls):
return {
"required": {
"api_key": ("STRING", {"multiline": False, "default": ""}),
"model": (["gemini-3-pro-image-preview", "gemini-2.5-flash-image"],),
"aspect_ratio": (["1:1", "2:3", "3:2", "3:4", "4:3", "9:16", "16:9", "21:9"],),
"temperature": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}),
"top_p": ("FLOAT", {"default": 0.85, "min": 0.0, "max": 1.0, "step": 0.01}),
"seed": ("INT", {"default": 69, "min": -1, "max": 2147483646, "step": 1}),
"resolution": (["1K", "2K", "4K"], {"default": "1K"}),
"prompt": ("STRING", {"multiline": True, "default": ""}),
"api_key": ("STRING", {"multiline": False, "default": ""}),
"model": (["gemini-3-pro-image-preview", "gemini-2.5-flash-image"],),
"aspect_ratio": (["1:1", "2:3", "3:2", "3:4", "4:3", "9:16", "16:9", "21:9"],),
"resolution": (["1K", "2K", "4K"], {"default": "1K"}),
"temperature": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}),
"top_p": ("FLOAT", {"default": 0.85, "min": 0.0, "max": 1.0, "step": 0.01}),
"seed": ("INT", {"default": 69, "min": -1, "max": 2147483646, "step": 1}),
},
"optional": {
"prompt": ("STRING", {"multiline": True, "default": ""}),
"system_instruction": ("STRING", {"multiline": True, "default": ""}),
"image_1": ("IMAGE",),
"image_2": ("IMAGE",),
"image_3": ("IMAGE",),
"image_4": ("IMAGE",),
"image_5": ("IMAGE",),
"system_instruction": ("STRING", {"multiline": True, "default": ""}),
"image_1": ("IMAGE",),
"image_2": ("IMAGE",),
"image_3": ("IMAGE",),
"image_4": ("IMAGE",),
"image_5": ("IMAGE",),
}
}
@@ -36,22 +36,17 @@ class NanoBananaNode:
FUNCTION = "generate"
CATEGORY = "image/generation"
def tensor_to_pil(self, tensor):
def _convert_tensor_to_bytes(self, tensor):
if tensor.dim() == 4:
tensor = tensor[0]
array = (tensor.cpu().numpy() * 255).astype(np.uint8)
return Image.fromarray(array)
def pil_to_tensor(self, image):
if image.mode != 'RGB':
image = image.convert('RGB')
array = np.array(image).astype(np.float32) / 255.0
tensor = torch.from_numpy(array)
return tensor.unsqueeze(0)
arr = (tensor.cpu().numpy() * 255).astype(np.uint8)
buf = io.BytesIO()
Image.fromarray(arr).save(buf, format='PNG')
return buf.getvalue()
def generate(self, api_key, model, aspect_ratio, resolution, temperature, top_p, seed,
prompt="", system_instruction="",
image_1=None, image_2=None, image_3=None, image_4=None, image_5=None):
prompt="", system_instruction="", **kwargs):
key = api_key.strip() or os.environ.get("GEMINI_API_KEY")
if not key:
@@ -59,69 +54,55 @@ class NanoBananaNode:
client = genai.Client(api_key=key)
# Build parts list
parts = []
input_images = [kwargs.get(f"image_{i}") for i in range(1, 6)]
for img in input_images:
if img is not None:
img_bytes = self._convert_tensor_to_bytes(img)
parts.append(types.Part.from_bytes(mime_type="image/png", data=img_bytes))
# Add images
for img_tensor in [image_1, image_2, image_3, image_4, image_5]:
if img_tensor is not None:
pil_img = self.tensor_to_pil(img_tensor)
buffer = io.BytesIO()
pil_img.save(buffer, format='PNG')
parts.append(types.Part.from_bytes(
mime_type="image/png",
data=buffer.getvalue()
))
# Add prompt if provided
if prompt.strip():
parts.append(types.Part.from_text(text=prompt))
if not parts:
raise ValueError("At least one image or prompt must be provided.")
contents = [types.Content(role="user", parts=parts)]
config_dict = {
"temperature": temperature,
"seed": seed,
"top_p": top_p,
"response_modalities": ["IMAGE"],
"image_config": types.ImageConfig(aspect_ratio=aspect_ratio),
}
img_config_params = {"aspect_ratio": aspect_ratio}
if "gemini-3-pro" in model:
config_dict["image_config"] = types.ImageConfig(aspect_ratio=aspect_ratio, image_size=resolution)
else:
print("gemini-2.5-flash-image does not support resolution parameter, using default resolution")
config = types.GenerateContentConfig(**config_dict)
if system_instruction.strip():
config.system_instruction = [types.Part.from_text(text=system_instruction)]
# Generate
response = client.models.generate_content(
model=model,
contents=contents,
config=config,
img_config_params["image_size"] = resolution
config = types.GenerateContentConfig(
temperature=temperature,
seed=seed,
top_p=top_p,
response_modalities=["IMAGE"],
image_config=types.ImageConfig(**img_config_params),
system_instruction=system_instruction.strip() if system_instruction.strip() else None
)
result_image = None
if response.candidates and response.candidates[0].content and response.candidates[0].content.parts:
for part in response.candidates[0].content.parts:
if part.inline_data and part.inline_data.data:
result_image = Image.open(io.BytesIO(part.inline_data.data))
break
try:
response = client.models.generate_content(
model=model,
contents=[types.Content(role="user", parts=parts)],
config=config,
)
except Exception as e:
raise RuntimeError(f"Gemini API Error: {str(e)}")
if result_image is None:
raise ValueError("No image generated by the API.")
return (self.pil_to_tensor(result_image),)
try:
img_data = response.candidates[0].content.parts[0].inline_data.data
result_pil = Image.open(io.BytesIO(img_data)).convert("RGB")
result_tensor = torch.from_numpy(np.array(result_pil).astype(np.float32) / 255.0).unsqueeze(0)
return (result_tensor,)
except (AttributeError, IndexError, TypeError):
raise ValueError("API returned a response, but no valid image data was found.")
@classmethod
def IS_CHANGED(cls, **kwargs):
return f"{kwargs.get('prompt', '')}-{kwargs.get('temperature', 0.5)}-{kwargs.get('top_p', 0.85)}-{kwargs.get('seed', 69)}-{kwargs.get('aspect_ratio', '1:1')}-{kwargs.get('model', 'gemini-3-pro-image-preview')}-{kwargs.get('image_1')}-{kwargs.get('image_2')}-{kwargs.get('image_3')}-{kwargs.get('image_4')}-{kwargs.get('image_5')}"
def IS_CHANGED(cls, seed, **kwargs):
return seed
NODE_CLASS_MAPPINGS = {"NanoBananaNode": NanoBananaNode}
NODE_DISPLAY_NAME_MAPPINGS = {"NanoBananaNode": "Nano Banana"}