优化GPT
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||

|
||||
|
||||
+4
-3
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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"))
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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'](
|
||||
|
||||
Reference in New Issue
Block a user