-
Notifications
You must be signed in to change notification settings - Fork 12
Expand file tree
/
Copy pathgenerator.py
More file actions
62 lines (52 loc) · 2.33 KB
/
Copy pathgenerator.py
File metadata and controls
62 lines (52 loc) · 2.33 KB
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
"""Wrapper class for accessing OpenAI API."""
import os
from typing import Any, ClassVar
import openai
from dotenv import load_dotenv
from .schemas import Message
class Generator:
"""Wrapper class for accessing OpenAI API."""
_default_generation_params: ClassVar[dict[str, Any]] = {
"max_tokens": 150,
"n": 1,
"stop": None,
"temperature": 0.7,
}
def __init__(self, base_url: str | None = None, model_name: str | None = None, **generation_params: Any) -> None: # noqa: ANN401
"""
Initialize the wrapper for LLM.
:param base_url: HTTP-endpoint for sending API requests to OpenAI API compatible server.
Omit this to infer OPENAI_BASE_URL from environment.
:param model_name: Name of LLM. Omit this to infer OPENAI_MODEL_NAME from environment.
:param generation_params: kwargs that will be sent with a request to the endpoint.
Omit this to use AutoIntent's default parameters.
"""
if not base_url:
load_dotenv()
base_url = os.environ["OPENAI_BASE_URL"]
if not model_name:
load_dotenv()
model_name = os.environ["OPENAI_MODEL_NAME"]
self.model_name = model_name
self.client = openai.OpenAI(base_url=base_url)
self.async_client = openai.AsyncOpenAI(base_url=base_url)
self.generation_params = {
**self._default_generation_params,
**generation_params,
} # https://stackoverflow.com/a/65539348
def get_chat_completion(self, messages: list[Message]) -> str:
"""Prompt LLM and return its answer synchronously."""
response = self.client.chat.completions.create(
messages=messages, # type: ignore[arg-type]
model=self.model_name,
**self.generation_params,
)
return response.choices[0].message.content # type: ignore[return-value]
async def get_chat_completion_async(self, messages: list[Message]) -> str:
"""Prompt LLM and return its answer asynchronously."""
response = await self.async_client.chat.completions.create(
messages=messages, # type: ignore[arg-type]
model=self.model_name,
**self.generation_params,
)
return response.choices[0].message.content # type: ignore[return-value]