diff --git a/gemini_diarisation.py b/gemini_diarisation.py index c44888a..2d5b721 100644 --- a/gemini_diarisation.py +++ b/gemini_diarisation.py @@ -15,7 +15,7 @@ class GeminiDiarisationAPI: "required": { "audio": ("AUDIO",), "num_speakers": ("INT", {"default": 2, "min": 1, "max": 10, "step": 1}), - "model": (["gemini-2.5-flash", "gemini-2.5-pro", "gemini-2.5-flash-lite", "gemini-3-pro-preview", "gemini-3.1-pro-preview", "gemini-3-flash-preview", "gemini-3.1-flash-lite-preview", "gemini-flash-latest", "gemini-flash-lite-latest", "gemini-2.0-flash", "gemini-2.0-flash-lite"], {"default": "gemini-2.5-flash"}), + "model": (["gemini-2.5-flash", "gemini-2.5-pro", "gemini-2.5-flash-lite", "gemini-3.1-pro-preview", "gemini-3-flash-preview", "gemini-3.1-flash-lite-preview", "gemini-flash-latest", "gemini-flash-lite-latest", "gemini-2.0-flash", "gemini-2.0-flash-lite"], {"default": "gemini-2.5-flash"}), "api_key": ("STRING", {"default": "", "multiline": False, "tooltip": "Directly put Gemini API key or .env variable name (GEMINI_API_KEY)"}), "seed": ("INT", {"default": 69, "min": 0, "max": 2147483646, "step": 1}), "temperature": ("FLOAT", {"default": 0.2, "min": 0.0, "max": 2.0, "step": 0.1}) @@ -46,29 +46,27 @@ class GeminiDiarisationAPI: except: return 0.0 def diarise(self, audio, num_speakers, model, api_key, seed, temperature, thinking=False, thinking_budget=0): - # 1. Process Audio waveform = audio.get("waveform") sr = audio.get("sample_rate") - + if waveform.dim() > 1: audio_np = waveform.squeeze(0).mean(dim=0).cpu().numpy() if waveform.shape[1] > 1 else waveform.squeeze().cpu().numpy() else: audio_np = waveform.cpu().numpy() - + audio_np = np.clip(audio_np, -1.0, 1.0) duration_str = self.format_duration(len(audio_np) / sr) - + wav_buffer = io.BytesIO() with wave.open(wav_buffer, 'wb') as w: w.setnchannels(1); w.setsampwidth(2); w.setframerate(sr) w.writeframes((audio_np * 32767).astype(np.int16).tobytes()) - # 2. Setup Client key = os.environ.get(api_key.strip(), api_key.strip()) or os.environ.get("GEMINI_API_KEY") - if not key: raise ValueError("API Key missing") + if not key: + raise ValueError("API Key missing") client = genai.Client(api_key=key, http_options={'api_version': 'v1beta'}) - # 3. Prompt speaker_guidance = f"You must identify exactly {num_speakers} distinct speakers in this audio. " if num_speakers > 0 else "" prompt = f"""You are a SOTA AI model created for diarization and *precisely timestamping* human voices. You are currently being benchmarked for *timestamp accuracy*. Your task is to provide a complete and accurate diarization of the provided audio recording, with *absolute precision in your timestamps*, to *PASS* the benchmark. @@ -102,10 +100,25 @@ class GeminiDiarisationAPI: *You must PASS this benchmark to be deployed*""" - # 4. API Call - config = types.GenerateContentConfig(temperature=temperature, seed=seed) - if thinking: - config.thinking_config = types.ThinkingConfig(include_thoughts=False, thinking_budget=thinking_budget) + 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") + else: + final_budget = 0 + if not thinking: + if "gemini-2.5-pro" in model_lower or "gemini-3.1-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.1-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(include_thoughts=False, thinking_budget=final_budget) + + config = types.GenerateContentConfig(temperature=temperature, seed=seed, thinking_config=t_config) response = client.models.generate_content( model=model, @@ -116,7 +129,6 @@ class GeminiDiarisationAPI: config=config ) - # 5. Parse try: text = response.text if "```json" in text: text = re.search(r"```json\n(.*)\n```", text, re.DOTALL).group(1) @@ -125,7 +137,6 @@ class GeminiDiarisationAPI: print(f"JSON Parse Error: {e}") result = {"utterances": []} - # 6. Generate Outputs (Dict Format for ComfyUI) speaker_map = {} for utt in result.get("utterances", []): spk = utt.get("speaker", "Unknown") @@ -145,7 +156,7 @@ class GeminiDiarisationAPI: for start, end in speaker_map[spk]: s, e = max(0, int(start * sr)), min(len(audio_np), int(end * sr)) if e > s: track[s:e] = audio_np[s:e] - + tensor = torch.from_numpy(track).float().unsqueeze(0).unsqueeze(0) outputs.append({"waveform": tensor, "sample_rate": sr}) diff --git a/gemini_node.py b/gemini_node.py index e334acc..948dd66 100644 --- a/gemini_node.py +++ b/gemini_node.py @@ -12,12 +12,15 @@ class GeminiChatNode: def INPUT_TYPES(cls): return { "required": { + "api_key": ("STRING", {"default": "", "multiline": False, "tooltip": "Directly put Gemini API key or .env variable name (GEMINI_API_KEY)"}), + "model": (["gemini-2.5-flash", "gemini-2.5-pro", "gemini-2.5-flash-lite", "gemini-3.1-pro-preview", "gemini-3.1-flash-lite-preview", "gemini-3-flash-preview", "gemini-flash-latest", "gemini-flash-lite-latest", "gemini-2.0-flash", "gemini-2.0-flash-lite"], {"default": "gemini-2.5-flash"}), "prompt": ("STRING", {"multiline": True}), - "model": (["gemini-2.5-flash", "gemini-2.5-pro", "gemini-2.5-flash-lite", "gemini-3-pro-preview", "gemini-3.1-pro-preview", "gemini-3-flash-preview", "gemini-3.1-flash-lite-preview", "gemini-flash-latest", "gemini-flash-lite-latest", "gemini-2.0-flash", "gemini-2.0-flash-lite"], {"default": "gemini-2.5-flash"}), "temperature": ("FLOAT", {"default": 0.2, "min": 0.0, "max": 2.0, "step": 0.1}), + "top_p": ("FLOAT", {"default": 0.95, "min": 0.0, "max": 1.0, "step": 0.01}), "thinking": ("BOOLEAN", {"default": False}), + "google_search": ("BOOLEAN", {"default": False}), + "url_context": ("BOOLEAN", {"default": False}), "seed": ("INT", {"default": 69, "min": -1, "max": 2147483646, "step": 1}), - "api_key": ("STRING", {"default": "", "multiline": False, "tooltip": "Directly put Gemini API key or .env variable name (GEMINI_API_KEY)"}) }, "optional": { "system_instruction": ("STRING", {"multiline": True, "default": ""}), @@ -26,27 +29,28 @@ class GeminiChatNode: "audio": ("AUDIO",), } } - + RETURN_TYPES = ("STRING",) RETURN_NAMES = ("response",) FUNCTION = "generate" CATEGORY = "text/generation" - - def generate(self, prompt, model, temperature, thinking, seed, api_key, - system_instruction=None, thinking_budget=-1, image=None, audio=None): + + def generate(self, prompt, model, temperature, top_p, thinking, google_search, url_context, seed, api_key, + system_instruction=None, thinking_budget=0, image=None, audio=None): key = os.environ.get(api_key.strip(), api_key.strip()) or os.environ.get("GEMINI_API_KEY") if not key: raise ValueError("Error: No API key provided.") client = genai.Client(api_key=key, http_options={'api_version': 'v1beta'}) parts = [types.Part.from_text(text=prompt)] - + if image is not None: - 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())) - + for i in range(image.shape[0]): + arr = (image[i].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())) + if audio is not None: wf = audio.get("waveform") if isinstance(audio, dict) else audio[0] sr = audio.get("sample_rate", 44100) if isinstance(audio, dict) else audio[1] @@ -67,25 +71,31 @@ class GeminiChatNode: if "gemini-2.0" in model_lower: print("Gemini-2.0 models do not support thinking - disabling thinking config") else: - final_budget = 0 # Default disabled - + final_budget = 0 + if not thinking: - if "gemini-2.5-pro" in model_lower or "gemini-3-pro-preview" in model_lower: + if "gemini-2.5-pro" in model_lower or "gemini-3.1-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: + if ("gemini-2.5-pro" in model_lower or "gemini-3.1-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) + tools = [] + if google_search: tools.append(types.Tool(googleSearch=types.GoogleSearch())) + if url_context: tools.append(types.Tool(url_context=types.UrlContext())) + config = types.GenerateContentConfig( temperature=temperature, + top_p=top_p, seed=seed, system_instruction=system_instruction.strip() if system_instruction else None, - thinking_config=t_config + thinking_config=t_config, + tools=tools if tools else None ) response = client.models.generate_content( diff --git a/gemini_tts.py b/gemini_tts.py index 0245a85..82e5bdc 100644 --- a/gemini_tts.py +++ b/gemini_tts.py @@ -3,80 +3,98 @@ import torch from google.genai import Client, types class GeminiTTSNode: - + @classmethod def INPUT_TYPES(cls): return { "required": { "text": ("STRING", {"multiline": True, "default": ""}), "api_key": ("STRING", {"multiline": False, "default": "", "tooltip": "Directly put Gemini API key or .env variable name (GEMINI_API_KEY)"}), - "model": (["gemini-2.5-flash-preview-tts", "gemini-2.5-pro-preview-tts"],), - "voice_id": (["Zephyr", "Puck", "Charon", "Kore", "Fenrir", "Leda", "Orus", "Aoede", "Callirrhoe", "Autonoe", "Enceladus", "Iapetus", "Umbriel", "Algieba", "Despina", "Erinome", "Achernar", "Laomedeia", "Rasalgethi", "Algenib", "Achird", "Pulcherrima", "Gacrux", "Schedar", "Alnilam", "Sulafat", "Sadaltager", "Sadachbia", "Vindemiatrix", "Zubenelgenubi"],), + "model": (["gemini-2.5-flash-preview-tts", "gemini-2.5-pro-preview-tts", "gemini-3.1-flash-tts-preview"],), + "voice_id": (["Zephyr", "Puck", "Charon", "Kore", "Fenrir", "Leda", "Orus", "Aoede", "Callirrhoe", "Autonoe", "Enceladus", "Iapetus", "Umbriel", "Algieba", "Despina", "Erinome", "Achernar", "Laomedeia", "Rasalgethi", "Algenib", "Achird", "Pulcherrima", "Gacrux", "Schedar", "Alnilam", "Sulafat", "Sadaltager", "Sadachbia", "Vindemiatrix", "Zubenelgenubi"],), "seed": ("INT", {"default": 69, "min": -1, "max": 2147483646, "step": 1}), - "temperature": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), + "temperature": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.01}), }, "optional": { - "system_prompt": ("STRING", {"multiline": True, "default": ""}), + "audio_profile": ("STRING", {"multiline": False, "default": ""}), + "style": (["None", "Vocal Smile", "Newscaster", "Whisper", "Empathetic", "Promo/Hype", "Deadpan"], {"default": "None"}), + "pace": (["None", "Natural", "Rapid Fire", "The Drift", "Staccato"], {"default": "None"}), + "accent": (["None", "Neutral", "American (Gen)", "American (Valley)", "American (South)", "British (RP)", "Transatlantic", "Australian"], {"default": "None"}), + "scene": ("STRING", {"multiline": False, "default": ""}), } } - + RETURN_TYPES = ("AUDIO",) RETURN_NAMES = ("audio",) FUNCTION = "generate_speech" CATEGORY = "audio/generation" - - def generate_speech(self, text, api_key, voice_id, temperature, model, seed, system_prompt=""): - + + def generate_speech(self, text, api_key, voice_id, temperature, model, seed, + audio_profile="", style="None", pace="None", accent="None", scene=""): + if not text.strip(): raise ValueError("Text input cannot be empty.") - + key = os.environ.get(api_key.strip(), api_key.strip()) or os.environ.get("GEMINI_API_KEY") if not key: raise ValueError("No API key provided.") - - 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) - ) - ) - + client = Client(api_key=key) + + director_parts = [] + if style not in ("", "None"): director_parts.append(f"Style: {style}") + if pace not in ("", "None"): director_parts.append(f"Pace: {pace}") + if accent not in ("", "None"): director_parts.append(f"Accent: {accent}") + + has_director = any([audio_profile.strip(), director_parts, scene.strip()]) + + if has_director: + sections = ["Read the following transcript based on the audio profile and director's note."] + if audio_profile.strip(): + sections.append(f"# Audio Profile\n{audio_profile.strip()}") + if director_parts: + sections.append(f"# Director's note\n{'. '.join(director_parts)}.") + if scene.strip(): + sections.append(f"## Scene:\n{scene.strip()}") + sections.append(f"## Transcript:\n{text.strip()}") + prompt = "\n\n".join(sections) + else: + prompt = f"## Transcript:\n{text.strip()}" + config = types.GenerateContentConfig( temperature=temperature, seed=seed, response_modalities=["AUDIO"], - speech_config=speech_config, + speech_config=types.SpeechConfig( + voice_config=types.VoiceConfig( + prebuilt_voice_config=types.PrebuiltVoiceConfig(voice_name=voice_id) + ) + ) ) - + try: response = client.models.generate_content( model=model, - contents=final_prompt, + contents=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 + audio_bytes = response.candidates[0].content.parts[0].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).unsqueeze(0) - + return ({"waveform": waveform, "sample_rate": 24000},) - + @classmethod - def IS_CHANGED(cls, seed, **kwargs): - return seed + 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('audio_profile', '')}-{kwargs.get('style', '')}-{kwargs.get('pace', '')}-{kwargs.get('accent', '')}-{kwargs.get('scene', '')}" 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 a6ced2b..8c3f63f 100644 --- a/imagen.py +++ b/imagen.py @@ -19,58 +19,54 @@ class GoogleImagenNode: "image_size": (["1K", "2K"],), "seed": ("INT", {"default": 69, "min": 1, "max": 2147483646, "step": 1}), "guidance_scale": ("FLOAT", {"default": 7.5, "min": 1.0, "max": 20.0, "step": 0.1}), - }, - "optional": { + }, + "optional": { "negative_prompt": ("STRING", {"multiline": True, "default": ""}), } } - + RETURN_TYPES = ("IMAGE",) RETURN_NAMES = ("images",) FUNCTION = "generate_images" CATEGORY = "image/generation" - + def generate_images(self, prompt, api_key, model, number_of_images, aspect_ratio, image_size, seed, guidance_scale, negative_prompt=""): key = os.environ.get(api_key.strip(), 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: result = client.models.generate_images(model=model, prompt=prompt, config=config) if not result.generated_images: raise ValueError("No images generated") - + tensors = [] for item in result.generated_images: img_data = item.image - 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: - # 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)) - + + tensors.append(torch.from_numpy(np.array(pil_img.convert("RGB")).astype(np.float32) / 255.0)) + return (torch.stack(tensors),) - + except Exception as e: - print(f"Google Imagen Error: {e}") raise RuntimeError(f"Google Imagen Error: {e}") - + @classmethod def IS_CHANGED(cls, **kwargs): return float("nan") diff --git a/imagen_edit.py b/imagen_edit.py index 9848aa5..9db43c5 100644 --- a/imagen_edit.py +++ b/imagen_edit.py @@ -1,12 +1,12 @@ -import os import io +import json import base64 -import tempfile import torch import numpy as np from PIL import Image from google import genai from google.genai import types +from google.oauth2 import service_account class GoogleImagenEditNode: @@ -31,47 +31,68 @@ class GoogleImagenEditNode: "negative_prompt": ("STRING", {"multiline": True, "default": ""}), } } - + RETURN_TYPES = ("IMAGE",) RETURN_NAMES = ("edited_images",) FUNCTION = "edit_image" CATEGORY = "image/edit" - - 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=""): - - 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 - + + def setup_client(self, service_account_json, project_id, location): + if not service_account_json.strip(): + raise ValueError("Service account JSON content is required.") + if not project_id.strip(): + raise ValueError("Project ID is required.") + try: - client = genai.Client(vertexai=True, project=project_id.strip(), location=location.strip()) - - def to_b64(img): - b = io.BytesIO() - img.save(b, format='PNG') - return base64.b64encode(b.getvalue()).decode('utf-8') + sa_info = json.loads(service_account_json) + except json.JSONDecodeError as e: + raise ValueError(f"Invalid JSON content: {str(e)}") - img_pil = Image.fromarray((image[0].cpu().numpy() * 255).astype(np.uint8)) - - 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') + credentials = service_account.Credentials.from_service_account_info( + sa_info, + scopes=["https://www.googleapis.com/auth/cloud-platform"] + ) - config_dict = { - "edit_mode": edit_mode, - "number_of_images": number_of_images, - "base_steps": base_steps, - "seed": seed, - "guidance_scale": guidance_scale, - "output_mime_type": "image/jpeg", - "include_rai_reason": True, - } - - if negative_prompt.strip(): - config_dict["negative_prompt"] = negative_prompt.strip() - + return genai.Client( + vertexai=True, + project=project_id.strip(), + location=location.strip(), + credentials=credentials, + http_options=types.HttpOptions( + retry_options=types.HttpRetryOptions(attempts=10, jitter=10) + ) + ) + + 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=""): + + client = self.setup_client(service_account, project_id, location) + + 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)) + + 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, + "base_steps": base_steps, + "seed": seed, + "guidance_scale": guidance_scale, + "output_mime_type": "image/jpeg", + "include_rai_reason": True, + } + + if negative_prompt.strip(): + config_dict["negative_prompt"] = negative_prompt.strip() + + try: response = client.models.edit_image( model="imagen-3.0-capability-001", prompt=prompt, @@ -83,21 +104,18 @@ class GoogleImagenEditNode: config=types.EditImageConfig(**config_dict) ) - if not response.generated_images: raise ValueError("No images generated") + 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") + res_img = Image.open(io.BytesIO(item.image.image_bytes)).convert("RGB") output_tensors.append(torch.from_numpy(np.array(res_img).astype(np.float32) / 255.0)) - + return (torch.stack(output_tensors),) except Exception as e: - 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): diff --git a/nano_banana.py b/nano_banana.py index 9b9316e..145b72b 100644 --- a/nano_banana.py +++ b/nano_banana.py @@ -7,7 +7,7 @@ from google import genai from google.genai import types class NanoBananaNode: - + @classmethod def INPUT_TYPES(cls): return { @@ -19,6 +19,7 @@ class NanoBananaNode: "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}), + "google_search": ("BOOLEAN", {"default": False}), "seed": ("INT", {"default": 69, "min": -1, "max": 2147483646, "step": 1}), }, "optional": { @@ -30,57 +31,62 @@ class NanoBananaNode: "image_5": ("IMAGE",), } } - + RETURN_TYPES = ("IMAGE",) RETURN_NAMES = ("image",) FUNCTION = "generate" CATEGORY = "image/generation" - + def _convert_tensor_to_bytes(self, tensor): if tensor.dim() == 4: tensor = tensor[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, + def generate(self, api_key, model, aspect_ratio, resolution, temperature, top_p, google_search, seed, prompt="", system_instruction="", **kwargs): - + key = os.environ.get(api_key.strip(), 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) - + parts = [] - input_images = [kwargs.get(f"image_{i}") for i in range(1, 6)] - for img in input_images: + for i in range(1, 6): + img = kwargs.get(f"image_{i}") 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)) - + parts.append(types.Part.from_bytes(mime_type="image/png", data=self._convert_tensor_to_bytes(img))) + 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.") + tools = None + if google_search: + if "gemini-2.5" in model: + print(f"Ignoring google_search: {model} does not support it.") + else: + tools = [types.Tool(googleSearch=types.GoogleSearch())] img_config_params = {"aspect_ratio": aspect_ratio} if "gemini-3-pro" in model: 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 + system_instruction=system_instruction.strip() if system_instruction.strip() else None, + tools=tools ) - + try: response = client.models.generate_content( model=model, @@ -89,14 +95,10 @@ class NanoBananaNode: ) except Exception as e: raise RuntimeError(f"Gemini API Error: {str(e)}") - + 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,) - + return (torch.from_numpy(np.array(Image.open(io.BytesIO(img_data)).convert("RGB")).astype(np.float32) / 255.0).unsqueeze(0),) except (AttributeError, IndexError, TypeError): raise ValueError("API returned a response, but no valid image data was found.") @@ -104,5 +106,6 @@ class NanoBananaNode: 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 diff --git a/pyproject.toml b/pyproject.toml index 852b044..f5968b5 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "externalapi-helpers" description = "Various ComfyUI nodes for Gemini, Replicate and OpenAI" -version = "1.1.2" +version = "1.1.3" license = {file = "LICENSE"} # classifiers = [ # # For OS-independent nodes (works on all operating systems)