Add files via upload
This commit is contained in:
+232
-2
@@ -7,6 +7,8 @@ import asyncio
|
||||
import requests
|
||||
from PIL import Image
|
||||
from io import BytesIO
|
||||
import tempfile
|
||||
import time
|
||||
from typing import List, Dict, Any, Optional, Union, Tuple
|
||||
from pathlib import Path
|
||||
from .send_request import send_request
|
||||
@@ -21,7 +23,12 @@ from .utils import (
|
||||
load_combo_settings,
|
||||
create_settings_from_ui,
|
||||
prepare_batch_images,
|
||||
process_auto_mode_images
|
||||
process_auto_mode_images,
|
||||
tensor_to_pil,
|
||||
gemini2_process_images,
|
||||
gemini2_prepare_response,
|
||||
gemini2_create_client,
|
||||
validate_gemini_key
|
||||
)
|
||||
import base64
|
||||
import numpy as np
|
||||
@@ -29,6 +36,15 @@ import codecs
|
||||
import random
|
||||
import math
|
||||
|
||||
# Add Google Gemini SDK imports
|
||||
try:
|
||||
from google import genai
|
||||
from google.genai import types
|
||||
GEMINI_SDK_AVAILABLE = True
|
||||
except ImportError:
|
||||
GEMINI_SDK_AVAILABLE = False
|
||||
print("Google Generative AI SDK not found. Install with: pip install google-generativeai")
|
||||
|
||||
# Add ComfyUI directory to path
|
||||
comfy_path = os.path.abspath(os.path.join(os.path.dirname(__file__), '..', '..'))
|
||||
if comfy_path not in sys.path:
|
||||
@@ -171,7 +187,7 @@ class IFLLM:
|
||||
},
|
||||
"optional": {
|
||||
"images": ("IMAGE", {"list": True}),
|
||||
"strategy": (["normal", "omost", "create", "edit", "variations"], {"default": "normal"}),
|
||||
"strategy": (["normal", "omost", "create", "edit", "variations", "gemini2_create"], {"default": "normal"}),
|
||||
"mask": ("MASK", {}),
|
||||
"prime_directives": ("STRING", {"forceInput": True, "tooltip": "The system prompt for the LLM."}),
|
||||
"profiles": (["None"] + list(cls().profiles.keys()), {"default": "None", "tooltip": "The pre-defined system_prompt from the json profile file on the presets folder you can edit or make your own will be listed here."}),
|
||||
@@ -410,6 +426,9 @@ class IFLLM:
|
||||
elif strategy_name == "edit":
|
||||
return await self.execute_edit_strategy(
|
||||
user_prompt, current_images, current_mask, **kwargs)
|
||||
elif strategy_name == "gemini2_create":
|
||||
return await self.execute_gemini2_create_strategy(
|
||||
user_prompt, current_images, current_mask, **kwargs)
|
||||
else:
|
||||
raise ValueError(f"Unsupported strategy: {strategy_name}")
|
||||
|
||||
@@ -1083,6 +1102,217 @@ class IFLLM:
|
||||
user_prompt
|
||||
)
|
||||
|
||||
async def execute_gemini2_create_strategy(self, user_prompt, current_images, current_mask=None, **kwargs):
|
||||
"""
|
||||
Execute Gemini 2.0 create strategy using the Google Gemini API SDK.
|
||||
Handles batches of images as input and returns generated images.
|
||||
|
||||
Args:
|
||||
user_prompt (str): The prompt for image generation
|
||||
current_images (torch.Tensor): Batch of input images [B,H,W,C]
|
||||
current_mask (torch.Tensor, optional): Mask tensor
|
||||
**kwargs: Additional arguments including API key, model settings, etc.
|
||||
|
||||
Returns:
|
||||
dict: Response dictionary with generated images and other metadata
|
||||
"""
|
||||
try:
|
||||
# Check if Gemini SDK is available
|
||||
if not GEMINI_SDK_AVAILABLE:
|
||||
error_msg = "Google Generative AI SDK not installed. Install with: pip install google-generativeai"
|
||||
logger.error(error_msg)
|
||||
return self.create_error_response(
|
||||
current_images,
|
||||
current_mask,
|
||||
error_msg,
|
||||
user_prompt
|
||||
)
|
||||
|
||||
# Initialize variables for response
|
||||
response_text = ""
|
||||
temp_img_paths = []
|
||||
|
||||
# Get API key
|
||||
if kwargs.get('external_api_key'):
|
||||
api_key = kwargs.get('external_api_key')
|
||||
else:
|
||||
api_key = kwargs.get('llm_api_key')
|
||||
|
||||
if not api_key:
|
||||
logger.error("No valid Gemini API key provided")
|
||||
return self.create_error_response(
|
||||
current_images,
|
||||
current_mask,
|
||||
"Error: No valid Gemini API key provided. Please set GEMINI_API_KEY in your environment or provide external_api_key.",
|
||||
user_prompt
|
||||
)
|
||||
|
||||
# Process parameters
|
||||
temperature = kwargs.get('temperature', 0.8)
|
||||
seed = kwargs.get('seed', 0)
|
||||
batch_count = kwargs.get('batch_count', 1)
|
||||
|
||||
# Use random seed if seed is 0 or random is True
|
||||
if seed == 0 or kwargs.get('random', False):
|
||||
import random
|
||||
seed = random.randint(1, 2**31 - 1)
|
||||
|
||||
logger.info(f"Using Gemini 2.0 Create strategy with seed: {seed}, temperature: {temperature}")
|
||||
|
||||
# Create Gemini client
|
||||
client = genai.Client(api_key=api_key)
|
||||
|
||||
# Process input images
|
||||
if current_images is not None and current_images.nelement() > 0:
|
||||
# Prepare input images for Gemini API
|
||||
input_images = prepare_batch_images(current_images)
|
||||
logger.info(f"Processing {len(input_images)} input images for Gemini")
|
||||
|
||||
# Convert images to format required by Gemini
|
||||
contents = []
|
||||
|
||||
# Add each image to the request
|
||||
for idx, img in enumerate(input_images):
|
||||
try:
|
||||
# Convert tensor to PIL image
|
||||
pil_image = tensor_to_pil(img)
|
||||
|
||||
# Save as temporary file
|
||||
temp_img_path = os.path.join(tempfile.gettempdir(), f"gemini_input_{idx}_{int(time.time())}.png")
|
||||
pil_image.save(temp_img_path)
|
||||
temp_img_paths.append(temp_img_path)
|
||||
|
||||
# Read image data
|
||||
with open(temp_img_path, "rb") as f:
|
||||
image_bytes = f.read()
|
||||
|
||||
# Add image to content
|
||||
contents.append({
|
||||
"inline_data": {
|
||||
"mime_type": "image/png",
|
||||
"data": image_bytes
|
||||
}
|
||||
})
|
||||
|
||||
except Exception as img_error:
|
||||
logger.error(f"Error processing input image {idx}: {str(img_error)}")
|
||||
|
||||
# Add the prompt after all images
|
||||
contents.append({"text": user_prompt})
|
||||
else:
|
||||
# No input images, just use the prompt
|
||||
contents = user_prompt
|
||||
logger.info("No input images provided, using text prompt only")
|
||||
|
||||
# Configure generation parameters
|
||||
gen_config = types.GenerateContentConfig(
|
||||
temperature=temperature,
|
||||
seed=seed,
|
||||
response_modalities=['Text', 'Image']
|
||||
# Request multiple images based on batch_count
|
||||
# generation_parameters parameter is not supported and causing errors
|
||||
# generation_parameters={
|
||||
# "num_iterations": batch_count
|
||||
# }
|
||||
)
|
||||
|
||||
# Note: Gemini 2.0 API doesn't support the num_iterations parameter directly
|
||||
# It can return multiple images for some prompts but doesn't guarantee batch_count
|
||||
# The API will decide how many images to return based on the prompt
|
||||
|
||||
# Call Gemini API
|
||||
logger.info(f"Calling Gemini API with {len(contents) if isinstance(contents, list) else 1} content parts")
|
||||
response = client.models.generate_content(
|
||||
model="models/gemini-2.0-flash-exp", # Using the latest image generation model
|
||||
contents=contents,
|
||||
config=gen_config
|
||||
)
|
||||
|
||||
logger.info("Received response from Gemini API")
|
||||
|
||||
# Process the response to extract generated images
|
||||
if not hasattr(response, 'candidates') or not response.candidates:
|
||||
logger.error("API response contained no candidates")
|
||||
return self.create_error_response(
|
||||
current_images,
|
||||
current_mask,
|
||||
"Error: Gemini API returned no candidates in the response",
|
||||
user_prompt
|
||||
)
|
||||
|
||||
# Extract generated images and text
|
||||
generated_images = []
|
||||
|
||||
for candidate_idx, candidate in enumerate(response.candidates):
|
||||
if not hasattr(candidate, 'content') or not hasattr(candidate.content, 'parts'):
|
||||
continue
|
||||
|
||||
for part in candidate.content.parts:
|
||||
# Extract text content
|
||||
if hasattr(part, 'text') and part.text:
|
||||
response_text += part.text + "\n"
|
||||
|
||||
# Extract image content
|
||||
if hasattr(part, 'inline_data') and part.inline_data:
|
||||
try:
|
||||
# Get binary image data
|
||||
image_binary = part.inline_data.data
|
||||
generated_images.append(image_binary)
|
||||
logger.info(f"Extracted image {len(generated_images)} from response")
|
||||
except Exception as img_error:
|
||||
logger.error(f"Error extracting image from response: {str(img_error)}")
|
||||
|
||||
# Clean up temporary files
|
||||
for temp_path in temp_img_paths:
|
||||
try:
|
||||
if os.path.exists(temp_path):
|
||||
os.remove(temp_path)
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to remove temporary file {temp_path}: {str(e)}")
|
||||
|
||||
# If no images were generated, return error
|
||||
if not generated_images:
|
||||
logger.warning("No images found in Gemini API response")
|
||||
return self.create_error_response(
|
||||
current_images,
|
||||
current_mask,
|
||||
f"No images generated. API response: {response_text[:500]}...",
|
||||
user_prompt
|
||||
)
|
||||
|
||||
# Process generated images for ComfyUI
|
||||
image_data = {
|
||||
"data": [{"b64_json": base64.b64encode(img).decode('utf-8')} for img in generated_images]
|
||||
}
|
||||
|
||||
# Convert binary image data to tensors
|
||||
images_tensor, mask_tensor = process_images_for_comfy(
|
||||
image_data,
|
||||
placeholder_image_path=self.placeholder_image_path,
|
||||
response_key="data",
|
||||
field_name="b64_json"
|
||||
)
|
||||
|
||||
logger.info(f"Successfully processed {len(generated_images)} generated images")
|
||||
|
||||
return {
|
||||
"Question": user_prompt,
|
||||
"Response": f"Generated {len(generated_images)} images with Gemini 2.0.\n\n{response_text}",
|
||||
"Negative": kwargs.get('neg_content', ''),
|
||||
"Tool_Output": generated_images,
|
||||
"Retrieved_Image": images_tensor,
|
||||
"Mask": mask_tensor
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error in Gemini 2.0 create strategy: {str(e)}", exc_info=True)
|
||||
return self.create_error_response(
|
||||
current_images,
|
||||
current_mask,
|
||||
f"Error in Gemini 2.0 create strategy: {str(e)}",
|
||||
user_prompt
|
||||
)
|
||||
|
||||
def get_models(self, engine, base_ip, port, api_key=None):
|
||||
return get_models(engine, base_ip, port, api_key)
|
||||
|
||||
|
||||
+42
-42
@@ -12,52 +12,52 @@ import aiohttp
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
async def send_anthropic_request(api_key, model, system_message, user_message, messages, temperature, max_tokens, base64_images, tools=None, tool_choice=None):
|
||||
client = AsyncAnthropic(
|
||||
api_key=api_key,
|
||||
base_url="https://api.anthropic.com",
|
||||
default_headers={
|
||||
"anthropic-version": "2023-06-01",
|
||||
"anthropic-beta": "prompt-caching-2024-07-31"
|
||||
}
|
||||
)
|
||||
|
||||
anthropic_messages = prepare_anthropic_messages(user_message, messages, base64_images)
|
||||
|
||||
data = {
|
||||
"model": model,
|
||||
"messages": anthropic_messages,
|
||||
"temperature": temperature,
|
||||
"max_tokens": max_tokens
|
||||
}
|
||||
|
||||
if system_message:
|
||||
data["system"] = system_message
|
||||
|
||||
if tools:
|
||||
data["tools"] = tools
|
||||
if tool_choice:
|
||||
data["tool_choice"] = tool_choice
|
||||
|
||||
try:
|
||||
response = await client.messages.create(**data)
|
||||
# Create client with minimal parameters
|
||||
client = AsyncAnthropic(
|
||||
api_key=api_key
|
||||
)
|
||||
|
||||
anthropic_messages = prepare_anthropic_messages(user_message, messages, base64_images)
|
||||
|
||||
data = {
|
||||
"model": model,
|
||||
"messages": anthropic_messages,
|
||||
"temperature": temperature,
|
||||
"max_tokens": max_tokens
|
||||
}
|
||||
|
||||
if system_message:
|
||||
data["system"] = system_message
|
||||
|
||||
if tools:
|
||||
# If tools were used, return the full response
|
||||
return response
|
||||
else:
|
||||
# If no tools were used, format the response to match the specified structure
|
||||
generated_text = response.content[0].text if response.content else ""
|
||||
return {
|
||||
"choices": [{
|
||||
"message": {
|
||||
"content": generated_text
|
||||
}
|
||||
}]
|
||||
}
|
||||
data["tools"] = tools
|
||||
if tool_choice:
|
||||
data["tool_choice"] = tool_choice
|
||||
|
||||
try:
|
||||
response = await client.messages.create(**data)
|
||||
|
||||
if tools:
|
||||
# If tools were used, return the full response
|
||||
return response
|
||||
else:
|
||||
# If no tools were used, format the response to match the specified structure
|
||||
generated_text = response.content[0].text if response.content else ""
|
||||
return {
|
||||
"choices": [{
|
||||
"message": {
|
||||
"content": generated_text
|
||||
}
|
||||
}]
|
||||
}
|
||||
except Exception as e:
|
||||
error_msg = f"Error: An exception occurred while processing the Anthropic request: {str(e)}"
|
||||
logger.error(error_msg)
|
||||
return {"choices": [{"message": {"content": error_msg}}]}
|
||||
except Exception as e:
|
||||
error_msg = f"Error: An exception occurred while processing the Anthropic request: {str(e)}"
|
||||
logger.error(error_msg)
|
||||
return {"choices": [{"message": {"content": error_msg}}]}
|
||||
logger.error(f"Error initializing Anthropic client: {str(e)}")
|
||||
return {"choices": [{"message": {"content": f"Error initializing Anthropic client: {str(e)}"}}]}
|
||||
|
||||
def detect_image_type(base64_string):
|
||||
"""
|
||||
|
||||
+19
-16
@@ -46,25 +46,28 @@ async def send_groq_request(
|
||||
Union[str, Dict[str, Any]]: Standardized response.
|
||||
"""
|
||||
try:
|
||||
# Initialize client with minimal parameters
|
||||
client = AsyncGroq(api_key=api_key)
|
||||
|
||||
# Prepare messages
|
||||
groq_messages = prepare_groq_messages(base64_images, user_message, messages)
|
||||
|
||||
# Create completion using AsyncGroq client
|
||||
completion = await client.chat.completions.create(
|
||||
model=model,
|
||||
messages=groq_messages,
|
||||
temperature=temperature,
|
||||
max_tokens=max_tokens,
|
||||
top_p=top_p,
|
||||
stream=False, # Assuming streaming is not required
|
||||
stop=None, # Adjust stop sequences if necessary
|
||||
)
|
||||
|
||||
# Convert completion to a serializable format
|
||||
completion_dict = completion.to_dict() if hasattr(completion, 'to_dict') else {}
|
||||
logger.debug(f"Received response: {json.dumps(completion_dict, indent=2)}")
|
||||
try:
|
||||
# Create completion using AsyncGroq client
|
||||
completion = await client.chat.completions.create(
|
||||
model=model,
|
||||
messages=groq_messages,
|
||||
temperature=temperature,
|
||||
max_tokens=max_tokens,
|
||||
top_p=top_p,
|
||||
stream=False, # Assuming streaming is not required
|
||||
stop=None, # Adjust stop sequences if necessary
|
||||
)
|
||||
|
||||
# Convert completion to a serializable format
|
||||
completion_dict = completion.to_dict() if hasattr(completion, 'to_dict') else {}
|
||||
logger.debug(f"Received response: {json.dumps(completion_dict, indent=2)}")
|
||||
|
||||
if tools:
|
||||
return completion_dict if completion_dict else completion
|
||||
else:
|
||||
@@ -88,8 +91,8 @@ async def send_groq_request(
|
||||
logger.error(f"Groq API error: {e}")
|
||||
return {"choices": [{"message": {"content": str(e)}}]}
|
||||
except Exception as e:
|
||||
logger.error(f"Unexpected error: {e}")
|
||||
return {"choices": [{"message": {"content": "An unexpected error occurred."}}]}
|
||||
logger.error(f"Error initializing Groq client: {str(e)}")
|
||||
return {"choices": [{"message": {"content": f"Error initializing Groq client: {str(e)}"}}]}
|
||||
|
||||
def prepare_groq_messages(
|
||||
base64_images: List[str],
|
||||
|
||||
@@ -1765,3 +1765,166 @@ def print_available_models():
|
||||
# Usage example:
|
||||
if __name__ == "__main__":
|
||||
print_available_models()
|
||||
|
||||
def gemini2_process_images(images, max_input_images=5, target_size=(768, 768)):
|
||||
"""
|
||||
Process a batch of images for Gemini 2.0 API.
|
||||
|
||||
Args:
|
||||
images (torch.Tensor or list): Image batch in ComfyUI format [B,H,W,C] or list of tensors
|
||||
max_input_images (int): Maximum number of images to include (Gemini may have limits)
|
||||
target_size (tuple): Target size for images (width, height)
|
||||
|
||||
Returns:
|
||||
list: List of processed PIL images ready for the Gemini API
|
||||
"""
|
||||
import torch
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
|
||||
# Handle different input types
|
||||
processed_images = []
|
||||
|
||||
if isinstance(images, torch.Tensor):
|
||||
# Handle 4D tensor [B,H,W,C]
|
||||
if images.dim() == 4:
|
||||
# Limit to max_input_images
|
||||
batch_size = min(images.shape[0], max_input_images)
|
||||
|
||||
for i in range(batch_size):
|
||||
# Get single image tensor [H,W,C]
|
||||
img_tensor = images[i].cpu()
|
||||
|
||||
# Convert to numpy and scale to 0-255
|
||||
img_np = (img_tensor.numpy() * 255).clip(0, 255).astype(np.uint8)
|
||||
|
||||
# Convert to PIL
|
||||
pil_img = Image.fromarray(img_np)
|
||||
|
||||
# Resize to target size if needed
|
||||
if pil_img.size != target_size:
|
||||
pil_img = pil_img.resize(target_size, Image.Resampling.LANCZOS)
|
||||
|
||||
processed_images.append(pil_img)
|
||||
|
||||
# Handle 3D tensor [H,W,C]
|
||||
elif images.dim() == 3:
|
||||
img_tensor = images.cpu()
|
||||
img_np = (img_tensor.numpy() * 255).clip(0, 255).astype(np.uint8)
|
||||
pil_img = Image.fromarray(img_np)
|
||||
|
||||
if pil_img.size != target_size:
|
||||
pil_img = pil_img.resize(target_size, Image.Resampling.LANCZOS)
|
||||
|
||||
processed_images.append(pil_img)
|
||||
|
||||
# Handle list of tensors
|
||||
elif isinstance(images, list):
|
||||
# Limit to max_input_images
|
||||
num_images = min(len(images), max_input_images)
|
||||
|
||||
for i in range(num_images):
|
||||
img = images[i]
|
||||
|
||||
if isinstance(img, torch.Tensor):
|
||||
img_tensor = img.cpu()
|
||||
|
||||
# Handle different tensor dimensions
|
||||
if img_tensor.dim() == 4 and img_tensor.shape[0] == 1: # [1,H,W,C]
|
||||
img_tensor = img_tensor.squeeze(0)
|
||||
|
||||
img_np = (img_tensor.numpy() * 255).clip(0, 255).astype(np.uint8)
|
||||
pil_img = Image.fromarray(img_np)
|
||||
|
||||
if pil_img.size != target_size:
|
||||
pil_img = pil_img.resize(target_size, Image.Resampling.LANCZOS)
|
||||
|
||||
processed_images.append(pil_img)
|
||||
|
||||
return processed_images
|
||||
|
||||
def gemini2_prepare_response(response, width=512, height=512):
|
||||
"""
|
||||
Extract and prepare images from Gemini 2.0 API response.
|
||||
|
||||
Args:
|
||||
response: Gemini API response object
|
||||
width (int): Target width for extracted images
|
||||
height (int): Target height for extracted images
|
||||
|
||||
Returns:
|
||||
tuple: (list of image binaries, response text)
|
||||
"""
|
||||
from io import BytesIO
|
||||
|
||||
images = []
|
||||
response_text = ""
|
||||
|
||||
# Handle empty response
|
||||
if not response or not hasattr(response, 'candidates') or not response.candidates:
|
||||
return images, "No response generated"
|
||||
|
||||
# Process each candidate
|
||||
for candidate in response.candidates:
|
||||
if not hasattr(candidate, 'content') or not hasattr(candidate.content, 'parts'):
|
||||
continue
|
||||
|
||||
for part in candidate.content.parts:
|
||||
# Process text parts
|
||||
if hasattr(part, 'text') and part.text:
|
||||
response_text += part.text + "\n"
|
||||
|
||||
# Process image parts
|
||||
if hasattr(part, 'inline_data') and part.inline_data:
|
||||
try:
|
||||
# Get binary image data
|
||||
image_binary = part.inline_data.data
|
||||
images.append(image_binary)
|
||||
except Exception as e:
|
||||
print(f"Error extracting image from response: {e}")
|
||||
|
||||
return images, response_text
|
||||
|
||||
def gemini2_create_client(api_key):
|
||||
"""
|
||||
Create and return a Gemini API client.
|
||||
|
||||
Args:
|
||||
api_key (str): The Gemini API key
|
||||
|
||||
Returns:
|
||||
Client: Gemini API client object
|
||||
"""
|
||||
try:
|
||||
from google import genai
|
||||
client = genai.Client(api_key=api_key)
|
||||
return client
|
||||
except ImportError:
|
||||
raise ImportError("The google-generativeai package is required. Please install it with: pip install google-generativeai")
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"Failed to create Gemini client: {str(e)}")
|
||||
|
||||
def validate_gemini_key(api_key):
|
||||
"""
|
||||
Validate a Gemini API key by making a simple test request.
|
||||
|
||||
Args:
|
||||
api_key (str): The Gemini API key to validate
|
||||
|
||||
Returns:
|
||||
bool: True if key is valid, False otherwise
|
||||
"""
|
||||
try:
|
||||
from google import genai
|
||||
|
||||
# Initialize client with the key
|
||||
client = genai.Client(api_key=api_key)
|
||||
|
||||
# Try a simple models list request
|
||||
models = client.models.list()
|
||||
|
||||
# If we get here, the key is valid
|
||||
return True
|
||||
except Exception as e:
|
||||
print(f"Invalid Gemini API key: {str(e)}")
|
||||
return False
|
||||
|
||||
Reference in New Issue
Block a user