From d5f67d1569fa21f768391dcab71a032b53a258e2 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=F0=9F=90=A6?= <95787+whatbirdisthat@users.noreply.github.com> Date: Thu, 12 Oct 2023 09:17:42 +1100 Subject: [PATCH] refactor for repetition --- cyberdolphin_openai_advanced.py | 66 +++++++++++++++---------------- cyberdolphin_openai_compatible.py | 56 +++++++++----------------- cyberdolphin_openai_simple.py | 25 ++++-------- openai_client.py | 22 +++++++++++ settings.py | 5 +++ 5 files changed, 85 insertions(+), 89 deletions(-) create mode 100644 openai_client.py diff --git a/cyberdolphin_openai_advanced.py b/cyberdolphin_openai_advanced.py index 24aea0d..86de3b1 100644 --- a/cyberdolphin_openai_advanced.py +++ b/cyberdolphin_openai_advanced.py @@ -1,5 +1,7 @@ import openai -from .settings import load_settings + +from .openai_client import OpenAiClient +from .settings import load_settings, api_settings class CyberdolphinOpenAIAdvanced: @@ -8,16 +10,17 @@ class CyberdolphinOpenAIAdvanced: @classmethod def INPUT_TYPES(s): the_settings = load_settings() - openai_settings = the_settings['openai'] - openai.api_base = openai_settings['api_base'] - openai.api_key = openai_settings['api_key'] - openai.organization = openai_settings['organisation'] + openai.api_base, openai.api_key, openai.organization = api_settings('openai') + openai_model_list = [m["id"] for m in openai.Model.list()['data']] example_system_prompt = the_settings['prompt_templates']['gpt-3.5-turbo']['system'] - example_user_prompt = f"{the_settings['prompt_templates']['gpt-3.5-turbo']['prefix']}{the_settings['example_user_prompt']}{the_settings['prompt_templates']['gpt-3.5-turbo']['suffix']}" + example_user_prompt = f"\ + {the_settings['prompt_templates']['gpt-3.5-turbo']['prefix']}\ + {the_settings['example_user_prompt']}\ + {the_settings['prompt_templates']['gpt-3.5-turbo']['suffix']}" return { "required": { - "model": ([m["id"] for m in openai.Model.list()['data']], { + "model": (openai_model_list, { "default": "gpt-3.5-turbo"}), "system_prompt": ('STRING', { "multiline": True, @@ -56,15 +59,11 @@ class CyberdolphinOpenAIAdvanced: if temperature < 0 or temperature > 2: errors_list.append( """Temperature should be a value between 0.0 and 2.0 - - higher values like 0.8 will make the output more random, + - openai says higher values like 0.8 will make the output more random, lower values like 0.2 make it more focused and deterministic. """) return errors_list - def api_settings(self): - openai_settings = load_settings()['openai'] - return openai_settings['api_base'], openai_settings['api_key'], openai_settings['organisation'] - def generate(self, model: str, system_prompt: str, user_prompt="", temperature: float | None = None, top_p: float | None = None): errors = self.validate_content(temperature, top_p) @@ -72,28 +71,29 @@ class CyberdolphinOpenAIAdvanced: error_report = "\n".join([e for e in errors]) raise RuntimeError(f"There were problems with the parameters:\n{error_report}") else: - system_content = system_prompt user_content = user_prompt - if top_p == 0.0: - openai.api_base, openai.api_key, openai.organization = self.api_settings() - response = openai.ChatCompletion.create( - model=model, - temperature=temperature, - messages=[ - {"role": "system", "content": system_content}, - {"role": "user", "content": user_content} - ] - ) - else: - openai.api_base, openai.api_key, openai.organization = self.api_settings() - response = openai.ChatCompletion.create( - model=model, - top_p=top_p, - messages=[ - {"role": "system", "content": system_content}, - {"role": "user", "content": user_content} - ] - ) + + if top_p == 0.0 or top_p == 1.0: + top_p = int(top_p) + + response = OpenAiClient.complete( + key="openai", + model=model, + temperature=temperature, + top_p=top_p, + system_content=system_content, + user_content=user_content) + + # openai.api_base, openai.api_key, openai.organization = api_settings() + # response = openai.ChatCompletion.create( + # model=model, + # temperature=temperature, + # top_p=top_p, + # messages=[ + # {"role": "system", "content": system_content}, + # {"role": "user", "content": user_content} + # ] + # ) return (f'{response.choices[0].message.content}',) diff --git a/cyberdolphin_openai_compatible.py b/cyberdolphin_openai_compatible.py index 50ea8c9..5983178 100644 --- a/cyberdolphin_openai_compatible.py +++ b/cyberdolphin_openai_compatible.py @@ -1,5 +1,7 @@ import openai -from .settings import load_settings + +from .openai_client import OpenAiClient +from .settings import load_settings, api_settings all_settings = load_settings() prompt_templates = all_settings['prompt_templates'] @@ -12,22 +14,19 @@ default_user_prompt = all_settings['example_user_prompt'] class CyberdolphinOpenAICompatible: the_settings = None - def api_settings(self): - openai_settings = load_settings()['openai_compatible'] - return openai_settings['api_base'], openai_settings['api_key'], openai_settings['organisation'] @classmethod def INPUT_TYPES(s): - openai.api_key = openai_settings['api_key'] - openai.organization = openai_settings['organisation'] - openai.api_base = openai_settings['api_base'] + openai.api_base, openai.api_key, openai.organization = api_settings('openai_compatible') available_models = [m["id"] for m in openai.Model.list()['data']] + available_templates = [t for t in prompt_templates] + return { "required": { - "api_base": ("STRING", { - "default": openai_settings['api_base'] - }), - "prompt_template": ([t for t in prompt_templates], { + # "api_base": ("STRING", { + # "default": openai_settings['api_base'] + # }), + "prompt_template": (available_templates, { "default": 'default' }), "model": (available_models, { @@ -69,11 +68,9 @@ class CyberdolphinOpenAICompatible: """) return errors_list - def generate(self, api_base: str, prompt_template: str, model: str, temperature: float | None = None, + def generate(self, prompt_template: str, model: str, temperature: float | None = None, top_p: float | None = None, user_prompt=""): errors = self.validate_content(temperature, top_p) - if top_p == 0.0: - top_p = None if errors: error_report = "\n".join([e for e in errors]) raise RuntimeError(f"There were problems with the parameters:\n{error_report}") @@ -83,29 +80,12 @@ class CyberdolphinOpenAICompatible: system_content = this_prompt['system'] user_content = f"{this_prompt['prefix']} {user_prompt} {this_prompt['suffix']}" - openai.api_base, openai.api_key, openai.organization = self.api_settings() - - # openai_settings = load_settings()['openai_compatible'] - # openai.api_base = openai_settings['api_base'] - # openai.api_key = openai_settings['api_key'] - # openai.organization = openai_settings['organisation'] - if top_p is None: - response = openai.ChatCompletion.create( - model=model, - temperature=temperature, - messages=[ - {"role": "system", "content": system_content}, - {"role": "user", "content": user_content} - ] - ) - else: - response = openai.ChatCompletion.create( - model=model, - top_p=top_p, - messages=[ - {"role": "system", "content": system_content}, - {"role": "user", "content": user_content} - ] - ) + response = OpenAiClient.complete( + key="openai_compatible", + model=model, + temperature=temperature, + top_p=top_p, + system_content=system_content, + user_content=user_content) return (f'{response.choices[0].message.content}',) diff --git a/cyberdolphin_openai_simple.py b/cyberdolphin_openai_simple.py index e212bb1..2a44880 100644 --- a/cyberdolphin_openai_simple.py +++ b/cyberdolphin_openai_simple.py @@ -1,4 +1,4 @@ -import openai +from .openai_client import OpenAiClient from .settings import load_settings @@ -27,28 +27,17 @@ class CyberdolphinOpenAISimple: # OUTPUT_NODE = False CATEGORY = "🐬 CyberDolphin" - def api_settings(self): - openai_settings = load_settings()['openai'] - return openai_settings['api_base'], openai_settings['api_key'], openai_settings['organisation'] def generate(self, user_prompt="", temperature: float = 1.0): prompts = load_settings()['prompt_templates'] - system_prompt = prompts['gpt-3.5-turbo']['system'] + system_content = prompts['gpt-3.5-turbo']['system'] user_content = f"{prompts['gpt-3.5-turbo']['prefix']} {user_prompt} {prompts['gpt-3.5-turbo']['suffix']}" - - # openai_settings = load_settings()['openai'] - # openai.api_base = openai_settings['api_base'] - # openai.api_key = openai_settings['api_key'] - # openai.organization = openai_settings['organisation'] - openai.api_base, openai.api_key, openai.organization = self.api_settings() - - response = openai.ChatCompletion.create( + response = OpenAiClient.complete( + key="openai_compatible", model=load_settings()['openai']['default_model'], temperature=temperature, - messages=[ - {"role": "system", "content": system_prompt}, - {"role": "user", "content": user_content} - ] - ) + top_p=1.0, + system_content=system_content, + user_content=user_content) return (f'{response.choices[0].message.content}',) diff --git a/openai_client.py b/openai_client.py new file mode 100644 index 0000000..2321a82 --- /dev/null +++ b/openai_client.py @@ -0,0 +1,22 @@ +import openai + +from custom_nodes.cyberdolphin.settings import api_settings + + +class OpenAiClient: + @staticmethod + def complete(key, model, temperature, top_p, system_content, user_content): + if top_p == 0.0: + top_p = 0.001 + + openai.api_base, openai.api_key, openai.organization = api_settings(key) + response = openai.ChatCompletion.create( + model=model, + temperature=temperature, + top_p=top_p, + messages=[ + {"role": "system", "content": system_content}, + {"role": "user", "content": user_content} + ] + ) + return response diff --git a/settings.py b/settings.py index a234dc2..56f99b4 100644 --- a/settings.py +++ b/settings.py @@ -38,3 +38,8 @@ def load_settings(): the_yaml = yaml.safe_load(settings) # print(f'LOADED: {the_yaml["cyberdolphin"]}') return the_yaml['cyberdolphin'] + + +def api_settings(section: str = "openai"): + openai_settings = load_settings()[section] + return openai_settings['api_base'], openai_settings['api_key'], openai_settings['organisation']