diff --git a/README.md b/README.md index e060be7..7250caa 100644 --- a/README.md +++ b/README.md @@ -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) diff --git a/__init__.py b/__init__.py index 4e2e9a9..54a202a 100644 --- a/__init__.py +++ b/__init__.py @@ -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的节点功能 diff --git a/nodes/ChatGPT.py b/nodes/ChatGPT.py index 3ce168c..f2fc80b 100644 --- a/nodes/ChatGPT.py +++ b/nodes/ChatGPT.py @@ -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,) + diff --git a/nodes/Clipseg.py b/nodes/Clipseg.py index e2252e6..a2e94c8 100644 --- a/nodes/Clipseg.py +++ b/nodes/Clipseg.py @@ -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") diff --git a/nodes/ImageNode.py b/nodes/ImageNode.py index 228b968..894927f 100644 --- a/nodes/ImageNode.py +++ b/nodes/ImageNode.py @@ -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,) diff --git a/nodes/PromptNode.py b/nodes/PromptNode.py index d71cdde..245dab3 100644 --- a/nodes/PromptNode.py +++ b/nodes/PromptNode.py @@ -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 diff --git a/nodes/ScreenShareNode.py b/nodes/ScreenShareNode.py index aa72b2f..6f2ffb6 100644 --- a/nodes/ScreenShareNode.py +++ b/nodes/ScreenShareNode.py @@ -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,) diff --git a/nodes/Vae.py b/nodes/Vae.py index 8efb0de..413cbc9 100644 --- a/nodes/Vae.py +++ b/nodes/Vae.py @@ -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")) diff --git a/web/javascript/checkVersion_mixlab.js b/web/javascript/checkVersion_mixlab.js index 1347b41..9225629 100644 --- a/web/javascript/checkVersion_mixlab.js +++ b/web/javascript/checkVersion_mixlab.js @@ -16,7 +16,7 @@ fetch(`https://api.github.com/repos/${repoOwner}/${repoName}/releases/latest`) app.ui.dialog.show(`

${repoName}
Latest release version: ${latestVersion}

Please proceed to the official repository to download the latest version.

- 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'](