mirror of
https://github.com/NicolasBohn/NexQuant.git
synced 2026-07-27 23:47:46 +00:00
6cc1e5da3c
* init commit * limit problem numbers * ensemble lower case * add runtime and spec to coder * submission check notice * sub EDA in sample execution * avoid lightgbm * add time limit to scenario * rephrase the submission check * give positive feedback when facing warning in check * ENABLE FEEDBACK * fix feedback bug --------- Co-authored-by: Xu Yang <peteryang@vip.qq.com>
470 lines
19 KiB
Python
470 lines
19 KiB
Python
from __future__ import annotations
|
|
|
|
import json
|
|
import re
|
|
import sqlite3
|
|
import time
|
|
import uuid
|
|
from abc import ABC, abstractmethod
|
|
from copy import deepcopy
|
|
from pathlib import Path
|
|
from typing import Any, Optional, cast
|
|
|
|
from pydantic import TypeAdapter
|
|
|
|
from rdagent.core.utils import LLM_CACHE_SEED_GEN, SingletonBaseClass
|
|
from rdagent.log import LogColors
|
|
from rdagent.log import rdagent_logger as logger
|
|
from rdagent.oai.llm_conf import LLM_SETTINGS
|
|
from rdagent.utils import md5_hash
|
|
|
|
|
|
class SQliteLazyCache(SingletonBaseClass):
|
|
def __init__(self, cache_location: str) -> None:
|
|
super().__init__()
|
|
self.cache_location = cache_location
|
|
db_file_exist = Path(cache_location).exists()
|
|
# TODO: sqlite3 does not support multiprocessing.
|
|
self.conn = sqlite3.connect(cache_location, timeout=20)
|
|
self.c = self.conn.cursor()
|
|
if not db_file_exist:
|
|
self.c.execute(
|
|
"""
|
|
CREATE TABLE chat_cache (
|
|
md5_key TEXT PRIMARY KEY,
|
|
chat TEXT
|
|
)
|
|
""",
|
|
)
|
|
self.c.execute(
|
|
"""
|
|
CREATE TABLE embedding_cache (
|
|
md5_key TEXT PRIMARY KEY,
|
|
embedding TEXT
|
|
)
|
|
""",
|
|
)
|
|
self.c.execute(
|
|
"""
|
|
CREATE TABLE message_cache (
|
|
conversation_id TEXT PRIMARY KEY,
|
|
message TEXT
|
|
)
|
|
""",
|
|
)
|
|
self.conn.commit()
|
|
|
|
def chat_get(self, key: str) -> str | None:
|
|
md5_key = md5_hash(key)
|
|
self.c.execute("SELECT chat FROM chat_cache WHERE md5_key=?", (md5_key,))
|
|
result = self.c.fetchone()
|
|
return None if result is None else result[0]
|
|
|
|
def embedding_get(self, key: str) -> list | dict | str | None:
|
|
md5_key = md5_hash(key)
|
|
self.c.execute("SELECT embedding FROM embedding_cache WHERE md5_key=?", (md5_key,))
|
|
result = self.c.fetchone()
|
|
return None if result is None else json.loads(result[0])
|
|
|
|
def chat_set(self, key: str, value: str) -> None:
|
|
md5_key = md5_hash(key)
|
|
self.c.execute(
|
|
"INSERT OR REPLACE INTO chat_cache (md5_key, chat) VALUES (?, ?)",
|
|
(md5_key, value),
|
|
)
|
|
self.conn.commit()
|
|
return None
|
|
|
|
def embedding_set(self, content_to_embedding_dict: dict) -> None:
|
|
for key, value in content_to_embedding_dict.items():
|
|
md5_key = md5_hash(key)
|
|
self.c.execute(
|
|
"INSERT OR REPLACE INTO embedding_cache (md5_key, embedding) VALUES (?, ?)",
|
|
(md5_key, json.dumps(value)),
|
|
)
|
|
self.conn.commit()
|
|
|
|
def message_get(self, conversation_id: str) -> list[dict[str, Any]]:
|
|
self.c.execute("SELECT message FROM message_cache WHERE conversation_id=?", (conversation_id,))
|
|
result = self.c.fetchone()
|
|
return [] if result is None else cast(list[dict[str, Any]], json.loads(result[0]))
|
|
|
|
def message_set(self, conversation_id: str, message_value: list[dict[str, Any]]) -> None:
|
|
self.c.execute(
|
|
"INSERT OR REPLACE INTO message_cache (conversation_id, message) VALUES (?, ?)",
|
|
(conversation_id, json.dumps(message_value)),
|
|
)
|
|
self.conn.commit()
|
|
return None
|
|
|
|
|
|
class SessionChatHistoryCache(SingletonBaseClass):
|
|
def __init__(self) -> None:
|
|
"""load all history conversation json file from self.session_cache_location"""
|
|
self.cache = SQliteLazyCache(cache_location=LLM_SETTINGS.prompt_cache_path)
|
|
|
|
def message_get(self, conversation_id: str) -> list[dict[str, Any]]:
|
|
return self.cache.message_get(conversation_id)
|
|
|
|
def message_set(self, conversation_id: str, message_value: list[dict[str, Any]]) -> None:
|
|
self.cache.message_set(conversation_id, message_value)
|
|
|
|
|
|
class ChatSession:
|
|
def __init__(self, api_backend: Any, conversation_id: str | None = None, system_prompt: str | None = None) -> None:
|
|
self.conversation_id = str(uuid.uuid4()) if conversation_id is None else conversation_id
|
|
self.system_prompt = system_prompt if system_prompt is not None else LLM_SETTINGS.default_system_prompt
|
|
self.api_backend = api_backend
|
|
|
|
def build_chat_completion_message(self, user_prompt: str) -> list[dict[str, Any]]:
|
|
history_message = SessionChatHistoryCache().message_get(self.conversation_id)
|
|
messages = history_message
|
|
if not messages:
|
|
messages.append({"role": LLM_SETTINGS.system_prompt_role, "content": self.system_prompt})
|
|
messages.append(
|
|
{
|
|
"role": "user",
|
|
"content": user_prompt,
|
|
},
|
|
)
|
|
return messages
|
|
|
|
def build_chat_completion_message_and_calculate_token(self, user_prompt: str) -> Any:
|
|
messages = self.build_chat_completion_message(user_prompt)
|
|
return self.api_backend._calculate_token_from_messages(messages)
|
|
|
|
def build_chat_completion(self, user_prompt: str, *args, **kwargs) -> str: # type: ignore[no-untyped-def]
|
|
"""
|
|
this function is to build the session messages
|
|
user prompt should always be provided
|
|
"""
|
|
messages = self.build_chat_completion_message(user_prompt)
|
|
|
|
with logger.tag(f"session_{self.conversation_id}"):
|
|
response: str = self.api_backend._try_create_chat_completion_or_embedding( # noqa: SLF001
|
|
*args,
|
|
messages=messages,
|
|
chat_completion=True,
|
|
**kwargs,
|
|
)
|
|
logger.log_object({"user": user_prompt, "resp": response}, tag="debug_llm")
|
|
|
|
messages.append(
|
|
{
|
|
"role": "assistant",
|
|
"content": response,
|
|
},
|
|
)
|
|
SessionChatHistoryCache().message_set(self.conversation_id, messages)
|
|
return response
|
|
|
|
def get_conversation_id(self) -> str:
|
|
return self.conversation_id
|
|
|
|
def display_history(self) -> None:
|
|
# TODO: Realize a beautiful presentation format for history messages
|
|
pass
|
|
|
|
|
|
class APIBackend(ABC):
|
|
"""
|
|
Abstract base class for LLM API backends
|
|
supporting auto retry, cache and auto continue
|
|
Inner api call should be implemented in the subclass
|
|
"""
|
|
|
|
def __init__(
|
|
self,
|
|
use_chat_cache: bool | None = None,
|
|
dump_chat_cache: bool | None = None,
|
|
use_embedding_cache: bool | None = None,
|
|
dump_embedding_cache: bool | None = None,
|
|
):
|
|
self.dump_chat_cache = LLM_SETTINGS.dump_chat_cache if dump_chat_cache is None else dump_chat_cache
|
|
self.use_chat_cache = LLM_SETTINGS.use_chat_cache if use_chat_cache is None else use_chat_cache
|
|
self.dump_embedding_cache = (
|
|
LLM_SETTINGS.dump_embedding_cache if dump_embedding_cache is None else dump_embedding_cache
|
|
)
|
|
self.use_embedding_cache = (
|
|
LLM_SETTINGS.use_embedding_cache if use_embedding_cache is None else use_embedding_cache
|
|
)
|
|
if self.dump_chat_cache or self.use_chat_cache or self.dump_embedding_cache or self.use_embedding_cache:
|
|
self.cache_file_location = LLM_SETTINGS.prompt_cache_path
|
|
self.cache = SQliteLazyCache(cache_location=self.cache_file_location)
|
|
|
|
self.retry_wait_seconds = LLM_SETTINGS.retry_wait_seconds
|
|
|
|
def build_chat_session(
|
|
self,
|
|
conversation_id: str | None = None,
|
|
session_system_prompt: str | None = None,
|
|
) -> ChatSession:
|
|
"""
|
|
conversation_id is a 256-bit string created by uuid.uuid4() and is also
|
|
the file name under session_cache_folder/ for each conversation
|
|
"""
|
|
return ChatSession(self, conversation_id, session_system_prompt)
|
|
|
|
def _build_messages(
|
|
self,
|
|
user_prompt: str,
|
|
system_prompt: str | None = None,
|
|
former_messages: list[dict[str, Any]] | None = None,
|
|
*,
|
|
shrink_multiple_break: bool = False,
|
|
) -> list[dict[str, Any]]:
|
|
"""
|
|
build the messages to avoid implementing several redundant lines of code
|
|
|
|
"""
|
|
if former_messages is None:
|
|
former_messages = []
|
|
# shrink multiple break will recursively remove multiple breaks(more than 2)
|
|
if shrink_multiple_break:
|
|
while "\n\n\n" in user_prompt:
|
|
user_prompt = user_prompt.replace("\n\n\n", "\n\n")
|
|
if system_prompt is not None:
|
|
while "\n\n\n" in system_prompt:
|
|
system_prompt = system_prompt.replace("\n\n\n", "\n\n")
|
|
system_prompt = LLM_SETTINGS.default_system_prompt if system_prompt is None else system_prompt
|
|
messages = [
|
|
{
|
|
"role": LLM_SETTINGS.system_prompt_role,
|
|
"content": system_prompt,
|
|
},
|
|
]
|
|
messages.extend(former_messages[-1 * LLM_SETTINGS.max_past_message_include :])
|
|
messages.append(
|
|
{
|
|
"role": "user",
|
|
"content": user_prompt,
|
|
},
|
|
)
|
|
return messages
|
|
|
|
def _build_log_messages(self, messages: list[dict[str, Any]]) -> str:
|
|
log_messages = ""
|
|
for m in messages:
|
|
log_messages += (
|
|
f"\n{LogColors.MAGENTA}{LogColors.BOLD}Role:{LogColors.END}"
|
|
f"{LogColors.CYAN}{m['role']}{LogColors.END}\n"
|
|
f"{LogColors.MAGENTA}{LogColors.BOLD}Content:{LogColors.END} "
|
|
f"{LogColors.CYAN}{m['content']}{LogColors.END}\n"
|
|
)
|
|
return log_messages
|
|
|
|
def build_messages_and_create_chat_completion( # type: ignore[no-untyped-def]
|
|
self,
|
|
user_prompt: str,
|
|
system_prompt: str | None = None,
|
|
former_messages: list | None = None,
|
|
chat_cache_prefix: str = "",
|
|
shrink_multiple_break: bool = False,
|
|
*args,
|
|
**kwargs,
|
|
) -> str:
|
|
if former_messages is None:
|
|
former_messages = []
|
|
messages = self._build_messages(
|
|
user_prompt,
|
|
system_prompt,
|
|
former_messages,
|
|
shrink_multiple_break=shrink_multiple_break,
|
|
)
|
|
|
|
resp = self._try_create_chat_completion_or_embedding( # type: ignore[misc]
|
|
*args,
|
|
messages=messages,
|
|
chat_completion=True,
|
|
chat_cache_prefix=chat_cache_prefix,
|
|
**kwargs,
|
|
)
|
|
if isinstance(resp, list):
|
|
raise ValueError("The response of _try_create_chat_completion_or_embedding should be a string.")
|
|
logger.log_object({"system": system_prompt, "user": user_prompt, "resp": resp}, tag="debug_llm")
|
|
return resp
|
|
|
|
def create_embedding(self, input_content: str | list[str], *args, **kwargs) -> list[float] | list[list[float]]: # type: ignore[no-untyped-def]
|
|
input_content_list = [input_content] if isinstance(input_content, str) else input_content
|
|
resp = self._try_create_chat_completion_or_embedding( # type: ignore[misc]
|
|
input_content_list=input_content_list,
|
|
embedding=True,
|
|
*args,
|
|
**kwargs,
|
|
)
|
|
if isinstance(input_content, str):
|
|
return resp[0] # type: ignore[return-value]
|
|
return resp # type: ignore[return-value]
|
|
|
|
def build_messages_and_calculate_token(
|
|
self,
|
|
user_prompt: str,
|
|
system_prompt: str | None,
|
|
former_messages: list[dict[str, Any]] | None = None,
|
|
*,
|
|
shrink_multiple_break: bool = False,
|
|
) -> int:
|
|
if former_messages is None:
|
|
former_messages = []
|
|
messages = self._build_messages(
|
|
user_prompt, system_prompt, former_messages, shrink_multiple_break=shrink_multiple_break
|
|
)
|
|
return self._calculate_token_from_messages(messages)
|
|
|
|
def _try_create_chat_completion_or_embedding( # type: ignore[no-untyped-def]
|
|
self,
|
|
max_retry: int = 10,
|
|
chat_completion: bool = False,
|
|
embedding: bool = False,
|
|
*args,
|
|
**kwargs,
|
|
) -> str | list[list[float]]:
|
|
assert not (chat_completion and embedding), "chat_completion and embedding cannot be True at the same time"
|
|
max_retry = LLM_SETTINGS.max_retry if LLM_SETTINGS.max_retry is not None else max_retry
|
|
for i in range(max_retry):
|
|
try:
|
|
if embedding:
|
|
return self._create_embedding_with_cache(*args, **kwargs)
|
|
if chat_completion:
|
|
return self._create_chat_completion_auto_continue(*args, **kwargs)
|
|
except Exception as e: # noqa: BLE001
|
|
if hasattr(e, "message") and (
|
|
"'messages' must contain the word 'json' in some form" in e.message
|
|
or "\\'messages\\' must contain the word \\'json\\' in some form" in e.message
|
|
):
|
|
kwargs["add_json_in_prompt"] = True
|
|
elif hasattr(e, "message") and embedding and "maximum context length" in e.message:
|
|
kwargs["input_content_list"] = [
|
|
content[: len(content) // 2] for content in kwargs.get("input_content_list", [])
|
|
]
|
|
else:
|
|
time.sleep(self.retry_wait_seconds)
|
|
logger.warning(str(e))
|
|
logger.warning(f"Retrying {i+1}th time...")
|
|
error_message = f"Failed to create chat completion after {max_retry} retries."
|
|
raise RuntimeError(error_message)
|
|
|
|
def _create_chat_completion_add_json_in_prompt(
|
|
self,
|
|
messages: list[dict[str, Any]],
|
|
add_json_in_prompt: bool = False,
|
|
json_mode: bool = False,
|
|
*args: Any,
|
|
**kwargs: Any,
|
|
) -> tuple[str, str | None]:
|
|
"""
|
|
add json related content in the prompt if add_json_in_prompt is True
|
|
"""
|
|
if json_mode and add_json_in_prompt:
|
|
for message in messages[::-1]:
|
|
message["content"] = message["content"] + "\nPlease respond in json format."
|
|
if message["role"] == LLM_SETTINGS.system_prompt_role:
|
|
# NOTE: assumption: systemprompt is always the first message
|
|
break
|
|
return self._create_chat_completion_inner_function(messages=messages, json_mode=json_mode, *args, **kwargs) # type: ignore[misc]
|
|
|
|
def _create_chat_completion_auto_continue(
|
|
self,
|
|
messages: list[dict[str, Any]],
|
|
*args: Any,
|
|
json_mode: bool = False,
|
|
chat_cache_prefix: str = "",
|
|
seed: Optional[int] = None,
|
|
json_target_type: Optional[str] = None,
|
|
**kwargs: Any,
|
|
) -> str:
|
|
"""
|
|
Call the chat completion function and automatically continue the conversation if the finish_reason is length.
|
|
"""
|
|
if seed is None and LLM_SETTINGS.use_auto_chat_cache_seed_gen:
|
|
seed = LLM_CACHE_SEED_GEN.get_next_seed()
|
|
input_content_json = json.dumps(messages)
|
|
input_content_json = (
|
|
chat_cache_prefix + input_content_json + f"<seed={seed}/>"
|
|
) # FIXME this is a hack to make sure the cache represents the round index
|
|
if self.use_chat_cache:
|
|
cache_result = self.cache.chat_get(input_content_json)
|
|
if cache_result is not None:
|
|
if LLM_SETTINGS.log_llm_chat_content:
|
|
logger.info(self._build_log_messages(messages), tag="llm_messages")
|
|
logger.info(f"{LogColors.CYAN}Response:{cache_result}{LogColors.END}", tag="llm_messages")
|
|
return cache_result
|
|
|
|
all_response = ""
|
|
new_messages = deepcopy(messages)
|
|
try_n = 6
|
|
for _ in range(try_n): # for some long code, 3 times may not enough for reasoning models
|
|
if "json_mode" in kwargs:
|
|
del kwargs["json_mode"]
|
|
response, finish_reason = self._create_chat_completion_add_json_in_prompt(
|
|
new_messages, json_mode=json_mode, *args, **kwargs
|
|
) # type: ignore[misc]
|
|
all_response += response
|
|
if finish_reason is None or finish_reason != "length":
|
|
if json_mode:
|
|
try:
|
|
json.loads(all_response)
|
|
except:
|
|
match = re.search(r"```json(.*)```", all_response, re.DOTALL)
|
|
all_response = match.groups()[0] if match else all_response
|
|
json.loads(all_response)
|
|
if json_target_type is not None:
|
|
TypeAdapter(json_target_type).validate_json(all_response)
|
|
if self.dump_chat_cache:
|
|
self.cache.chat_set(input_content_json, all_response)
|
|
return all_response
|
|
new_messages.append({"role": "assistant", "content": response})
|
|
raise RuntimeError(f"Failed to continue the conversation after {try_n} retries.")
|
|
|
|
def _create_embedding_with_cache(
|
|
self, input_content_list: list[str], *args: Any, **kwargs: Any
|
|
) -> list[list[float]]:
|
|
content_to_embedding_dict = {}
|
|
filtered_input_content_list = []
|
|
if self.use_embedding_cache:
|
|
for content in input_content_list:
|
|
cache_result = self.cache.embedding_get(content)
|
|
if cache_result is not None:
|
|
content_to_embedding_dict[content] = cache_result
|
|
else:
|
|
filtered_input_content_list.append(content)
|
|
else:
|
|
filtered_input_content_list = input_content_list
|
|
|
|
if len(filtered_input_content_list) > 0:
|
|
resp = self._create_embedding_inner_function(input_content_list=filtered_input_content_list)
|
|
for index, data in enumerate(resp):
|
|
content_to_embedding_dict[filtered_input_content_list[index]] = data
|
|
if self.dump_embedding_cache:
|
|
self.cache.embedding_set(content_to_embedding_dict)
|
|
return [content_to_embedding_dict[content] for content in input_content_list] # type: ignore[misc]
|
|
|
|
@abstractmethod
|
|
def _calculate_token_from_messages(self, messages: list[dict[str, Any]]) -> int:
|
|
"""
|
|
Calculate the token count from messages
|
|
"""
|
|
raise NotImplementedError("Subclasses must implement this method")
|
|
|
|
@abstractmethod
|
|
def _create_embedding_inner_function( # type: ignore[no-untyped-def]
|
|
self, input_content_list: list[str], *args, **kwargs
|
|
) -> list[list[float]]: # noqa: ARG002
|
|
"""
|
|
Call the embedding function
|
|
"""
|
|
raise NotImplementedError("Subclasses must implement this method")
|
|
|
|
@abstractmethod
|
|
def _create_chat_completion_inner_function( # type: ignore[no-untyped-def] # noqa: C901, PLR0912, PLR0915
|
|
self,
|
|
messages: list[dict[str, Any]],
|
|
json_mode: bool = False,
|
|
*args,
|
|
**kwargs,
|
|
) -> tuple[str, str | None]:
|
|
"""
|
|
Call the chat completion function
|
|
"""
|
|
raise NotImplementedError("Subclasses must implement this method")
|