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"}
|
||||
Reference in New Issue
Block a user