Initial Commit

This commit is contained in:
Michael Poutre
2023-11-06 22:48:57 -08:00
commit bb2d4c4940
7 changed files with 93 additions and 0 deletions
+2
View File
@@ -0,0 +1,2 @@
__pycache__/
+9
View File
@@ -0,0 +1,9 @@
from .nodes import ImageWithPrompt
NODE_CLASS_MAPPINGS = {
"KepOpenAI_ImageWithPrompt": ImageWithPrompt,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"KepOpenAI_ImageWithPrompt": "Image With Prompt",
}
View File
+5
View File
@@ -0,0 +1,5 @@
import os
def get_open_ai_api_key() -> str:
return os.environ.get("OPEN_AI_API_KEY", None)
+17
View File
@@ -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")
+59
View File
@@ -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,)
+1
View File
@@ -0,0 +1 @@
openai