Update nodes.py

This commit is contained in:
shi3z
2024-03-13 17:25:42 +09:00
committed by GitHub
parent d82a8e8508
commit 7fd63b3e73
+97 -8
View File
@@ -9,15 +9,98 @@ import io
import numpy as np
import json
import os
import server
#OpenAIのAPIキーを入力
#OpenAI API Key
os.environ["OPENAI_API_KEY"] = "sk-"
from openai import OpenAI
# MemeplexのUIDとAPIKEYを入力
# Memeplex UID/APIKey (Optional)
uid=""
apikey=""
def generate_image(prompt,negative,model="custom_sdxl_anime1",qty=1,width=1024,height=1024):
def gpt(utterance,model="gpt-3.5-turbo"):
messages=[
{"role": "system", "content": "You are specialist of LLM and Stable Diffusion,and artist."},
]
messages.append({"role": "user", "content": utterance})
#response = openai.chat(
print("call gpt")
response = client.chat.completions.create(
model=model,
messages=messages,
#model="gpt-3.5-turbo-1106",
#model="gpt-3.5-turbo",
response_format={"type":"json_object"},
)
messages.append({"role": "assistant", "content": response.choices[0].message.content})
return response.choices[0].message.content
import requests
class GPT:
@classmethod
def INPUT_TYPES(s):
return {"required": {"text": ("STRING",{"forceInput": True}),
"prompt": ("STRING", {"multiline": True}),
"model": (["gpt-3.5-turbo","gpt-4-1106-preview"],)},
"optional": {"result": ("STRING", {"multiline": True})}
}
RETURN_TYPES = ("STRING",)
FUNCTION = "run"
CATEGORY = "Memeplex"
def run(self, text, model,prompt,result=None):
for i in range(5):
prompt=("以下の条件を満たすStableDiffusion用のプロンプトをできるだけ詳細にかつ、masterpieceなどのいい感じのワードを加え、高画質にできるだけ近づけるようにして最低30語は使え。全て英語で考えろ。JSON形式で、回答はpromptというプロパティに格納しろ。日本語は一切使わず、プロンプト以外の余計なことも言うな\n"+
prompt+"\n"+text)
response=gpt(prompt,model=model)
response=json.loads(response)
if "prompt" in response:
response=response["prompt"]
break
time.sleep(5)
if i >= 4:
print("gpt response error")
throw("gpt response error")
url = "http://localhost:8188/memeplex/update_text"
headers = {'Content-Type': 'application/json'}
data = {
"text": response,
}
response = requests.post(url, headers=headers, json=data)
if response.status_code == 200:
print("データの送信に成功しました。")
else:
print(f"データの送信に失敗しました。ステータスコード: {response.status_code}")
return (text,)
class TextSend:
@classmethod
def INPUT_TYPES(s):
return {"required": {"text": ("STRING", {"forceInput": True})}}
OUTPUT_NODE = True
RETURN_TYPES = ()
FUNCTION = "run"
CATEGORY = "MyCustomClient"
def run(self, text):
# テキストをクライアントに送信する(コンソールアプリの方で受け取る)
text = text + " (from TextSend node)"
server.PromptServer.instance.send_sync("send_text", {"text": text})
return ()
NODE_CLASS_MAPPINGS = {
"TextSend": TextSend,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"TextSend": "TextSend",
}
def generate_image(prompt,negative,model="custom_sdxl_anime2",qty=1,width=1024,height=1024):
"""
response = client.images.generate(
model="dall-e-3",
@@ -36,6 +119,10 @@ def generate_image(prompt,negative,model="custom_sdxl_anime1",qty=1,width=1024,h
return result["result"]
class TextInput:
@classmethod
def INPUT_TYPES(s):
@@ -108,10 +195,10 @@ class MemeplexRender:
def INPUT_TYPES(s):
return {"required": {
"prompt": ("STRING", {"forceInput": True}),
"width": ("INT", {"default": 512, "min": 512, "max": 1024, "step": 64}),
"height": ("INT", {"default": 512, "min": 512, "max": 1024, "step": 64}),
"width": ("INT", {"default": 1024, "min": 512, "max": 2048, "step": 64}),
"height": ("INT", {"default": 1024, "min": 512, "max": 2048, "step": 64}),
"qty": ("INT", {"default": 9, "min": 1, "max": 9, "step": 1}),
"model": (["trinart","StableDiffusion-v1-5","StableDiffusion-v2-0"],)
"model": (["SDXL1.0","trinart","StableDiffusion-v1-5","StableDiffusion-v2-0"],)
}
,
"optional":{
@@ -218,12 +305,14 @@ NODE_CLASS_MAPPINGS = {
"TextInput": TextInput,
"MemeplexCustomSDXLRender": MemeplexCustomSDXLRender,
"MemeplexRender": MemeplexRender,
"DallERender":DallERender
"DallERender":DallERender,
"GPT":GPT
}
NODE_DISPLAY_NAME_MAPPINGS = {
"TextInput": "TextInput",
"MemeplexCustomSDXLRender": "MemeplexCustomSDXLRender",
"MemeplexRender":"MemeplexRender",
"DallERender":"DallERender"
"DallERender":"DallERender",
"GPT":"GPT"
}