diff --git a/IFLLMNode.py b/IFLLMNode.py index a7350f3..2660fe8 100644 --- a/IFLLMNode.py +++ b/IFLLMNode.py @@ -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) diff --git a/anthropic_api.py b/anthropic_api.py index 91a584a..8b443ca 100644 --- a/anthropic_api.py +++ b/anthropic_api.py @@ -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): """ diff --git a/groq_api.py b/groq_api.py index 4be39b7..feaf647 100644 --- a/groq_api.py +++ b/groq_api.py @@ -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], diff --git a/utils.py b/utils.py index 76ccbf3..0fe76df 100644 --- a/utils.py +++ b/utils.py @@ -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