Added ClaudeAPI

This commit is contained in:
daxcay
2024-11-12 15:16:29 +05:30
parent 17be43a2e2
commit bc57dc27d3
7 changed files with 166 additions and 6 deletions
+6
View File
@@ -34,6 +34,8 @@ from .classes.DataSet_OpenAIChatImage import N_CLASS_MAPPINGS as OpenAIChatImage
from .classes.DataSet_OpenAIChatImageBatch import N_CLASS_MAPPINGS as OpenAIChatImageBatchMappings, N_DISPLAY_NAME_MAPPINGS as OpenAIChatImageBatchNameMappings
from .classes.DataSet_GroqChat import N_CLASS_MAPPINGS as GroqChatMappings, N_DISPLAY_NAME_MAPPINGS as GroqChatNameMappings
from .classes.DataSet_GroqChatImage import N_CLASS_MAPPINGS as GroqChatImageMappings, N_DISPLAY_NAME_MAPPINGS as GroqChatImageNameMappings
from .classes.DataSet_ClaudeAIChat import N_CLASS_MAPPINGS as ClaudeChatMappings, N_DISPLAY_NAME_MAPPINGS as ClaudeChatNameMappings
from .classes.DataSet_ClaudeAIChatImage import N_CLASS_MAPPINGS as ClaudeChatImageMappings, N_DISPLAY_NAME_MAPPINGS as ClaudeChatImageNameMappings
NODE_CLASS_MAPPINGS = {}
NODE_CLASS_MAPPINGS.update(VisualizerMappings)
@@ -53,6 +55,8 @@ NODE_CLASS_MAPPINGS.update(OpenAIChatImageMappings)
NODE_CLASS_MAPPINGS.update(OpenAIChatImageBatchMappings)
NODE_CLASS_MAPPINGS.update(GroqChatMappings)
NODE_CLASS_MAPPINGS.update(GroqChatImageMappings)
NODE_CLASS_MAPPINGS.update(ClaudeChatMappings)
NODE_CLASS_MAPPINGS.update(ClaudeChatImageMappings)
NODE_DISPLAY_NAME_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS.update(VisualizerNameMappings)
@@ -72,5 +76,7 @@ NODE_DISPLAY_NAME_MAPPINGS.update(OpenAIChatImageNameMappings)
NODE_DISPLAY_NAME_MAPPINGS.update(OpenAIChatImageBatchNameMappings)
NODE_DISPLAY_NAME_MAPPINGS.update(GroqChatNameMappings)
NODE_DISPLAY_NAME_MAPPINGS.update(GroqChatImageNameMappings)
NODE_DISPLAY_NAME_MAPPINGS.update(ClaudeChatNameMappings)
NODE_DISPLAY_NAME_MAPPINGS.update(ClaudeChatImageNameMappings)
WEB_DIRECTORY = "./web"
+59
View File
@@ -0,0 +1,59 @@
import os
import anthropic
class DataSet_ClaudeAIChat:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
api_models = [ "claude-3-5-sonnet-latest", "claude-3-5-haiku-latest", "claude-3-opus-latest" ]
return {
"required": {
"model": (api_models, {"default": api_models[0]}),
# "system_prompt": ("STRING", {"multiline": True, "default": ""}),
"user_prompt": ("STRING", {"multiline": True, "default": ""}),
"max_tokens": ("INT", {"default": 1024})
}
}
FUNCTION = "generate"
RETURN_TYPES = ("STRING",)
def generate(self, model, user_prompt, max_tokens):
try:
api_client = anthropic.Anthropic(api_key=os.environ.get("ANTHROPIC_API_KEY"))
chat_completion = api_client.messages.create(
model=model,
messages=[
{
"role": "user",
"content": user_prompt,
}
],
max_tokens=max_tokens,
)
msg = ""
for message in chat_completion.content:
if message.type == 'text':
msg = message.text
return (msg,)
except Exception as e:
return (f"Error: {str(e)}",)
N_CLASS_MAPPINGS = {
"DataSet_ClaudeAIChat": DataSet_ClaudeAIChat,
}
N_DISPLAY_NAME_MAPPINGS = {
"DataSet_ClaudeAIChat": "DataSet_ClaudeAIChat",
}
+90
View File
@@ -0,0 +1,90 @@
import base64
import io
from PIL import Image
import numpy as np
import anthropic
import os
api_key = os.environ.get("ANTHROPIC_API_KEY")
class DataSet_ClaudeAIChatImage:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
api_models = ["claude-3-5-sonnet-latest",
"claude-3-5-haiku-latest", "claude-3-opus-latest"]
return {
"required": {
"image": ("IMAGE",),
"model": (api_models, {"default": api_models[0]}),
# "system_prompt": ("STRING", {"multiline": True, "default": ""}),
"user_prompt": ("STRING", {"multiline": True, "default": ""}),
"max_tokens": ("INT", {"default": 1024})
}
}
RETURN_TYPES = ("STRING",)
FUNCTION = "generate"
CATEGORY = "🔶DATASET🔶"
def to_base64(self, image):
image = image[0]
i = 255. * image.cpu().numpy()
img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8))
buffered = io.BytesIO()
img.save(buffered, format="PNG")
return base64.b64encode(buffered.getvalue()).decode("utf-8")
def generate(self, image, model, user_prompt, max_tokens):
try:
base64img = self.to_base64(image)
api_client = anthropic.Anthropic(api_key=os.environ.get("ANTHROPIC_API_KEY"))
chat_completion = api_client.messages.create(
model=model,
max_tokens=max_tokens,
messages=[
{
"role": "user",
"content": [
{
"type": "image",
"source": {
"type": "base64",
"media_type": "image/png",
"data": base64img,
},
},
{"type": "text", "text": user_prompt}
],
}
],
)
msg = ""
for message in chat_completion.content:
if message.type == 'text':
msg = message.text
return (msg,)
except Exception as e:
return (f"Error: {str(e)}",)
N_CLASS_MAPPINGS = {
"DataSet_ClaudeAIChatImage": DataSet_ClaudeAIChatImage,
}
N_DISPLAY_NAME_MAPPINGS = {
"DataSet_ClaudeAIChatImage": "DataSet_ClaudeAIChatImage",
}
+3 -2
View File
@@ -1,4 +1,5 @@
from openai import OpenAI
import os
class DataSet_OpenAIChat:
@@ -11,7 +12,6 @@ class DataSet_OpenAIChat:
"required": {
"model": (["gpt-4", "gpt-4-32k", "gpt-3.5-turbo", "gpt-4-0125-preview", "gpt-4-turbo-preview", "gpt-4-1106-preview", "gpt-4-0613"], {"default": "gpt-3.5-turbo"}),
"api_url": ("STRING", {"multiline": False, "default": "https://api.openai.com/v1"}),
"api_key": ("STRING", {"multiline": False}),
"prompt": ("STRING", {"multiline": True, "default": ""}),
"token_length": ("INT", {"default": 1024})
}
@@ -21,8 +21,9 @@ class DataSet_OpenAIChat:
FUNCTION = "generate"
CATEGORY = "🔶DATASET🔶"
def generate(self, model, api_url, api_key, prompt, token_length):
def generate(self, model, api_url, prompt, token_length):
try:
api_key = os.environ.get("OPENAI_API_KEY")
ai = OpenAI(api_key=api_key, base_url=api_url)
if not api_key:
return "OpenAI API key is required."
+4 -2
View File
@@ -3,6 +3,7 @@ import io
from PIL import Image
import numpy as np
from openai import OpenAI
import os
class DataSet_OpenAIChatImage:
@@ -18,7 +19,7 @@ class DataSet_OpenAIChatImage:
"prompt": ("STRING", {"multiline": True, "default": ""}),
"model": (["gpt-4o","gpt-4", "gpt-4-32k", "gpt-3.5-turbo", "gpt-4-0125-preview", "gpt-4-turbo-preview", "gpt-4-1106-preview", "gpt-4-0613"], {"default": "gpt-4o"}),
"api_url": ("STRING", {"multiline": False, "default": "https://api.openai.com/v1"}),
"api_key": ("STRING", {"multiline": False}),
# "api_key": ("STRING", {"multiline": False}),
"token_length": ("INT", {"default": 1024})
}
}
@@ -35,9 +36,10 @@ class DataSet_OpenAIChatImage:
img.save(buffered, format="PNG")
return base64.b64encode(buffered.getvalue()).decode("utf-8")
def generate(self, image, image_detail, model, api_url, api_key, prompt, token_length):
def generate(self, image, image_detail, model, api_url, prompt, token_length):
try:
api_key = os.environ.get("OPENAI_API_KEY")
ai = OpenAI(api_key=api_key, base_url=api_url)
base64img = self.to_base64(image)
if not api_key:
+2 -1
View File
@@ -38,7 +38,7 @@ class DataSet_OpenAIChatImageBatch:
img.save(buffered, format="PNG")
return base64.b64encode(buffered.getvalue()).decode("utf-8")
def generate(self, images, image_detail, model, api_url, api_key, prompt, token_length):
def generate(self, images, image_detail, model, api_url, prompt, token_length):
try:
@@ -53,6 +53,7 @@ class DataSet_OpenAIChatImageBatch:
for image in images:
api_key = os.environ.get("OPENAI_API_KEY")
ai = OpenAI(api_key=api_key, base_url=api_url)
base64img = self.to_base64(image)
if not api_key:
+2 -1
View File
@@ -3,4 +3,5 @@ wordcloud
networkx
pandas
openai
groq
groq
anthropic