From 0d907bf82a4dac1e911fe494b9ea11d66a7f6a2f Mon Sep 17 00:00:00 2001 From: Aryan185 Date: Fri, 2 Jan 2026 13:54:01 +0000 Subject: [PATCH] Streamlined and cleaned up code --- gemini_node.py | 119 ++++++++++----------------- gemini_segment.py | 203 +++++++++++++++------------------------------- gemini_tts.py | 75 ++++++++--------- imagen.py | 115 ++++++++------------------ imagen_edit.py | 167 +++++++++++--------------------------- nano_banana.py | 137 ++++++++++++++----------------- 6 files changed, 282 insertions(+), 534 deletions(-) diff --git a/gemini_node.py b/gemini_node.py index bf5d70f..0e7fc86 100644 --- a/gemini_node.py +++ b/gemini_node.py @@ -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"} \ No newline at end of file diff --git a/gemini_segment.py b/gemini_segment.py index 59f3414..f51829e 100644 --- a/gemini_segment.py +++ b/gemini_segment.py @@ -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"} \ No newline at end of file diff --git a/gemini_tts.py b/gemini_tts.py index 6a819ee..9582b1b 100644 --- a/gemini_tts.py +++ b/gemini_tts.py @@ -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"} \ No newline at end of file diff --git a/imagen.py b/imagen.py index 1a7141a..315a9db 100644 --- a/imagen.py +++ b/imagen.py @@ -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"} \ No newline at end of file diff --git a/imagen_edit.py b/imagen_edit.py index cedc128..2d9c8ce 100644 --- a/imagen_edit.py +++ b/imagen_edit.py @@ -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)"} \ No newline at end of file diff --git a/nano_banana.py b/nano_banana.py index 7eb7f68..4a315eb 100644 --- a/nano_banana.py +++ b/nano_banana.py @@ -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"} \ No newline at end of file