commit bb2d4c49409a80c671be5c9162bdee5645147f65 Author: Michael Poutre Date: Mon Nov 6 22:48:57 2023 -0800 Initial Commit diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..1b05740 --- /dev/null +++ b/.gitignore @@ -0,0 +1,2 @@ + +__pycache__/ diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..b58201a --- /dev/null +++ b/__init__.py @@ -0,0 +1,9 @@ +from .nodes import ImageWithPrompt + +NODE_CLASS_MAPPINGS = { + "KepOpenAI_ImageWithPrompt": ImageWithPrompt, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "KepOpenAI_ImageWithPrompt": "Image With Prompt", +} diff --git a/lib/__init__.py b/lib/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/lib/credentials.py b/lib/credentials.py new file mode 100644 index 0000000..1638f6f --- /dev/null +++ b/lib/credentials.py @@ -0,0 +1,5 @@ +import os + + +def get_open_ai_api_key() -> str: + return os.environ.get("OPEN_AI_API_KEY", None) diff --git a/lib/image.py b/lib/image.py new file mode 100644 index 0000000..cdbf6dd --- /dev/null +++ b/lib/image.py @@ -0,0 +1,17 @@ +import base64 + +import PIL +import numpy as np +from PIL import Image +from torch import Tensor + + +def tensor2pil(image: Tensor) -> PIL.Image.Image: + return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8)) + + +def pil2base64(image: PIL.Image.Image) -> str: + from io import BytesIO + buffered = BytesIO() + image.save(buffered, format="JPEG") + return base64.b64encode(buffered.getvalue()).decode("utf-8") diff --git a/nodes.py b/nodes.py new file mode 100644 index 0000000..0be59b6 --- /dev/null +++ b/nodes.py @@ -0,0 +1,59 @@ +from typing import Tuple + +import torch +from openai import Client as OpenAIClient + +from .lib import credentials, image + + +class ImageWithPrompt: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "Image": ("IMAGE", {}), + "prompt": ( + "STRING", + { + "multiline": True, + "default": "Generate a high quality prompt to be used for image generation.", + }, + ), + "max_tokens": ("INT", {"min": 1, "max": 2048, "default": 77}), + } + } + + RETURN_TYPES = ("STRING",) + FUNCTION = "generate_completion" + + CATEGORY = "OpenAI" + + def __init__(self): + self.open_ai_client: OpenAIClient = OpenAIClient( + api_key=credentials.get_open_ai_api_key() + ) + + def generate_completion( + self, Image: torch.Tensor, prompt: str, max_tokens: int + ) -> Tuple[str]: + b64image = image.pil2base64(image.tensor2pil(Image)) + response = self.open_ai_client.chat.completions.create( + model="gpt-4-vision-preview", + max_tokens=max_tokens, + messages=[ + { + "role": "user", + "content": [ + {"type": "text", "text": prompt}, + { + "type": "image_url", + "image_url": {"url": f"data:image/jpeg;base64,{b64image}"}, + }, + ], + } + ], + ) + if len(response.choices) == 0: + raise Exception("No response from OpenAI API") + + return (response.choices[0].message.content,) diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..ec838c5 --- /dev/null +++ b/requirements.txt @@ -0,0 +1 @@ +openai