From bb2d4c49409a80c671be5c9162bdee5645147f65 Mon Sep 17 00:00:00 2001 From: Michael Poutre Date: Mon, 6 Nov 2023 22:48:57 -0800 Subject: [PATCH] Initial Commit --- .gitignore | 2 ++ __init__.py | 9 +++++++ lib/__init__.py | 0 lib/credentials.py | 5 ++++ lib/image.py | 17 +++++++++++++ nodes.py | 59 ++++++++++++++++++++++++++++++++++++++++++++++ requirements.txt | 1 + 7 files changed, 93 insertions(+) create mode 100644 .gitignore create mode 100644 __init__.py create mode 100644 lib/__init__.py create mode 100644 lib/credentials.py create mode 100644 lib/image.py create mode 100644 nodes.py create mode 100644 requirements.txt 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