From 35a7f6e157391f6d3886985fad5279b9af12754d Mon Sep 17 00:00:00 2001 From: chflame Date: Thu, 29 Feb 2024 19:10:30 +0800 Subject: [PATCH] update requirements.txt --- .gitignore | 3 +- py/imagefunc.py | 35 ++++++++++++++++ py/prompt_inference.py | 88 +++++++++++++++++++++++++++++++++++++++ requirements.txt | 4 +- resource/inference.prompt | 5 +-- 5 files changed, 129 insertions(+), 6 deletions(-) create mode 100644 py/prompt_inference.py diff --git a/.gitignore b/.gitignore index 2887b2a..868a782 100644 --- a/.gitignore +++ b/.gitignore @@ -2,4 +2,5 @@ _test_*.* __pycache__ .venv .idea -*.pth \ No newline at end of file +*.pth +*.ini diff --git a/py/imagefunc.py b/py/imagefunc.py index 58585eb..7650228 100644 --- a/py/imagefunc.py +++ b/py/imagefunc.py @@ -1086,6 +1086,13 @@ def has_letters(string:str) -> bool: else: return False + +def replace_case(old:str, new:str, text:str) -> str: + index = text.lower().find(old.lower()) + if index == -1: + return text + return replace_case(old, new, text[:index] + new + text[index + len(old):]) + def random_numbers(total:int, random_range:int, seed:int=0, sum_of_numbers:int=0) -> list: random.seed(seed) numbers = [random.randint(-random_range//2, random_range//2) for _ in range(total - 1)] @@ -1170,6 +1177,34 @@ chop_mode = ['normal', 'multply', 'screen', 'add', 'subtract', 'difference', 'da default_lut_dir = os.path.join(os.path.dirname(os.path.dirname(os.path.normpath(__file__))), 'lut') default_font_dir = os.path.join(os.path.dirname(os.path.dirname(os.path.normpath(__file__))), 'font') resource_dir_ini_file = os.path.join(os.path.dirname(os.path.dirname(os.path.normpath(__file__))), "resource_dir.ini") +api_key_ini_file = os.path.join(os.path.dirname(os.path.dirname(os.path.normpath(__file__))), "api_key.ini") +inference_prompt_file = os.path.join(os.path.dirname(os.path.dirname(os.path.normpath(__file__))), "resource", "inference.prompt") + +def load_inference_prompt() -> str: + ret_value = '' + try: + with open(inference_prompt_file, 'r') as f: + ret_value = f.readlines() + except Exception as e: + log(f'Warning: {inference_prompt_file} ' + repr(e) + f", check it to be correct. ", message_type='warning') + return ''.join(ret_value) + +def get_api_key(api_name:str) -> str: + ret_value = '' + try: + with open(api_key_ini_file, 'r') as f: + ini = f.readlines() + for line in ini: + if line.startswith(f'{api_name}='): + ret_value = line[line.find('=') + 1:].rstrip().lstrip() + break + except Exception as e: + log(f'Warning: {api_key_ini_file} ' + repr(e) + f", check it to be correct. ", message_type='warning') + remove_char = ['"', "'", '“', '”', '‘', '’'] + for i in remove_char: + if i in ret_value: + ret_value = ret_value.replace(i, '') + return ret_value try: with open(resource_dir_ini_file, 'r') as f: diff --git a/py/prompt_inference.py b/py/prompt_inference.py new file mode 100644 index 0000000..b0d9dd0 --- /dev/null +++ b/py/prompt_inference.py @@ -0,0 +1,88 @@ +from .imagefunc import * +import google.generativeai as genai + +NODE_NAME = 'PromptInference' + +class PromptInference: + + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(self): + api_list = ['gemini-pro-vision'] + + return { + "required": { + "image": ("IMAGE", ), + "api": (api_list,), + "use_default_request": ("BOOLEAN", {"default": True}), + "custom_request": ("STRING", {"default": ""}), + "key_word": ("STRING", {"default": ""}), + "exclude_word": ("STRING", {"default": ""}), + "token_limit": ("INT", {"default": 80, "min": 2, "max": 1024, "step": 1}), + }, + "optional": { + } + } + + RETURN_TYPES = ("STRING",) + RETURN_NAMES = ("text",) + FUNCTION = 'prompt_inference' + CATEGORY = '😺dzNodes/LayerUtility' + OUTPUT_NODE = True + + def prompt_inference(self, image, api, use_default_request, custom_request, key_word, exclude_word, token_limit): + custom_request = custom_request.strip() + key_word = key_word.strip() + exclude_word = exclude_word.strip() + key_word_list = [] + if len(key_word) > 0: + key_word_list = list(re.split(r'[,,\s*]', key_word)) + key_word_list = [x for x in key_word_list if x != ''] # 去除空字符 + key_word_list = [f"({i})" for i in key_word_list] + + token_prompt = f"as needed keep it under {token_limit} tokens" + _image = tensor2pil(image).convert('RGB') + prompt = "" + ret_text = "" + + if use_default_request: + prompt = f"{load_inference_prompt()}" + if len(custom_request) > 0: + prompt = f"{prompt}{custom_request}, " + prompt = f"{prompt}{token_prompt}." + + if api == 'gemini-pro-vision': + model = genai.GenerativeModel(api) + genai.configure(api_key=get_api_key('google_api_key'), transport='rest') + response = model.generate_content([prompt, _image]) + ret_text = response.text + if len(exclude_word) > 0: + exclude_word_list = list(re.split(r'[,,\s*]', exclude_word)) + exclude_word_list = [x for x in exclude_word_list if x != ''] # 去除空字符 + print(f"exclude_words={exclude_word_list}") + if len(key_word_list) > 0: + ret_text = replace_case(exclude_word_list[0], key_word_list[0], ret_text) + for i in exclude_word_list: + ret_text = replace_case(i, '', ret_text) + print(f"after exclude_word ret_text = {ret_text}") + refine_model = genai.GenerativeModel('gemini-pro') + response = refine_model.generate_content(f"Please correct the grammar errors in the following text:{ret_text}") + ret_text = response.text + + if len(key_word) > 0: + ret_text = f"A photo of {', '.join(key_word_list)}, {replace_case('A photo of ', '', ret_text)}, " + + log(f"{NODE_NAME} request to gemini-pro-vision, prompt=\n\033[1;36m{prompt}\033[m\nresponse=\n\033[1;36m{ret_text}\033[m") + + log(f"{NODE_NAME} Processed.", message_type='finish') + return (ret_text,) + +NODE_CLASS_MAPPINGS = { + "LayerUtility: PromptInference": PromptInference +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "LayerUtility: PromptInference": "LayerUtility: PromptInference" +} \ No newline at end of file diff --git a/requirements.txt b/requirements.txt index be65d72..a9c5d0d 100644 --- a/requirements.txt +++ b/requirements.txt @@ -15,4 +15,6 @@ wget mediapipe loguru typer_config -fastapi \ No newline at end of file +fastapi +rich +google-generativeai \ No newline at end of file diff --git a/resource/inference.prompt b/resource/inference.prompt index 1c51494..fea56ed 100644 --- a/resource/inference.prompt +++ b/resource/inference.prompt @@ -1,4 +1 @@ -You are creating a prompt for Stable Diffusion to generate an image. -First step: describe this image, then put description into text. -Second step: generate a text prompt for based on first step. -Only respond with the prompt itself, but embellish it as needed keep it under 80 tokens. \ No newline at end of file +You are creating a prompt for Stable Diffusion to generate an image. First step: describe this image, then put description into text. Second step: generate a text prompt for based on first step. Only respond with the prompt itself, but embellish it. \ No newline at end of file