优化GPT

This commit is contained in:
shadowcz007
2023-12-05 13:47:08 +08:00
parent dfe720f3ec
commit d7d46682fc
10 changed files with 87 additions and 54 deletions
+1 -1
View File
@@ -24,7 +24,7 @@ https://github.com/shadowcz007/comfyui-mixlab-nodes/assets/12645064/e7e77f90-e43
### GPT
> ChatGPT、ChatGLM3 , Some code provided by rui.
> ChatGPT、ChatGLM3 , Some code provided by rui. If you are using OpenAI's service, fill in https://api.openai.com/v1 . If you are using a local LLM service, fill in http://127.0.0.1:xxxx/v1
![gpt-workflow.svg](./assets/gpt-workflow.svg)
+4 -3
View File
@@ -260,7 +260,7 @@ from .nodes.ImageNode import TransparentImage,LoadImagesFromPath,AreaToMask,Smoo
from .nodes.Vae import VAELoader,VAEDecode
from .nodes.ScreenShareNode import ScreenShareNode,FloatingVideo
from .nodes.Clipseg import CLIPSeg,CombineMasks
from .nodes.ChatGPT import ChatGPTNode,SessionHistory
from .nodes.ChatGPT import ChatGPTNode,ShowTextForGPT,CharacterInText
# 要导出的所有节点及其名称的字典
# 注意:名称应全局唯一
@@ -282,7 +282,8 @@ NODE_CLASS_MAPPINGS = {
"CLIPSeg_":CLIPSeg,
"CombineMasks_":CombineMasks,
"ChatGPT":ChatGPTNode,
"SessionHistory":SessionHistory
"ShowTextForGPT":ShowTextForGPT,
"CharacterInText":CharacterInText
}
# 一个包含节点友好/可读的标题的字典
@@ -294,7 +295,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"ScreenShare":"ScreenShare #Mixlab",
"FloatingVideo":"FloatingVideo #Mixlab",
"ChatGPT":"ChatGPT #Mixlab",
"SessionHistory":"SessionHistory #Mixlab"
"ShowTextForGPT":"ShowTextForGPT #Mixlab"
}
# web ui的节点功能
+53 -22
View File
@@ -66,7 +66,7 @@ class ChatGPTNode:
def __init__(self):
# self.__client = OpenAI()
self.session_history = [] # 用于存储会话历史的列表
self.seed=0
# self.seed=0
self.system_content="You are ChatGPT, a large language model trained by OpenAI. Answer as concisely as possible."
@classmethod
@@ -84,15 +84,16 @@ class ChatGPTNode:
"model": (["gpt-3.5-turbo", "gpt-3.5-turbo-16k-0613", "gpt-4-0613","gpt-4-1106-preview"],
{"default": "gpt-3.5-turbo"}),
"seed": ("INT", {"default": 0, "min": 0, "max": 10000, "step": 1}),
"context_size":("INT", {"default": 1, "min": 0, "max":30, "step": 1}),
},
}
RETURN_TYPES = ("STRING","STRING",)
RETURN_NAMES = ("text","session_history",)
RETURN_TYPES = ("STRING","STRING","STRING",)
RETURN_NAMES = ("text","messages","session_history",)
FUNCTION = "generate_contextual_text"
CATEGORY = "Mixlab/GPT"
CATEGORY = "♾️Mixlab/GPT"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,False,)
OUTPUT_IS_LIST = (False,False,False,)
def generate_contextual_text(self,
@@ -101,18 +102,13 @@ class ChatGPTNode:
prompt,
system_content,
model,
seed):
print(api_key,
api_url,
prompt,
system_content,
model,
seed)
seed,context_size):
# print(api_key!='',api_url,prompt,system_content,model,seed)
# 可以选择保留会话历史以维持上下文记忆
# 或者在此处清除会话历史 self.session_history.clear()
if seed!=self.seed:
self.seed=seed
self.session_history=[]
# if seed!=self.seed:
# self.seed=seed
# self.session_history=[]
# 把系统信息和初始信息添加到会话历史中
if system_content:
@@ -129,21 +125,32 @@ class ChatGPTNode:
# 把用户的提示添加到会话历史中
# 调用API时传递整个会话历史
messages=[{"role": "system", "content": self.system_content}]+self.session_history+[{"role": "user", "content": prompt}]
def crop_list_tail(lst, size):
if size >= len(lst):
return lst
elif size==0:
return []
else:
return lst[-size:]
session_history=crop_list_tail(self.session_history,context_size)
messages=[{"role": "system", "content": self.system_content}]+session_history+[{"role": "user", "content": prompt}]
response_content = chat(client,model,messages)
self.session_history=self.session_history+[{"role": "user", "content": prompt}]+[{'role':'assistant',"content":response_content}]
return (response_content,json.dumps(self.session_history, indent=4),)
return (response_content,json.dumps(messages, indent=4),json.dumps(self.session_history, indent=4),)
class SessionHistory:
class ShowTextForGPT:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"session_history": ("STRING", {"forceInput": True}),
"text": ("STRING", {"forceInput": True}),
}
}
@@ -153,8 +160,32 @@ class SessionHistory:
OUTPUT_NODE = True
OUTPUT_IS_LIST = (True,)
CATEGORY = "Mixlab/GPT"
CATEGORY = "♾️Mixlab/GPT"
def run(self, session_history):
def run(self, text):
# print(session_history)
return {"ui": {"text": session_history}, "result": (session_history,)}
return {"ui": {"text": text}, "result": (text,)}
class CharacterInText:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"text": ("STRING", {"multiline": True}),
"character": ("STRING", {"multiline": True}),
}
}
INPUT_IS_LIST = False
RETURN_TYPES = ("INT",)
FUNCTION = "run"
# OUTPUT_NODE = True
OUTPUT_IS_LIST = (False,)
CATEGORY = "♾️Mixlab/GPT"
def run(self, text,character):
b=[1 if character in text else 0]
return (b,)
+2 -2
View File
@@ -99,7 +99,7 @@ class CLIPSeg:
}
}
CATEGORY = "Mixlab/mask"
CATEGORY = "♾️Mixlab/mask"
RETURN_TYPES = ("MASK", "IMAGE", "IMAGE",)
RETURN_NAMES = ("Mask","Heatmap Mask", "BW Mask")
@@ -204,7 +204,7 @@ class CombineMasks:
},
}
CATEGORY = "Mixlab/mask"
CATEGORY = "♾️Mixlab/mask"
RETURN_TYPES = ("MASK", "IMAGE", "IMAGE",)
RETURN_NAMES = ("Combined Mask","Heatmap Mask", "BW Mask")
+9 -9
View File
@@ -341,7 +341,7 @@ class SmoothMask:
FUNCTION = "run"
CATEGORY = "Mixlab/mask"
CATEGORY = "♾️Mixlab/mask"
INPUT_IS_LIST = False
@@ -389,7 +389,7 @@ class FeatheredMask:
FUNCTION = "run"
CATEGORY = "Mixlab/mask"
CATEGORY = "♾️Mixlab/mask"
OUTPUT_IS_LIST = (False,)
@@ -456,7 +456,7 @@ class SplitLongMask:
FUNCTION = "run"
CATEGORY = "Mixlab/mask"
CATEGORY = "♾️Mixlab/mask"
OUTPUT_IS_LIST = (True,)
@@ -499,7 +499,7 @@ class TransparentImage:
FUNCTION = "run"
CATEGORY = "Mixlab/image"
CATEGORY = "♾️Mixlab/image"
# INPUT_IS_LIST = True, 一个batch传进来
OUTPUT_IS_LIST = (True,True,True,)
@@ -568,7 +568,7 @@ class EnhanceImage:
FUNCTION = "run"
CATEGORY = "Mixlab/image"
CATEGORY = "♾️Mixlab/image"
INPUT_IS_LIST = False
@@ -627,7 +627,7 @@ class LoadImagesFromPath:
FUNCTION = "run"
CATEGORY = "Mixlab/image"
CATEGORY = "♾️Mixlab/image"
# INPUT_IS_LIST = True
OUTPUT_IS_LIST = (True,True,False,)
@@ -688,7 +688,7 @@ class ImageCropByAlpha:
FUNCTION = "run"
CATEGORY = "Mixlab/image"
CATEGORY = "♾️Mixlab/image"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,)
@@ -719,7 +719,7 @@ class AreaToMask:
FUNCTION = "run"
CATEGORY = "Mixlab/mask"
CATEGORY = "♾️Mixlab/mask"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,)
@@ -754,7 +754,7 @@ class FaceToMask:
FUNCTION = "run"
CATEGORY = "Mixlab/mask"
CATEGORY = "♾️Mixlab/mask"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,)
+2 -2
View File
@@ -79,7 +79,7 @@ class RandomPrompt:
FUNCTION = "run"
CATEGORY = "Mixlab/prompt"
CATEGORY = "♾️Mixlab/prompt"
OUTPUT_IS_LIST = (True,)
OUTPUT_NODE = True
@@ -158,7 +158,7 @@ class RunWorkflow:
FUNCTION = "run"
CATEGORY = "Mixlab/workflow"
CATEGORY = "♾️Mixlab/workflow"
OUTPUT_IS_LIST = (True,)
OUTPUT_NODE = True
+2 -2
View File
@@ -89,7 +89,7 @@ class ScreenShareNode:
FUNCTION = "run"
CATEGORY = "Mixlab/image"
CATEGORY = "♾️Mixlab/image"
# INPUT_IS_LIST = True
OUTPUT_IS_LIST = (False,False,False)
@@ -114,7 +114,7 @@ class FloatingVideo:
OUTPUT_NODE = True
FUNCTION = "run"
CATEGORY = "Mixlab/image"
CATEGORY = "♾️Mixlab/image"
# INPUT_IS_LIST = True
# OUTPUT_IS_LIST = (False,False,)
+2 -2
View File
@@ -145,7 +145,7 @@ class VAELoader:
RETURN_TYPES = ("VAE",)
FUNCTION = "load_vae"
CATEGORY = "Mixlab/ConsistencyDecoder"
CATEGORY = "♾️Mixlab/ConsistencyDecoder"
#TODO: scale factor?
def load_vae(self, vae_name):
@@ -165,7 +165,7 @@ class VAEDecode:
RETURN_TYPES = ("IMAGE",)
FUNCTION = "decode"
CATEGORY = "Mixlab/ConsistencyDecoder"
CATEGORY = "♾️Mixlab/ConsistencyDecoder"
def decode(self, vae, samples):
image = vae.decode(samples["samples"].to("cuda:0"))
+1 -1
View File
@@ -16,7 +16,7 @@ fetch(`https://api.github.com/repos/${repoOwner}/${repoName}/releases/latest`)
app.ui.dialog.show(`<h4 style="font-size: 18px;">${repoName} <br>
Latest release version: ${latestVersion}</h4>
<p>Please proceed to the official repository to download the latest version.</p>
<a style=" color: #2196F3;
<a style="color: #2196F3;
font-size: 18px;
font-weight: 800;
letter-spacing: 2px;
+11 -10
View File
@@ -61,7 +61,8 @@ app.registerExtension({
return [128, 24] // a method to compute the current size of the widget
},
async serializeValue (nodeId, widgetIndex) {
return localStorage.getItem('_mixlab_api_key') || ''
//localStorage.getItem('_mixlab_api_key') || ''
return 'by Mixlab'
}
}
// widget.something = something; // maybe adds stuff to it
@@ -81,7 +82,7 @@ app.registerExtension({
return [128, 24] // a method to compute the current size of the widget
},
async serializeValue (nodeId, widgetIndex) {
return localStorage.getItem('_mixlab_api_url') || ''
return localStorage.getItem('_mixlab_api_url') || 'https://api.openai.com/v1'
}
}
// widget.something = something; // maybe adds stuff to it
@@ -101,7 +102,7 @@ app.registerExtension({
const api_key = this.widgets.filter(w => w.name == 'api_key')[0]
const api_url = this.widgets.filter(w => w.name == 'api_url')[0]
console.log('api_key', api_key, api_url)
// console.log('api_key', api_key, api_url)
const widget = {
type: 'div',
@@ -129,19 +130,19 @@ app.registerExtension({
inputKey.style = `margin:4px 48px;`
inputUrl.style = `margin:4px 48px`
inputUrl.value = localStorage.getItem('_mixlab_api_url') || ''
inputKey.value = localStorage.getItem('_mixlab_api_key') || ''
inputUrl.value = localStorage.getItem('_mixlab_api_url') || 'https://api.openai.com/v1'
inputKey.value = localStorage.getItem('_mixlab_api_key') || 'by Mixlab'
widget.div.appendChild(inputKey)
widget.div.appendChild(inputUrl)
inputKey.addEventListener('change', () => {
api_key.serializeValue = () => inputKey.value || ''
api_key.serializeValue = () => inputKey.value || 'by Mixlab'
localStorage.setItem('_mixlab_api_key', inputKey.value)
})
inputUrl.addEventListener('change', () => {
api_url.serializeValue = () => inputUrl.value || ''
api_url.serializeValue = () => inputUrl.value || 'https://api.openai.com/v1'
localStorage.setItem('_mixlab_api_url', inputUrl.value)
})
@@ -165,9 +166,9 @@ app.registerExtension({
})
app.registerExtension({
name: 'Mixlab.GPT.SessionHistory',
name: 'Mixlab.GPT.ShowTextForGPT',
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeData.name === 'SessionHistory') {
if (nodeData.name === 'ShowTextForGPT') {
function populate (text) {
if (this.widgets) {
// const pos = this.widgets.findIndex((w) => w.name === "text");
@@ -184,7 +185,7 @@ app.registerExtension({
this.widgets.length = 0
}
console.log('SessionHistory', this.widgets, text)
console.log('ShowTextForGPT', this.widgets, text)
for (const list of text) {
const w = ComfyWidgets['STRING'](