Initial Commit
This commit is contained in:
@@ -0,0 +1,2 @@
|
||||
|
||||
__pycache__/
|
||||
@@ -0,0 +1,9 @@
|
||||
from .nodes import ImageWithPrompt
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"KepOpenAI_ImageWithPrompt": ImageWithPrompt,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"KepOpenAI_ImageWithPrompt": "Image With Prompt",
|
||||
}
|
||||
@@ -0,0 +1,5 @@
|
||||
import os
|
||||
|
||||
|
||||
def get_open_ai_api_key() -> str:
|
||||
return os.environ.get("OPEN_AI_API_KEY", None)
|
||||
@@ -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")
|
||||
@@ -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,)
|
||||
@@ -0,0 +1 @@
|
||||
openai
|
||||
Reference in New Issue
Block a user