From 1e9ffc5ffcafe1bc961841a13ef3d5759e49ed93 Mon Sep 17 00:00:00 2001 From: yolain Date: Mon, 3 Jun 2024 14:44:28 +0800 Subject: [PATCH] add:auto translate chinese prompt to english --- README.en.md | 2 +- README.md | 2 + install.bat | 16 +++ prestartup_script.py | 1 + py/api.py | 10 ++ py/easyNodes.py | 46 +++++++- py/libs/conditioning.py | 6 + py/libs/translate.py | 238 ++++++++++++++++++++++++++++++++++++++++ requirements.txt | 3 +- 9 files changed, 316 insertions(+), 8 deletions(-) create mode 100644 install.bat create mode 100644 py/libs/translate.py diff --git a/README.en.md b/README.en.md index 75f152a..69e2aa1 100644 --- a/README.en.md +++ b/README.en.md @@ -30,7 +30,7 @@ - Background removal nodes for the RMBG-1.4 model supporting BriaAI, [BriaAI Guide](https://huggingface.co/briaai/RMBG-1.4) - Forcibly cleared the memory usage of the comfy UI model are supported - Stable Diffusion 3 multi-account API nodes are supported -- + ## Changelog **v1.1.8** diff --git a/README.md b/README.md index f75aec7..9fd23a6 100644 --- a/README.md +++ b/README.md @@ -36,11 +36,13 @@ - 支持 强制清理comfyUI模型显存占用 - 支持Stable Diffusion 3 多账号API节点 - 支持IC-Light的应用 [示例参考](https://github.com/yolain/ComfyUI-Yolain-Workflows?tab=readme-ov-file#2-5-ic-light) | [代码整合来源](https://github.com/huchenlei/ComfyUI-IC-Light) | [技术参考](https://github.com/lllyasviel/IC-Light) +- 中文提示词自动识别,使用[opus-mt-zh-en模型](https://huggingface.co/Helsinki-NLP/opus-mt-zh-en) ## 更新日志 **v1.1.8** +- 增加中文提示词自动翻译,使用[opus-mt-zh-en模型](https://huggingface.co/Helsinki-NLP/opus-mt-zh-en), 默认已对wildcard、lora正则处理, 其他需要保留的中文,可使用`@你的提示词@`包裹 (若依赖安装完成后报错, 请重启),测算大约会占0.3GB显存 - 增加 `easy controlnetStack` - controlnet堆 - 增加 `easy applyBrushNet` - [示例参考](https://github.com/yolain/ComfyUI-Yolain-Workflows/blob/main/workflows/2_advanced/2-4inpainting/2-4brushnet_1.1.8.json) - 增加 `easy applyPowerPaint` - [示例参考](https://github.com/yolain/ComfyUI-Yolain-Workflows/blob/main/workflows/2_advanced/2-4inpainting/2-4powerpaint_outpaint_1.1.8.json) diff --git a/install.bat b/install.bat new file mode 100644 index 0000000..b9027d4 --- /dev/null +++ b/install.bat @@ -0,0 +1,16 @@ +@echo off + +set "requirements_txt=%~dp0\requirements.txt" +set "python_exec=..\..\..\python_embeded\python.exe" + +echo Installing EasyUse Requirements... + +if exist "%python_exec%" ( + echo Installing with ComfyUI Portable + "%python_exec%" -s -m pip install -r "%requirements_txt%" +) else ( + echo Installing with system Python + pip install -r "%requirements_txt%" +) + +pause \ No newline at end of file diff --git a/prestartup_script.py b/prestartup_script.py index 1a82c0f..c3a59be 100644 --- a/prestartup_script.py +++ b/prestartup_script.py @@ -28,6 +28,7 @@ add_folder_path_and_extensions("ipadapter", [os.path.join(model_path, "ipadapter add_folder_path_and_extensions("dynamicrafter_models", [os.path.join(model_path, "dynamicrafter_models")], folder_paths.supported_pt_extensions) add_folder_path_and_extensions("mediapipe", [os.path.join(model_path, "mediapipe")], set(['.tflite','.pth'])) add_folder_path_and_extensions("inpaint", [os.path.join(model_path, "inpaint")], folder_paths.supported_pt_extensions) +add_folder_path_and_extensions("prompt_generator", [os.path.join(model_path, "prompt_generator")], folder_paths.supported_pt_extensions) add_folder_path_and_extensions("checkpoints_thumb", [os.path.join(model_path, "checkpoints")], image_suffixs) add_folder_path_and_extensions("loras_thumb", [os.path.join(model_path, "loras")], image_suffixs) \ No newline at end of file diff --git a/py/api.py b/py/api.py index 2931401..b32e277 100644 --- a/py/api.py +++ b/py/api.py @@ -11,6 +11,7 @@ from .logic import ConvertAnything from .libs.model import easyModelManager from .libs.utils import getMetadata, cleanGPUUsedForce, get_local_filepath from .libs.cache import remove_cache +from .libs.translate import has_chinese, zh_to_en try: import aiohttp @@ -30,6 +31,15 @@ def cleanGPU(request): return web.Response(status=500) pass +@PromptServer.instance.routes.post("/easyuse/translate") +async def translate(request): + post = await request.post() + text = post.get("text") + if has_chinese(text): + return web.json_response({"text": zh_to_en([text])[0]}) + else: + return web.json_response({"text": text}) + @PromptServer.instance.routes.get("/easyuse/reboot") def reboot(request): try: diff --git a/py/easyNodes.py b/py/easyNodes.py index 7a28921..ff90628 100644 --- a/py/easyNodes.py +++ b/py/easyNodes.py @@ -31,6 +31,7 @@ from .libs.xyplot import easyXYPlot from .libs.controlnet import easyControlnet from .libs.conditioning import prompt_to_cond, set_cond from .libs.easing import EasingBase +from .libs.translate import has_chinese, zh_to_en from .libs import cache as backend_cache sampler = easySampler() @@ -57,6 +58,8 @@ class positivePrompt: @staticmethod def main(positive): + if has_chinese(positive): + return zh_to_en([positive]) return positive, # 通配符提示词 @@ -86,8 +89,13 @@ class wildcardsPrompt: CATEGORY = "EasyUse/Prompt" - @staticmethod - def main(*args, **kwargs): + def translate(self, text): + if has_chinese(text): + return zh_to_en([text])[0] + else: + return text + + def main(self, *args, **kwargs): prompt = kwargs["prompt"] if "prompt" in kwargs else None seed = kwargs["seed"] @@ -98,10 +106,15 @@ class wildcardsPrompt: text = kwargs['text'] if "multiline_mode" in kwargs and kwargs["multiline_mode"]: populated_text = [] + _text = [] text = text.split("\n") for t in text: + t = self.translate(t) + _text.append(t) populated_text.append(process(t, seed)) + text = _text else: + text = self.translate(text) populated_text = [process(text, seed)] text = [text] return {"ui": {"value": [seed]}, "result": (text, populated_text)} @@ -126,7 +139,10 @@ class negativePrompt: @staticmethod def main(negative): - return negative, + if has_chinese(negative): + return zh_to_en([negative]) + else: + return negative, # 风格提示词选择器 class stylesPromptSelector: @@ -262,6 +278,8 @@ class prompt: CATEGORY = "EasyUse/Prompt" def doit(self, prompt, main, lighting): + if has_chinese(prompt): + prompt = zh_to_en([prompt])[0] if lighting != 'none' and main != 'none': prompt = main + ',' + lighting + ',' + prompt elif lighting != 'none' and main == 'none': @@ -305,6 +323,8 @@ class promptList: # Only process string input ports. if isinstance(v, str) and v != '': + if has_chinese(v): + v = zh_to_en([v])[0] prompts.append(v) return (prompts, prompts) @@ -332,6 +352,7 @@ class promptLine: def generate_strings(self, prompt, start_index, max_rows, workflow_prompt=None, my_unique_id=None): lines = prompt.split('\n') + lines = [zh_to_en([v])[0] if has_chinese(v) else v for v in lines] start_index = max(0, min(start_index, len(lines) - 1)) @@ -339,7 +360,6 @@ class promptLine: rows = lines[start_index:end_index] - return (rows, rows) class promptConcat: @@ -906,11 +926,11 @@ class fullLoader: "empty_latent_width": ("INT", {"default": 512, "min": 64, "max": MAX_RESOLUTION, "step": 8}), "empty_latent_height": ("INT", {"default": 512, "min": 64, "max": MAX_RESOLUTION, "step": 8}), - "positive": ("STRING", {"default":"", "placeholder": "Positive", "multiline": True}), + "positive": ("STRING", {"default": "", "placeholder": "Positive", "multiline": True}), "positive_token_normalization": (["none", "mean", "length", "length+mean"],), "positive_weight_interpretation": (["comfy", "A1111", "comfy++", "compel", "fixed attention"],), - "negative": ("STRING", {"default":"", "placeholder": "Negative", "multiline": True}), + "negative": ("STRING", {"default": "", "placeholder": "Negative", "multiline": True}), "negative_token_normalization": (["none", "mean", "length", "length+mean"],), "negative_weight_interpretation": (["comfy", "A1111", "comfy++", "compel", "fixed attention"],), @@ -1197,6 +1217,9 @@ class cascadeLoader: log_node_warn("正在处理提示词...") positive_seed = find_wildcards_seed(my_unique_id, positive, prompt) + # Translate cn to en + if has_chinese(positive): + positive = zh_to_en([positive])[0] model_c, clip, positive, positive_decode, show_positive_prompt, pipe_lora_stack = process_with_loras(positive, model_c, clip, "positive", @@ -1206,6 +1229,9 @@ class cascadeLoader: easyCache) positive_wildcard_prompt = positive_decode if show_positive_prompt or is_positive_linked_styles_selector else "" negative_seed = find_wildcards_seed(my_unique_id, negative, prompt) + # Translate cn to en + if has_chinese(negative): + negative = zh_to_en([negative])[0] model_c, clip, negative, negative_decode, show_negative_prompt, pipe_lora_stack = process_with_loras(negative, model_c, clip, "negative", @@ -1572,11 +1598,15 @@ class svdLoader: if clip_name == 'None': raise Exception("You need choose a open_clip model when positive is not empty") clip = easyCache.load_clip(clip_name) + if has_chinese(optional_positive): + optional_positive = zh_to_en([optional_positive])[0] positive_embeddings_final, = CLIPTextEncode().encode(clip, optional_positive) positive, = ConditioningConcat().concat(positive, positive_embeddings_final) if optional_negative is not None and optional_negative != '': if clip_name == 'None': raise Exception("You need choose a open_clip model when negative is not empty") + if has_chinese(optional_negative): + optional_positive = zh_to_en([optional_negative])[0] negative_embeddings_final, = CLIPTextEncode().encode(clip, optional_negative) negative, = ConditioningConcat().concat(negative, negative_embeddings_final) @@ -1741,8 +1771,12 @@ class dynamiCrafterLoader(DynamiCrafter): clipped.clip_layer(clip_skip) if positive is not None and positive != '': + if has_chinese(positive): + positive = zh_to_en([positive])[0] positive_embeddings_final, = CLIPTextEncode().encode(clipped, positive) if negative is not None and negative != '': + if has_chinese(negative): + negative = zh_to_en([negative])[0] negative_embeddings_final, = CLIPTextEncode().encode(clipped, negative) image = easySampler.pil2tensor(Image.new('RGB', (1, 1), (0, 0, 0))) diff --git a/py/libs/conditioning.py b/py/libs/conditioning.py index b5fba3f..38a5b8f 100644 --- a/py/libs/conditioning.py +++ b/py/libs/conditioning.py @@ -1,5 +1,6 @@ from .utils import find_wildcards_seed, find_nearest_steps, is_linked_styles_selector from .log import log_node_warn +from .translate import zh_to_en, has_chinese from .wildcards import process_with_loras from .adv_encode import advanced_encode @@ -9,6 +10,11 @@ def prompt_to_cond(type, model, clip, clip_skip, lora_stack, text, prompt_token_ styles_selector = is_linked_styles_selector(prompt, my_unique_id, type) title = "正面提示词" if type == 'positive' else "负面提示词" log_node_warn("正在进行" + title + "...") + + # Translate cn to en + if has_chinese(text): + text = zh_to_en([text])[0] + positive_seed = find_wildcards_seed(my_unique_id, text, prompt) model, clip, text, cond_decode, show_prompt, pipe_lora_stack = process_with_loras( text, model, clip, type, positive_seed, can_load_lora, lora_stack, easyCache) diff --git a/py/libs/translate.py b/py/libs/translate.py new file mode 100644 index 0000000..4f867d5 --- /dev/null +++ b/py/libs/translate.py @@ -0,0 +1,238 @@ +import re +import os +import folder_paths + +import comfy.utils +import torch +from transformers import AutoModelForSeq2SeqLM, AutoTokenizer + +from .utils import install_package +try: + from lark import Lark, Transformer, v_args +except: + print('install lark-parser...') + install_package('lark-parser') + from lark import Lark, Transformer, v_args + +model_path = os.path.join(folder_paths.models_dir, 'prompt_generator') +zh_en_model_path = os.path.join(model_path, 'opus-mt-zh-en') +zh_en_model, zh_en_tokenizer = None, None + +def correct_prompt_syntax(prompt=""): + # print("input prompt",prompt) + corrected_elements = [] + # 处理成统一的英文标点 + prompt = prompt.replace('(', '(').replace(')', ')').replace(',', ',').replace(';', ',').replace('。', '.').replace(':',':') + # 删除多余的空格 + prompt = re.sub(r'\s+', ' ', prompt).strip() + prompt = prompt.replace("< ","<").replace(" >",">").replace("( ","(").replace(" )",")").replace("[ ","[").replace(' ]',']') + + # 分词 + prompt_elements = prompt.split(',') + + def balance_brackets(element, open_bracket, close_bracket): + open_brackets_count = element.count(open_bracket) + close_brackets_count = element.count(close_bracket) + return element + close_bracket * (open_brackets_count - close_brackets_count) + + for element in prompt_elements: + element = element.strip() + + # 处理空元素 + if not element: + continue + + # 检查并处理圆括号、方括号、尖括号 + if element[0] in '([': + corrected_element = balance_brackets(element, '(', ')') if element[0] == '(' else balance_brackets(element, '[', ']') + elif element[0] == '<': + corrected_element = balance_brackets(element, '<', '>') + else: + # 删除开头的右括号或右方括号 + corrected_element = element.lstrip(')]') + + corrected_elements.append(corrected_element) + + # 重组修正后的prompt + return ','.join(corrected_elements) + +def detect_language(input_str): + # 统计中文和英文字符的数量 + count_cn = count_en = 0 + for char in input_str: + if '\u4e00' <= char <= '\u9fff': + count_cn += 1 + elif char.isalpha(): + count_en += 1 + + # 根据统计的字符数量判断主要语言 + if count_cn > count_en: + return "cn" + elif count_en > count_cn: + return "en" + else: + return "unknow" + +def has_chinese(text): + has_cn = False + _text = text + _text = re.sub(r'<.*?>', '', _text) + _text = re.sub(r'__.*?__', '', _text) + _text = re.sub(r'embedding:.*?(\d+)?', '', _text) + for char in _text: + if '\u4e00' <= char <= '\u9fff': + has_cn = True + break + elif char.isalpha(): + continue + return has_cn + +def translate(text): + global zh_en_model_path, zh_en_model, zh_en_tokenizer + + if not os.path.exists(zh_en_model_path): + zh_en_model_path = 'Helsinki-NLP/opus-mt-zh-en' + + if zh_en_model is None: + + zh_en_model = AutoModelForSeq2SeqLM.from_pretrained(zh_en_model_path).eval() + zh_en_tokenizer = AutoTokenizer.from_pretrained(zh_en_model_path, padding=True, truncation=True) + + zh_en_model.to("cuda" if torch.cuda.is_available() else "cpu") + with torch.no_grad(): + encoded = zh_en_tokenizer([text], return_tensors="pt") + encoded.to(zh_en_model.device) + sequences = zh_en_model.generate(**encoded) + return zh_en_tokenizer.batch_decode(sequences, skip_special_tokens=True)[0] + +@v_args(inline=True) # Decorator to flatten the tree directly into the function arguments +class ChinesePromptTranslate(Transformer): + + def sentence(self, *args): + return ", ".join(args) + + def phrase(self, *args): + return "".join(args) + + def emphasis(self, *args): + # Reconstruct the emphasis with translated content + return "(" + "".join(args) + ")" + + def weak_emphasis(self, *args): + print('weak_emphasis:', args) + return "[" + "".join(args) + "]" + + def embedding(self, *args): + print('prompt embedding', args[0]) + if len(args) == 1: + # print('prompt embedding',str(args[0])) + # 只传递了一个参数,意味着只有embedding名称没有数字 + embedding_name = str(args[0]) + return f"embedding:{embedding_name}" + elif len(args) > 1: + embedding_name, *numbers = args + + if len(numbers) == 2: + return f"embedding:{embedding_name}:{numbers[0]}:{numbers[1]}" + elif len(numbers) == 1: + return f"embedding:{embedding_name}:{numbers[0]}" + else: + return f"embedding:{embedding_name}" + + def lora(self, *args): + if len(args) == 1: + return f"" + elif len(args) > 1: + # print('lora', args) + _, loar_name, *numbers = args + loar_name = str(loar_name).strip() + if len(numbers) == 2: + return f"" + elif len(numbers) == 1: + return f"" + else: + return f"" + + def weight(self, word, number): + translated_word = translate(str(word)).rstrip('.') + return f"({translated_word}:{str(number).strip()})" + + def schedule(self, *args): + print('prompt schedule', args) + data = [str(arg).strip() for arg in args] + + return f"[{':'.join(data)}]" + + def word(self, word): + # Translate each word using the dictionary + if re.search(r'__.*?__', str(word)): + return str(word).rstrip('.') + elif re.search(r'@.*?@', str(word)): + return str(word).replace('@', '').rstrip('.') + elif detect_language(str(word)) == "cn": + return translate(str(word)).rstrip('.') + else: + return str(word).rstrip('.') + + +#定义Prompt文法 +grammar = """ +start: sentence +sentence: phrase ("," phrase)* +phrase: emphasis | weight | word | lora | embedding | schedule +emphasis: "(" sentence ")" -> emphasis + | "[" sentence "]" -> weak_emphasis +weight: "(" word ":" NUMBER ")" +schedule: "[" word ":" word ":" NUMBER "]" +lora: "<" WORD ":" WORD (":" NUMBER)? (":" NUMBER)? ">" +embedding: "embedding" ":" WORD (":" NUMBER)? (":" NUMBER)? +word: WORD + +NUMBER: /\s*-?\d+(\.\d+)?\s*/ +WORD: /[^,:\(\)\[\]<>]+/ +""" +def zh_to_en(text): + global zh_en_model_path, zh_en_model, zh_en_tokenizer + # 进度条 + pbar = comfy.utils.ProgressBar(len(text) + 1) + texts = [correct_prompt_syntax(t) for t in text] + + install_package('sentencepiece', '0.2.0') + + if not os.path.exists(zh_en_model_path): + zh_en_model_path = 'Helsinki-NLP/opus-mt-zh-en' + + if zh_en_model is None: + zh_en_model = AutoModelForSeq2SeqLM.from_pretrained(zh_en_model_path).eval() + zh_en_tokenizer = AutoTokenizer.from_pretrained(zh_en_model_path, padding=True, truncation=True) + + zh_en_model.to("cuda" if torch.cuda.is_available() else "cpu") + + prompt_result = [] + + en_texts = [] + + for t in texts: + if t: + # translated_text = translated_word = translate(zh_en_tokenizer,zh_en_model,str(t)) + parser = Lark(grammar, start="start", parser="lalr", transformer=ChinesePromptTranslate()) + # print('t',t) + result = parser.parse(t).children + # print('en_result',result) + # en_text=translate(zh_en_tokenizer,zh_en_model,text_without_syntax) + en_texts.append(result[0]) + + zh_en_model.to('cpu') + # print("test en_text", en_texts) + # en_text.to("cuda" if torch.cuda.is_available() else "cpu") + + pbar.update(1) + for t in en_texts: + prompt_result.append(t) + pbar.update(1) + + # print('prompt_result', prompt_result, ) + if len(prompt_result) == 0: + prompt_result = [""] + + return prompt_result \ No newline at end of file diff --git a/requirements.txt b/requirements.txt index efc5863..1967577 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,4 +1,5 @@ diffusers>=0.25.0 clip_interrogator>=0.6.0 +sentencepiece==0.2.0 +lark-parser onnxruntime -aiohttp \ No newline at end of file