From 10080ad3d6dd8a831f4a1f87a343fc6d50d04676 Mon Sep 17 00:00:00 2001 From: WQJ Date: Tue, 24 Dec 2024 21:41:14 +0800 Subject: [PATCH] first commit --- .gitignore | 1 + __init__.py | 11 +++++ llm_node.py | 124 +++++++++++++++++++++++++++++++++++++++++++++++ requirements.txt | 4 ++ 4 files changed, 140 insertions(+) create mode 100644 .gitignore create mode 100644 __init__.py create mode 100644 llm_node.py create mode 100644 requirements.txt diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..ef81b1e --- /dev/null +++ b/.gitignore @@ -0,0 +1 @@ +/.venv/ diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..95e550b --- /dev/null +++ b/__init__.py @@ -0,0 +1,11 @@ +from .llm_node import LLMImageDescription + +NODE_CLASS_MAPPINGS = { + "LLMImageDescription": LLMImageDescription +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "LLMImageDescription": "LLM Image Description" +} + +__version__ = "1.0.0" diff --git a/llm_node.py b/llm_node.py new file mode 100644 index 0000000..6db846f --- /dev/null +++ b/llm_node.py @@ -0,0 +1,124 @@ +import numpy as np +import base64 +from PIL import Image +from io import BytesIO +from openai import OpenAI + + +class LLMImageDescription: + def __init__(self): + self.output_dir = "output" + self.type = "output" + self._client = None + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + "model": (["gpt-4o", "claude-3-5-sonnet-20241022", "gemini-1.5-pro-latest", "gemini-2.0-flash-exp"],), + "api_url": ("STRING", { + "default": "", + "multiline": False + }), + "api_key": ("STRING", { + "default": "", + "multiline": False + }), + "prompt_template": ("STRING", { + "default": "Please describe this image in detail:", + "multiline": True + }) + } + } + + RETURN_TYPES = ("STRING",) + RETURN_NAMES = ("description",) + FUNCTION = "process_image" + CATEGORY = "image/text" + + def get_client(self, api_key, api_url=None): + """Get or create OpenAI client""" + if not self._client: + kwargs = {"api_key": api_key} + if api_url: + kwargs["base_url"] = api_url + self._client = OpenAI(**kwargs) + return self._client + + @staticmethod + def convert_image_to_base64(image): + # Convert PyTorch tensor to PIL Image + image = image.cpu().numpy() + image = (image * 255).astype(np.uint8) + if image.shape[0] == 3: # If image is in CHW format + image = np.transpose(image, (1, 2, 0)) + pil_image = Image.fromarray(image) + + # Convert PIL Image to base64 + buffered = BytesIO() + pil_image.save(buffered, format="PNG") + img_str = base64.b64encode(buffered.getvalue()).decode() + return img_str + + @staticmethod + def process_with_openai_compatible(prompt_template, base64_image, client, model): + """Process image using OpenAI API format""" + try: + response = client.chat.completions.create( + model=model, + messages=[ + { + "role": "user", + "content": [ + { + "type": "text", + "text": prompt_template + }, + { + "type": "image_url", + "image_url": { + "url": f"data:image/png;base64,{base64_image}" + } + } + ] + } + ], + max_tokens=300 + ) + return response.choices[0].message.content + except Exception as e: + print(f"API error: {str(e)}") + return f"Error: API request failed - {str(e)}" + + def process_image(self, image, model, api_url, api_key, prompt_template): + try: + # Convert the first image in batch to base64 + if len(image.shape) == 4: + image = image[0] + base64_image = self.convert_image_to_base64(image) + + # Get appropriate API URL based on model + if not api_url: + api_urls = { + "gpt-4o": "https://api.openai.com/v1/chat/completions", + "claude-3-5-sonnet-20241022": "https://api.anthropic.com/v1/messages", + "gemini-1.5-pro-latest": "https://generativelanguage.googleapis.com/v1beta/models/gemini-1.5-pro-latest:generateContent", + "gemini-2.0-flash-exp": "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.0-flash-exp:generateContent" + } + api_url = api_urls.get(model, "") + + # Process using unified OpenAI SDK + client = self.get_client(api_key, api_url) + description = self.process_with_openai_compatible(prompt_template, base64_image, client, model) + + return (description,) + + except Exception as e: + print(f"Error in image description: {str(e)}") + return (f"Error: Failed to generate image description. {str(e)}",) + + def __del__(self): + """Cleanup client on deletion""" + if self._client: + self._client.close() diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..7a0fed2 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,4 @@ +openai +numpy +requests +pillow