Streamlined and cleaned up code
This commit is contained in:
+41
-78
@@ -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
@@ -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
@@ -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"}
|
||||
@@ -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
@@ -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
@@ -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"}
|
||||
Reference in New Issue
Block a user