Spaces:
Sleeping
Sleeping
File size: 4,691 Bytes
338036b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 | from ..LLMInterface import LLMInterface
from ..LLMEnums import OpenAIEnums
from openai import OpenAI, APIConnectionError
import logging
from typing import Union, List
class OpenAIProvider(LLMInterface):
def __init__(self, api_key: str, api_url: str=None,
default_input_max_characters: int=1000,
default_generation_max_output_tokens: int=1000,
default_generation_tempreature: float=0.1):
self.api_key = api_key
self.api_url = api_url
self.default_input_max_characters = default_input_max_characters
self.default_generation_max_output_tokens = default_generation_max_output_tokens
self.default_generation_tempreature = default_generation_tempreature
self.generation_moddel_id = None
self.embedding_model_id = None
self.embedding_size = None
self.client = OpenAI(
api_key = self.api_key if self.api_key else "placeholder",
base_url = self.api_url if self.api_url and len(self.api_url) else None
)
self.enums = OpenAIEnums
self.logger = logging.getLogger(__name__)
def set_generation_model(self, model_id: str):
self.generation_moddel_id = model_id
def set_embedding_model(self, model_id: str, embedding_size: int):
self.embedding_model_id = model_id
self.embedding_size = embedding_size
def process_text(self, text: str):
return text[:self.default_input_max_characters].strip()
def generate_text(self, prompt: str, chat_history: list=[],
max_output_token: int=None,temperature: float = None):
if not self.client :
self.logger.error("client was not set!")
return None
if not self.generation_moddel_id :
self.logger.error("embedding model was not set")
return None
max_output_token = max_output_token if max_output_token else self.default_generation_max_output_tokens
temperature = temperature if temperature else self.default_generation_tempreature
chat_history.append(
self.construct_prompt(prompt=prompt, role=OpenAIEnums.USER))
try:
response = self.client.chat.completions.create(
model = self.generation_moddel_id,
messages = chat_history,
max_tokens = max_output_token,
temperature = temperature
)
except APIConnectionError as e:
self.logger.error(
f"LLM connection error: could not reach '{self.api_url}'. "
f"Make sure Ollama (or your LLM server) is running. Details: {e}"
)
return None
if not response or not response.choices or len(response.choices) == 0 or not response.choices[0]:
self.logger.error("Error while generating text with OpenAI")
return None
return response.choices[0].message.content
def embed_text(self, text: Union[str, List[str]], document_type: str =None):
if not self.client :
self.logger.error("OpenAI client was not set!")
return None
if isinstance(text, str):
text = [text]
if not self.embedding_model_id:
self.logger.error("embedding model for OpenAI was not set")
return None
try:
response = self.client.embeddings.create(
model = self.embedding_model_id,
input=text
)
except APIConnectionError as e:
self.logger.error(
f"LLM connection error: could not reach '{self.api_url}'. "
f"Make sure Ollama (or your LLM server) is running. Details: {e}"
)
return None
# response validation
if not response or not response.data or len(response.data) == 0 or not response.data[0].embedding:
self.logger.error("Error while embedding text with OpenAI")
return None
# Return all vectors so the controller can embed many chunks in one request.
return [item.embedding for item in response.data]
def construct_prompt(self, prompt: str, role: str):
return {
"role": role,
"content": prompt
}
|