Add files via upload

This commit is contained in:
ImpactFrames
2025-03-15 13:37:21 +00:00
committed by GitHub
parent 027e7724c6
commit dd7a35de2b
4 changed files with 456 additions and 60 deletions
+232 -2
View File
@@ -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
View File
@@ -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
View File
@@ -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],
+163
View File
@@ -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