diff --git a/rdagent/app/qlib_rd_loop/factor_from_report.py b/rdagent/app/qlib_rd_loop/factor_from_report.py
index b89bb88a..16aa699e 100644
--- a/rdagent/app/qlib_rd_loop/factor_from_report.py
+++ b/rdagent/app/qlib_rd_loop/factor_from_report.py
@@ -1,6 +1,6 @@
import json
from pathlib import Path
-from typing import Any, Tuple
+from typing import Any, Dict, Tuple
import fire
from jinja2 import Environment, StrictUndefined
@@ -49,6 +49,7 @@ def generate_hypothesis(factor_result: dict, report_content: str) -> str:
user_prompt=user_prompt,
system_prompt=system_prompt,
json_mode=True,
+ json_target_type=Dict[str, str],
)
response_json = json.loads(response)
diff --git a/rdagent/components/coder/CoSTEER/knowledge_management.py b/rdagent/components/coder/CoSTEER/knowledge_management.py
index fa473448..eb675e7d 100644
--- a/rdagent/components/coder/CoSTEER/knowledge_management.py
+++ b/rdagent/components/coder/CoSTEER/knowledge_management.py
@@ -6,7 +6,7 @@ import random
import re
from itertools import combinations
from pathlib import Path
-from typing import Union
+from typing import List, Union
from jinja2 import Environment, StrictUndefined
@@ -339,6 +339,7 @@ class CoSTEERRAGStrategyV2(RAGStrategy):
system_prompt=analyze_component_system_prompt,
user_prompt=analyze_component_user_prompt,
json_mode=True,
+ json_target_type=List[int],
),
)["component_no_list"]
return [all_component_nodes[index - 1] for index in sorted(list(set(component_no_list)))]
diff --git a/rdagent/components/coder/data_science/ensemble/__init__.py b/rdagent/components/coder/data_science/ensemble/__init__.py
index 25515e75..ccf7654d 100644
--- a/rdagent/components/coder/data_science/ensemble/__init__.py
+++ b/rdagent/components/coder/data_science/ensemble/__init__.py
@@ -12,6 +12,7 @@ File structure
"""
import json
+from typing import Dict
from rdagent.components.coder.CoSTEER import CoSTEER
from rdagent.components.coder.CoSTEER.evaluators import (
@@ -85,7 +86,10 @@ class EnsembleMultiProcessEvolvingStrategy(MultiProcessEvolvingStrategy):
for _ in range(5):
ensemble_code = json.loads(
APIBackend().build_messages_and_create_chat_completion(
- user_prompt=user_prompt, system_prompt=system_prompt, json_mode=True
+ user_prompt=user_prompt,
+ system_prompt=system_prompt,
+ json_mode=True,
+ json_target_type=Dict[str, str],
)
)["code"]
if ensemble_code != workspace.file_dict.get("ensemble.py"):
diff --git a/rdagent/components/coder/data_science/feature/__init__.py b/rdagent/components/coder/data_science/feature/__init__.py
index 3871215e..ec0d4d37 100644
--- a/rdagent/components/coder/data_science/feature/__init__.py
+++ b/rdagent/components/coder/data_science/feature/__init__.py
@@ -1,4 +1,5 @@
import json
+from typing import Dict
from rdagent.components.coder.CoSTEER import CoSTEER
from rdagent.components.coder.CoSTEER.evaluators import (
@@ -70,7 +71,10 @@ class FeatureMultiProcessEvolvingStrategy(MultiProcessEvolvingStrategy):
for _ in range(5):
feature_code = json.loads(
APIBackend().build_messages_and_create_chat_completion(
- user_prompt=user_prompt, system_prompt=system_prompt, json_mode=True
+ user_prompt=user_prompt,
+ system_prompt=system_prompt,
+ json_mode=True,
+ json_target_type=Dict[str, str],
)
)["code"]
if feature_code != workspace.file_dict.get("feature.py"):
diff --git a/rdagent/components/coder/data_science/model/__init__.py b/rdagent/components/coder/data_science/model/__init__.py
index 7bd41eeb..50594559 100644
--- a/rdagent/components/coder/data_science/model/__init__.py
+++ b/rdagent/components/coder/data_science/model/__init__.py
@@ -1,3 +1,5 @@
+from typing import Dict
+
from rdagent.components.coder.CoSTEER import CoSTEER
from rdagent.components.coder.CoSTEER.evaluators import (
CoSTEERMultiEvaluator,
@@ -83,6 +85,7 @@ class ModelMultiProcessEvolvingStrategy(MultiProcessEvolvingStrategy):
user_prompt=user_prompt,
system_prompt=system_prompt,
json_mode=BatchEditOut.json_mode,
+ json_target_type=Dict[str, str],
)
)
diff --git a/rdagent/components/coder/data_science/raw_data_loader/__init__.py b/rdagent/components/coder/data_science/raw_data_loader/__init__.py
index 2d30ec3d..e1b84044 100644
--- a/rdagent/components/coder/data_science/raw_data_loader/__init__.py
+++ b/rdagent/components/coder/data_science/raw_data_loader/__init__.py
@@ -24,6 +24,7 @@ File structure
import json
import re
+from typing import Dict
from rdagent.app.data_science.conf import DS_RD_SETTING
from rdagent.components.coder.CoSTEER import CoSTEER
@@ -108,20 +109,30 @@ class DataLoaderMultiProcessEvolvingStrategy(MultiProcessEvolvingStrategy):
spec_session = APIBackend().build_chat_session(session_system_prompt=system_prompt)
data_loader_spec = json.loads(
- spec_session.build_chat_completion(user_prompt=data_loader_prompt, json_mode=True)
+ spec_session.build_chat_completion(
+ user_prompt=data_loader_prompt, json_mode=True, json_target_type=Dict[str, str]
+ )
+ )["spec"]
+ feature_spec = json.loads(
+ spec_session.build_chat_completion(
+ user_prompt=feature_prompt, json_mode=True, json_target_type=Dict[str, str]
+ )
+ )["spec"]
+ model_spec = json.loads(
+ spec_session.build_chat_completion(
+ user_prompt=model_prompt, json_mode=True, json_target_type=Dict[str, str]
+ )
+ )["spec"]
+ ensemble_spec = json.loads(
+ spec_session.build_chat_completion(
+ user_prompt=ensemble_prompt, json_mode=True, json_target_type=Dict[str, str]
+ )
+ )["spec"]
+ workflow_spec = json.loads(
+ spec_session.build_chat_completion(
+ user_prompt=workflow_prompt, json_mode=True, json_target_type=Dict[str, str]
+ )
)["spec"]
- feature_spec = json.loads(spec_session.build_chat_completion(user_prompt=feature_prompt, json_mode=True))[
- "spec"
- ]
- model_spec = json.loads(spec_session.build_chat_completion(user_prompt=model_prompt, json_mode=True))[
- "spec"
- ]
- ensemble_spec = json.loads(spec_session.build_chat_completion(user_prompt=ensemble_prompt, json_mode=True))[
- "spec"
- ]
- workflow_spec = json.loads(spec_session.build_chat_completion(user_prompt=workflow_prompt, json_mode=True))[
- "spec"
- ]
else:
data_loader_spec = workspace.file_dict["spec/data_loader.md"]
feature_spec = workspace.file_dict["spec/feature.md"]
@@ -146,7 +157,10 @@ class DataLoaderMultiProcessEvolvingStrategy(MultiProcessEvolvingStrategy):
for _ in range(5):
data_loader_code = json.loads(
APIBackend().build_messages_and_create_chat_completion(
- user_prompt=user_prompt, system_prompt=system_prompt, json_mode=True
+ user_prompt=user_prompt,
+ system_prompt=system_prompt,
+ json_mode=True,
+ json_target_type=Dict[str, str],
)
)["code"]
if data_loader_code != workspace.file_dict.get("load_data.py"):
diff --git a/rdagent/components/coder/data_science/workflow/__init__.py b/rdagent/components/coder/data_science/workflow/__init__.py
index 78179560..aac664b2 100644
--- a/rdagent/components/coder/data_science/workflow/__init__.py
+++ b/rdagent/components/coder/data_science/workflow/__init__.py
@@ -1,4 +1,5 @@
import json
+from typing import Dict
from rdagent.components.coder.CoSTEER import CoSTEER
from rdagent.components.coder.CoSTEER.evaluators import (
@@ -73,7 +74,10 @@ class WorkflowMultiProcessEvolvingStrategy(MultiProcessEvolvingStrategy):
for _ in range(5):
workflow_code = json.loads(
APIBackend().build_messages_and_create_chat_completion(
- user_prompt=user_prompt, system_prompt=system_prompt, json_mode=True
+ user_prompt=user_prompt,
+ system_prompt=system_prompt,
+ json_mode=True,
+ json_target_type=Dict[str, str],
)
)["code"]
if workflow_code != workspace.file_dict.get("main.py"):
diff --git a/rdagent/components/coder/factor_coder/eva_utils.py b/rdagent/components/coder/factor_coder/eva_utils.py
index 40f85bc7..8c7705a4 100644
--- a/rdagent/components/coder/factor_coder/eva_utils.py
+++ b/rdagent/components/coder/factor_coder/eva_utils.py
@@ -2,7 +2,7 @@ import io
import json
from abc import abstractmethod
from pathlib import Path
-from typing import Tuple
+from typing import Dict, Tuple
import pandas as pd
from jinja2 import Environment, StrictUndefined
@@ -212,7 +212,10 @@ class FactorOutputFormatEvaluator(FactorEvaluator):
try:
api = APIBackend() if attempts == 0 else APIBackend(use_chat_cache=False)
resp = api.build_messages_and_create_chat_completion(
- user_prompt=gen_df_info_str, system_prompt=system_prompt, json_mode=True
+ user_prompt=gen_df_info_str,
+ system_prompt=system_prompt,
+ json_mode=True,
+ json_target_type=Dict[str, str | bool | int],
)
resp_dict = json.loads(resp)
resp_dict["output_format_decision"] = str(resp_dict["output_format_decision"]).lower() in ["true", "1"]
@@ -556,6 +559,7 @@ class FactorFinalDecisionEvaluator(FactorEvaluator):
system_prompt=system_prompt,
json_mode=True,
seed=attempts, # in case of useless retrying when cache enabled.
+ json_target_type=Dict[str, str | bool | int],
),
)
final_decision = final_evaluation_dict["final_decision"]
diff --git a/rdagent/components/coder/factor_coder/evolving_strategy.py b/rdagent/components/coder/factor_coder/evolving_strategy.py
index 37b132e7..05804a84 100644
--- a/rdagent/components/coder/factor_coder/evolving_strategy.py
+++ b/rdagent/components/coder/factor_coder/evolving_strategy.py
@@ -2,6 +2,7 @@ from __future__ import annotations
import json
from pathlib import Path
+from typing import Dict
from jinja2 import Environment, StrictUndefined
@@ -168,7 +169,10 @@ class FactorMultiProcessEvolvingStrategy(MultiProcessEvolvingStrategy):
APIBackend(
use_chat_cache=FACTOR_COSTEER_SETTINGS.coder_use_cache
).build_messages_and_create_chat_completion(
- user_prompt=user_prompt, system_prompt=system_prompt, json_mode=True
+ user_prompt=user_prompt,
+ system_prompt=system_prompt,
+ json_mode=True,
+ json_target_type=Dict[str, str],
)
)["code"]
return code
diff --git a/rdagent/components/coder/model_coder/eva_utils.py b/rdagent/components/coder/model_coder/eva_utils.py
index fbff6aaf..25ac7aa7 100644
--- a/rdagent/components/coder/model_coder/eva_utils.py
+++ b/rdagent/components/coder/model_coder/eva_utils.py
@@ -1,6 +1,6 @@
import json
from pathlib import Path
-from typing import Tuple
+from typing import Dict, Tuple
import numpy as np
from jinja2 import Environment, StrictUndefined
@@ -177,6 +177,7 @@ class ModelFinalEvaluator(CoSTEEREvaluator):
user_prompt=user_prompt,
system_prompt=system_prompt,
json_mode=True,
+ json_target_type=Dict[str, str | bool | int],
),
)
if isinstance(final_evaluation_dict["final_decision"], str) and final_evaluation_dict[
diff --git a/rdagent/components/coder/model_coder/evolving_strategy.py b/rdagent/components/coder/model_coder/evolving_strategy.py
index aa22ac7e..1081e6d0 100644
--- a/rdagent/components/coder/model_coder/evolving_strategy.py
+++ b/rdagent/components/coder/model_coder/evolving_strategy.py
@@ -1,5 +1,6 @@
import json
from pathlib import Path
+from typing import Dict
from jinja2 import Environment, StrictUndefined
@@ -96,6 +97,7 @@ class ModelMultiProcessEvolvingStrategy(MultiProcessEvolvingStrategy):
user_prompt=user_prompt,
system_prompt=system_prompt,
json_mode=True,
+ json_target_type=Dict[str, str],
),
)["code"]
return code
diff --git a/rdagent/core/experiment.py b/rdagent/core/experiment.py
index 767365e1..d49dd408 100644
--- a/rdagent/core/experiment.py
+++ b/rdagent/core/experiment.py
@@ -18,7 +18,7 @@ from rdagent.utils import filter_progress_bar
from rdagent.utils.fmt import shrink_text
if typing.TYPE_CHECKING:
- from rdagent.core.proposal import ExperimentFeedback, Hypothesis
+ from rdagent.core.proposal import Hypothesis
from rdagent.utils.env import Env
"""
@@ -225,7 +225,7 @@ class FBWorkspace(Workspace):
"""
for name, code in workspace.file_dict.items():
self.inject_files(**{name: code})
-
+
def copy(self) -> FBWorkspace:
"""
copy the workspace from the original one
diff --git a/rdagent/oai/backend/base.py b/rdagent/oai/backend/base.py
index 8bb20a19..5a0b83ce 100644
--- a/rdagent/oai/backend/base.py
+++ b/rdagent/oai/backend/base.py
@@ -1,53 +1,454 @@
+from __future__ import annotations
+
+import json
+import re
+import sqlite3
+import time
+import uuid
from abc import ABC, abstractmethod
-from typing import Any, Dict, List, Optional, Tuple, Union
+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"""
+ """
+ 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
- @abstractmethod
def build_chat_session(
- self, conversation_id: Optional[str] = None, session_system_prompt: Optional[str] = None
- ) -> Any:
- """Create a new chat session"""
- pass
+ 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)
- @abstractmethod
- def build_messages_and_create_chat_completion(
+ def _build_messages(
self,
user_prompt: str,
- system_prompt: Optional[str] = None,
- former_messages: Optional[List[Any]] = None,
+ 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_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: Any,
- **kwargs: Any,
+ *args,
+ **kwargs,
) -> str:
- """Build messages and get chat completion"""
- pass
+ if former_messages is None:
+ former_messages = []
+ messages = self._build_messages(
+ user_prompt,
+ system_prompt,
+ former_messages,
+ shrink_multiple_break=shrink_multiple_break,
+ )
- @abstractmethod
- def create_embedding(
- self, input_content: Union[str, List[str]], *args: Any, **kwargs: Any
- ) -> Union[List[Any], Any]:
- """Create embeddings for input text"""
- pass
+ 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[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,
+ )
+ return resp # type: ignore[return-value]
- @abstractmethod
def build_messages_and_calculate_token(
self,
user_prompt: str,
- system_prompt: Optional[str],
- former_messages: Optional[List[Dict[str, Any]]] = None,
+ system_prompt: str | None,
+ former_messages: list[dict[str, Any]] | None = None,
*,
shrink_multiple_break: bool = False,
) -> int:
- """Build messages and calculate their token count"""
- pass
+ 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)
-# TODO: seperate cache layer. try to be tranparent.
-class CachedAPIBackend(APIBackend):
- ...
- # @abstractmethod
- # def none_cache_function ...
+ 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""
+ ) # 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(f"{LogColors.CYAN}Response:{cache_result}{LogColors.END}", tag="llm_messages")
+ return cache_result
+
+ all_response = ""
+ new_messages = deepcopy(messages)
+ for _ in range(3):
+ 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("Failed to continue the conversation after 3 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")
diff --git a/rdagent/oai/backend/deprec.py b/rdagent/oai/backend/deprec.py
index a6caeed8..3011385f 100644
--- a/rdagent/oai/backend/deprec.py
+++ b/rdagent/oai/backend/deprec.py
@@ -93,153 +93,6 @@ class ConvManager:
# TODO: reseve line breaks to make it more convient to edit file directly.
-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 DeprecBackend(APIBackend):
"""
This is a unified interface for different backends.
@@ -252,20 +105,10 @@ class DeprecBackend(APIBackend):
# Instead, we can use LLM_SETTINGS directly. If it's difficult to support different backend settings, we can split them into multiple BaseSettings.
def __init__( # noqa: C901, PLR0912, PLR0915
self,
- *,
- chat_api_key: str | None = None,
- chat_model: str | None = None,
- chat_api_base: str | None = None,
- chat_api_version: str | None = None,
- embedding_api_key: str | None = None,
- embedding_model: str | None = None,
- embedding_api_base: str | None = None,
- embedding_api_version: str | None = None,
- use_chat_cache: bool | None = None,
- dump_chat_cache: bool | None = None,
- use_embedding_cache: bool | None = None,
- dump_embedding_cache: bool | None = None,
+ *args: Any,
+ **kwargs: Any,
) -> None:
+ super().__init__(*args, **kwargs)
if LLM_SETTINGS.use_llama2:
self.generator = Llama.build(
ckpt_dir=LLM_SETTINGS.llama2_ckpt_dir,
@@ -310,7 +153,7 @@ class DeprecBackend(APIBackend):
if not os.environ.get("PYTHONHTTPSVERIFY", "") and hasattr(ssl, "_create_unverified_context"):
ssl._create_default_https_context = ssl._create_unverified_context # noqa: SLF001
self.chat_model_map = json.loads(LLM_SETTINGS.chat_model_map)
- self.chat_model = LLM_SETTINGS.chat_model if chat_model is None else chat_model
+ self.chat_model = LLM_SETTINGS.chat_model
self.encoder = None
elif LLM_SETTINGS.chat_use_azure_deepseek:
self.client = ChatCompletionsClient(
@@ -331,37 +174,25 @@ class DeprecBackend(APIBackend):
# Priority: chat_api_key/embedding_api_key > openai_api_key > os.environ.get("OPENAI_API_KEY")
# TODO: Simplify the key design. Consider Pandatic's field alias & priority.
self.chat_api_key = (
- chat_api_key
- or LLM_SETTINGS.chat_openai_api_key
- or LLM_SETTINGS.openai_api_key
- or os.environ.get("OPENAI_API_KEY")
+ LLM_SETTINGS.chat_openai_api_key or LLM_SETTINGS.openai_api_key or os.environ.get("OPENAI_API_KEY")
)
self.embedding_api_key = (
- embedding_api_key
- or LLM_SETTINGS.embedding_openai_api_key
- or LLM_SETTINGS.openai_api_key
- or os.environ.get("OPENAI_API_KEY")
+ LLM_SETTINGS.embedding_openai_api_key or LLM_SETTINGS.openai_api_key or os.environ.get("OPENAI_API_KEY")
)
- self.chat_model = LLM_SETTINGS.chat_model if chat_model is None else chat_model
+ self.chat_model = LLM_SETTINGS.chat_model
self.chat_model_map = json.loads(LLM_SETTINGS.chat_model_map)
self.encoder = self._get_encoder()
self.chat_openai_base_url = LLM_SETTINGS.chat_openai_base_url
self.embedding_openai_base_url = LLM_SETTINGS.embedding_openai_base_url
- self.chat_api_base = LLM_SETTINGS.chat_azure_api_base if chat_api_base is None else chat_api_base
- self.chat_api_version = (
- LLM_SETTINGS.chat_azure_api_version if chat_api_version is None else chat_api_version
- )
+ self.chat_api_base = LLM_SETTINGS.chat_azure_api_base
+ self.chat_api_version = LLM_SETTINGS.chat_azure_api_version
self.chat_stream = LLM_SETTINGS.chat_stream
self.chat_seed = LLM_SETTINGS.chat_seed
- self.embedding_model = LLM_SETTINGS.embedding_model if embedding_model is None else embedding_model
- self.embedding_api_base = (
- LLM_SETTINGS.embedding_azure_api_base if embedding_api_base is None else embedding_api_base
- )
- self.embedding_api_version = (
- LLM_SETTINGS.embedding_azure_api_version if embedding_api_version is None else embedding_api_version
- )
+ self.embedding_model = LLM_SETTINGS.embedding_model
+ self.embedding_api_base = LLM_SETTINGS.embedding_azure_api_base
+ self.embedding_api_version = LLM_SETTINGS.embedding_azure_api_version
if (self.chat_use_azure or self.embedding_use_azure) and (
self.chat_use_azure_token_provider or self.embedding_use_azure_token_provider
@@ -396,23 +227,10 @@ class DeprecBackend(APIBackend):
else openai.OpenAI(api_key=self.embedding_api_key, base_url=self.embedding_openai_base_url)
)
- 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)
-
# transfer the config to the class if the config is not supposed to change during the runtime
self.use_llama2 = LLM_SETTINGS.use_llama2
self.use_gcr_endpoint = LLM_SETTINGS.use_gcr_endpoint
self.chat_use_azure_deepseek = LLM_SETTINGS.chat_use_azure_deepseek
- self.retry_wait_seconds = LLM_SETTINGS.retry_wait_seconds
def _get_encoder(self) -> tiktoken.Encoding:
"""
@@ -442,222 +260,27 @@ class DeprecBackend(APIBackend):
raise
return encoding
- 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_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[Any] | Any: # 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]
- return resp
-
- 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,
- **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""
- ) # 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(f"{LogColors.CYAN}Response:{cache_result}{LogColors.END}", tag="llm_messages")
- return cache_result
-
- all_response = ""
- new_messages = deepcopy(messages)
- for _ in range(3):
- if "json_mode" in kwargs:
- del kwargs["json_mode"]
- response, finish_reason = self._create_chat_completion_inner_function(
- new_messages, json_mode=json_mode, *args, **kwargs
- ) # type: ignore[misc]
- all_response += response
- if finish_reason is None or finish_reason != "length":
- if self.chat_use_azure_deepseek:
- match = re.search(r"(.*?)(.*)", all_response, re.DOTALL)
- think_part, all_response = match.groups() if match else ("", all_response)
- if LLM_SETTINGS.log_llm_chat_content:
- logger.info(f"{LogColors.CYAN}Think:{think_part}{LogColors.END}", tag="llm_messages")
- logger.info(f"{LogColors.CYAN}Response:{all_response}{LogColors.END}", tag="llm_messages")
-
- 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 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("Failed to continue the conversation after 3 retries.")
-
- 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[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_inner_function(*args, **kwargs)
- if chat_completion:
- return self._create_chat_completion_auto_continue(*args, **kwargs)
- except openai.BadRequestError as e: # noqa: PERF203
- logger.warning(str(e))
- logger.warning(f"Retrying {i+1}th time...")
- if (
- "'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 embedding and "maximum context length" in e.message:
- kwargs["input_content_list"] = [
- content[: len(content) // 2] for content in kwargs.get("input_content_list", [])
- ]
- except Exception as e: # noqa: BLE001
- logger.warning(str(e))
- logger.warning(f"Retrying {i+1}th time...")
- time.sleep(self.retry_wait_seconds)
- error_message = f"Failed to create chat completion after {max_retry} retries."
- raise RuntimeError(error_message)
-
def _create_embedding_inner_function( # type: ignore[no-untyped-def]
self, input_content_list: list[str], *args, **kwargs
- ) -> list[Any]: # noqa: ARG002
+ ) -> list[list[float]]: # noqa: ARG002
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
+ for sliced_filtered_input_content_list in [
+ input_content_list[i : i + LLM_SETTINGS.embedding_max_str_num]
+ for i in range(0, len(input_content_list), LLM_SETTINGS.embedding_max_str_num)
+ ]:
+ if self.embedding_use_azure:
+ response = self.embedding_client.embeddings.create(
+ model=self.embedding_model,
+ input=sliced_filtered_input_content_list,
+ )
+ else:
+ response = self.embedding_client.embeddings.create(
+ model=self.embedding_model,
+ input=sliced_filtered_input_content_list,
+ )
+ for index, data in enumerate(response.data):
+ content_to_embedding_dict[sliced_filtered_input_content_list[index]] = data.embedding
- if len(filtered_input_content_list) > 0:
- for sliced_filtered_input_content_list in [
- filtered_input_content_list[i : i + LLM_SETTINGS.embedding_max_str_num]
- for i in range(0, len(filtered_input_content_list), LLM_SETTINGS.embedding_max_str_num)
- ]:
- if self.embedding_use_azure:
- response = self.embedding_client.embeddings.create(
- model=self.embedding_model,
- input=sliced_filtered_input_content_list,
- )
- else:
- response = self.embedding_client.embeddings.create(
- model=self.embedding_model,
- input=sliced_filtered_input_content_list,
- )
- for index, data in enumerate(response.data):
- content_to_embedding_dict[sliced_filtered_input_content_list[index]] = data.embedding
-
- 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]
def _build_log_messages(self, messages: list[dict[str, Any]]) -> str:
@@ -674,10 +297,6 @@ class DeprecBackend(APIBackend):
def _create_chat_completion_inner_function( # type: ignore[no-untyped-def] # noqa: C901, PLR0912, PLR0915
self,
messages: list[dict[str, Any]],
- temperature: float | None = None,
- max_tokens: int | None = None,
- frequency_penalty: float | None = None,
- presence_penalty: float | None = None,
json_mode: bool = False,
add_json_in_prompt: bool = False,
*args,
@@ -695,14 +314,10 @@ class DeprecBackend(APIBackend):
logger.info(self._build_log_messages(messages), tag="llm_messages")
# TODO: fail to use loguru adaptor due to stream response
- if temperature is None:
- temperature = LLM_SETTINGS.chat_temperature
- if max_tokens is None:
- max_tokens = LLM_SETTINGS.chat_max_tokens
- if frequency_penalty is None:
- frequency_penalty = LLM_SETTINGS.chat_frequency_penalty
- if presence_penalty is None:
- presence_penalty = LLM_SETTINGS.chat_presence_penalty
+ temperature = LLM_SETTINGS.chat_temperature
+ max_tokens = LLM_SETTINGS.chat_max_tokens
+ frequency_penalty = LLM_SETTINGS.chat_frequency_penalty
+ presence_penalty = LLM_SETTINGS.chat_presence_penalty
# Use index 4 to skip the current function and intermediate calls,
# and get the locals of the caller's frame.
@@ -784,6 +399,11 @@ class DeprecBackend(APIBackend):
finish_reason = response.choices[0].finish_reason
if LLM_SETTINGS.log_llm_chat_content:
logger.info(f"{LogColors.CYAN}Response:{resp}{LogColors.END}", tag="llm_messages")
+ match = re.search(r"(.*?)(.*)", resp, re.DOTALL)
+ think_part, resp = match.groups() if match else ("", resp)
+ if LLM_SETTINGS.log_llm_chat_content:
+ logger.info(f"{LogColors.CYAN}Think:{think_part}{LogColors.END}", tag="llm_messages")
+ logger.info(f"{LogColors.CYAN}Response:{resp}{LogColors.END}", tag="llm_messages")
else:
call_kwargs = dict(
model=model,
@@ -869,18 +489,3 @@ class DeprecBackend(APIBackend):
num_tokens += tokens_per_name
num_tokens += 3 # every reply is primed with assistant
return num_tokens
-
- 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)
diff --git a/rdagent/oai/backend/litellm.py b/rdagent/oai/backend/litellm.py
index 998273c6..c13f976e 100644
--- a/rdagent/oai/backend/litellm.py
+++ b/rdagent/oai/backend/litellm.py
@@ -1,30 +1,20 @@
-import os
-import uuid
-from pathlib import Path
-from typing import Any, Dict, List, Optional, Union
+from typing import Any
-from litellm import acompletion, completion
-from litellm import encode as encode_litellm
-from litellm import token_counter
+from litellm import completion, embedding, token_counter
-from rdagent.core.conf import ExtendedBaseSettings
-from rdagent.core.utils import LLM_CACHE_SEED_GEN, SingletonBaseClass, import_class
from rdagent.log import LogColors
from rdagent.log import rdagent_logger as logger
from rdagent.oai.backend.base import APIBackend
-from rdagent.oai.llm_conf import LLM_SETTINGS
+from rdagent.oai.llm_conf import LLMSettings
-class LiteLLMSettings(ExtendedBaseSettings):
+class LiteLLMSettings(LLMSettings):
class Config:
env_prefix = "LITELLM_"
"""Use `LITELLM_` as prefix for environment variables"""
- # LiteLLM backend related config
- chat_model: str = "openai/gpt-4o"
- # LiteLLM embedding related config
- embedding_model: str = "openai/text-embedding-3-small"
+ # Placeholder for LiteLLM specific settings, so far it's empty
LITELLM_SETTINGS = LiteLLMSettings()
@@ -33,81 +23,28 @@ LITELLM_SETTINGS = LiteLLMSettings()
class LiteLLMAPIBackend(APIBackend):
"""LiteLLM implementation of APIBackend interface"""
- def __init__(self, litellm_model_name: str = "", litellm_api_key: str = "", *args: Any, **kwargs: Any) -> None:
- super().__init__()
- if len(args) > 0 or len(kwargs) > 0:
- logger.warning("LiteLLM backend does not support any additional arguments")
+ def __init__(self, *args: Any, **kwargs: Any) -> None:
+ super().__init__(*args, **kwargs)
- def build_chat_session(
- self, conversation_id: Optional[str] = None, session_system_prompt: Optional[str] = None
- ) -> Any:
- """Create a new chat session using LiteLLM"""
- # return {
- # "conversation_id": conversation_id or str(uuid.uuid4()),
- # "system_prompt": session_system_prompt,
- # "messages": []
- # }
- raise NotImplementedError("LiteLLM backend does not support chat session creation")
- # TODO: Implement the chat session creation logic , with ChatSession class
-
- def build_messages_and_create_chat_completion(
- self,
- user_prompt: str,
- system_prompt: Optional[str] = None,
- former_messages: Optional[List[Any]] = None,
- chat_cache_prefix: str = "",
- shrink_multiple_break: bool = False,
- *args: Any,
- **kwargs: Any,
- ) -> str:
- """Build messages and get LiteLLM chat completion"""
- messages = []
-
- if system_prompt:
- messages.append({"role": "system", "content": system_prompt})
-
- if former_messages:
- messages.extend(former_messages)
-
- messages.append({"role": "user", "content": user_prompt})
- model_name = LITELLM_SETTINGS.chat_model
- # Call LiteLLM completion
- response = completion(
- model=model_name,
+ def _calculate_token_from_messages(self, messages: list[dict[str, Any]]) -> int:
+ """
+ Calculate the token count from messages
+ """
+ num_tokens = token_counter(
+ model=LITELLM_SETTINGS.chat_model,
messages=messages,
- stream=kwargs.get("stream", False),
- temperature=kwargs.get("temperature", 0.7),
- max_tokens=kwargs.get("max_tokens", 1000),
- **kwargs,
- )
- logger.info(
- f"{LogColors.GREEN}Using chat model{LogColors.END} {model_name}",
- tag="debug_llm",
)
+ logger.info(f"{LogColors.CYAN}Token count: {LogColors.END} {num_tokens}", tag="debug_litellm_token")
+ return num_tokens
- if system_prompt:
- logger.info(f"{LogColors.RED}system:{LogColors.END} {system_prompt}", tag="debug_llm")
- if former_messages:
- for message in former_messages:
- logger.info(f"{LogColors.CYAN}{message['role']}:{LogColors.END} {message['content']}", tag="debug_llm")
- else:
- logger.info(
- f"{LogColors.RED}user:{LogColors.END} {user_prompt}\n{LogColors.BLUE}resp(next row):\n{LogColors.END} {response.choices[0].message.content}",
- tag="debug_llm",
- )
-
- return str(response.choices[0].message.content)
-
- def create_embedding(self, input_content: str | list[str], *args: Any, **kwargs: Any) -> list[Any] | Any:
- """Create embeddings using LiteLLM"""
- from litellm import embedding
-
- single_input = False
- if isinstance(input_content, str):
- input_content = [input_content]
- single_input = True
+ def _create_embedding_inner_function(
+ self, input_content_list: list[str], *args: Any, **kwargs: Any
+ ) -> list[list[float]]: # noqa: ARG002
+ """
+ Call the embedding function
+ """
response_list = []
- for input_content_iter in input_content:
+ for input_content_iter in input_content_list:
model_name = LITELLM_SETTINGS.embedding_model or "azure/text-embedding-3-small"
logger.info(f"{LogColors.GREEN}Using emb model{LogColors.END} {model_name}", tag="debug_litellm_emb")
logger.info(f"Creating embedding for: {input_content_iter}", tag="debug_litellm_emb")
@@ -116,34 +53,56 @@ class LiteLLMAPIBackend(APIBackend):
response = embedding(
model=model_name,
input=input_content_iter,
+ *args,
**kwargs,
)
response_list.append(response.data[0]["embedding"])
- if single_input:
- return response_list[0]
return response_list
- def build_messages_and_calculate_token(
+ def _create_chat_completion_inner_function( # type: ignore[no-untyped-def] # noqa: C901, PLR0912, PLR0915
self,
- user_prompt: str,
- system_prompt: Optional[str],
- former_messages: Optional[List[Dict[str, Any]]] = None,
- shrink_multiple_break: bool = False,
- ) -> int:
- """Build messages and calculate their token count using LiteLLM"""
- messages = []
+ messages: list[dict[str, Any]],
+ json_mode: bool = False,
+ *args,
+ **kwargs,
+ ) -> tuple[str, str | None]:
+ """
+ Call the chat completion function
+ """
+ if json_mode:
+ kwargs["response_format"] = {"type": "json_object"}
- if system_prompt:
- messages.append({"role": "system", "content": system_prompt})
-
- if former_messages:
- messages.extend(former_messages)
-
- messages.append({"role": "user", "content": user_prompt})
-
- num_tokens = token_counter(
+ # Call LiteLLM completion
+ response = completion(
model=LITELLM_SETTINGS.chat_model,
messages=messages,
+ stream=LITELLM_SETTINGS.chat_stream,
+ temperature=LITELLM_SETTINGS.chat_temperature,
+ max_tokens=LITELLM_SETTINGS.chat_max_tokens,
+ **kwargs,
)
- logger.info(f"{LogColors.CYAN}Token count: {LogColors.END} {num_tokens}", tag="debug_litellm_token")
- return num_tokens
+ logger.info(
+ f"{LogColors.GREEN}Using chat model{LogColors.END} {LITELLM_SETTINGS.chat_model}", tag="llm_messages"
+ )
+
+ if LITELLM_SETTINGS.chat_stream:
+ logger.info(f"{LogColors.BLUE}assistant:{LogColors.END}", tag="llm_messages")
+ content = ""
+ finish_reason = None
+ for message in response:
+ if message["choices"][0]["finish_reason"]:
+ finish_reason = message["choices"][0]["finish_reason"]
+ if "content" in message["choices"][0]["delta"]:
+ chunk = (
+ message["choices"][0]["delta"]["content"] or ""
+ ) # when finish_reason is "stop", content is None
+ content += chunk
+ logger.info(LogColors.CYAN + chunk + LogColors.END, raw=True, tag="llm_messages")
+
+ logger.info("\n", raw=True, tag="llm_messages")
+ else:
+ content = str(response.choices[0].message.content)
+ finish_reason = response.choices[0].finish_reason
+ logger.info(f"{LogColors.BLUE}assistant:{LogColors.END} {content}", tag="llm_messages")
+
+ return content, finish_reason
diff --git a/rdagent/scenarios/data_mining/developer/feedback.py b/rdagent/scenarios/data_mining/developer/feedback.py
index 271efd16..e1ca2927 100644
--- a/rdagent/scenarios/data_mining/developer/feedback.py
+++ b/rdagent/scenarios/data_mining/developer/feedback.py
@@ -3,6 +3,7 @@
import json
from pathlib import Path
+from typing import Dict
from jinja2 import Environment, StrictUndefined
@@ -59,6 +60,7 @@ class DMModelExperiment2Feedback(Experiment2Feedback):
user_prompt=user_prompt,
system_prompt=system_prompt,
json_mode=True,
+ json_target_type=Dict[str, str | bool | int],
)
# Parse the JSON response to extract the feedback
diff --git a/rdagent/scenarios/data_science/dev/feedback.py b/rdagent/scenarios/data_science/dev/feedback.py
index 0ab65c75..5734974f 100644
--- a/rdagent/scenarios/data_science/dev/feedback.py
+++ b/rdagent/scenarios/data_science/dev/feedback.py
@@ -1,20 +1,17 @@
import json
+from typing import Dict
import pandas as pd
-from rdagent.components.knowledge_management.graph import UndirectedNode
-from rdagent.core.experiment import Experiment
-from rdagent.core.prompts import Prompts
from rdagent.core.proposal import (
Experiment2Feedback,
ExperimentFeedback,
HypothesisFeedback,
)
-from rdagent.log import rdagent_logger as logger
from rdagent.oai.llm_utils import APIBackend
from rdagent.scenarios.data_science.experiment.experiment import DSExperiment
from rdagent.scenarios.data_science.proposal.exp_gen import DSTrace
-from rdagent.utils import convert2bool, remove_path_info_from_str
+from rdagent.utils import convert2bool
from rdagent.utils.agent.tpl import T
from rdagent.utils.repo.diff import generate_diff_from_dict
@@ -76,6 +73,7 @@ class DSExperiment2Feedback(Experiment2Feedback):
user_prompt=user_prompt,
system_prompt=system_prompt,
json_mode=True,
+ json_target_type=Dict[str, str | bool | int],
)
)
diff --git a/rdagent/scenarios/data_science/dev/runner/__init__.py b/rdagent/scenarios/data_science/dev/runner/__init__.py
index efb88196..a51cca4e 100644
--- a/rdagent/scenarios/data_science/dev/runner/__init__.py
+++ b/rdagent/scenarios/data_science/dev/runner/__init__.py
@@ -1,4 +1,5 @@
from pathlib import Path
+from typing import Dict
import pandas as pd
@@ -54,6 +55,7 @@ class DSRunnerMultiProcessEvolvingStrategy(MultiProcessEvolvingStrategy):
user_prompt=user_prompt,
system_prompt=system_prompt,
json_mode=BatchEditOut.json_mode,
+ json_target_type=Dict[str, str],
)
)
diff --git a/rdagent/scenarios/data_science/proposal/exp_gen.py b/rdagent/scenarios/data_science/proposal/exp_gen.py
index 261d0a1c..4cae7546 100644
--- a/rdagent/scenarios/data_science/proposal/exp_gen.py
+++ b/rdagent/scenarios/data_science/proposal/exp_gen.py
@@ -1,5 +1,5 @@
import json
-from typing import Literal
+from typing import Dict, Literal
import pandas as pd
@@ -212,7 +212,7 @@ class DSExpGen(ExpGen):
resp_dict = json.loads(
APIBackend().build_messages_and_create_chat_completion(
- user_prompt=user_prompt, system_prompt=system_prompt, json_mode=True
+ user_prompt=user_prompt, system_prompt=system_prompt, json_mode=True, json_target_type=dict
)
)
@@ -365,7 +365,7 @@ class DSExpGen(ExpGen):
resp_dict_component: dict = json.loads(
APIBackend().build_messages_and_create_chat_completion(
- component_user_prompt, component_sys_prompt, json_mode=True
+ component_user_prompt, component_sys_prompt, json_mode=True, json_target_type=Dict[str, str]
)
)
@@ -417,7 +417,10 @@ class DSExpGen(ExpGen):
def _f(user_prompt):
resp_dict = json.loads(
APIBackend().build_messages_and_create_chat_completion(
- user_prompt=user_prompt, system_prompt=system_prompt, json_mode=True
+ user_prompt=user_prompt,
+ system_prompt=system_prompt,
+ json_mode=True,
+ direct_exp_gen=Dict[str, Dict[str, str]],
)
)
assert "hypothesis_proposal" in resp_dict, "Hypothesis proposal not provided."
diff --git a/rdagent/scenarios/data_science/scen/__init__.py b/rdagent/scenarios/data_science/scen/__init__.py
index 586ac8b4..b64b82ee 100644
--- a/rdagent/scenarios/data_science/scen/__init__.py
+++ b/rdagent/scenarios/data_science/scen/__init__.py
@@ -1,6 +1,7 @@
import json
import os
from pathlib import Path
+from typing import Dict
import pandas as pd
from PIL import Image, TiffTags
@@ -258,6 +259,7 @@ class DataScienceScen(Scenario):
user_prompt=user_prompt,
system_prompt=sys_prompt,
json_mode=True,
+ json_target_type=Dict[str, str | int | bool],
)
response_json_analysis = json.loads(response_analysis)
diff --git a/rdagent/scenarios/kaggle/developer/coder.py b/rdagent/scenarios/kaggle/developer/coder.py
index a04054c5..dd1b4475 100644
--- a/rdagent/scenarios/kaggle/developer/coder.py
+++ b/rdagent/scenarios/kaggle/developer/coder.py
@@ -1,5 +1,6 @@
import json
from pathlib import Path
+from typing import Dict, List
from jinja2 import Environment, StrictUndefined
@@ -58,7 +59,10 @@ class KGModelFeatureSelectionCoder(Developer[KGModelExperiment]):
chosen_index = json.loads(
APIBackend().build_messages_and_create_chat_completion(
- user_prompt=user_prompt, system_prompt=system_prompt, json_mode=True
+ user_prompt=user_prompt,
+ system_prompt=system_prompt,
+ json_mode=True,
+ json_target_type=Dict[str, List[int]],
)
).get("Selected Group Index", [i + 1 for i in range(len(exp.experiment_workspace.data_description))])
chosen_index_to_list_index = [i - 1 for i in chosen_index]
diff --git a/rdagent/scenarios/kaggle/developer/feedback.py b/rdagent/scenarios/kaggle/developer/feedback.py
index 39f0acac..d023ef91 100644
--- a/rdagent/scenarios/kaggle/developer/feedback.py
+++ b/rdagent/scenarios/kaggle/developer/feedback.py
@@ -1,5 +1,6 @@
import json
from pathlib import Path
+from typing import Dict
import pandas as pd
from jinja2 import Environment, StrictUndefined
@@ -148,6 +149,7 @@ class KGExperiment2Feedback(Experiment2Feedback):
user_prompt=usr_prompt,
system_prompt=sys_prompt,
json_mode=True,
+ json_target_type=Dict[str, str | bool | int],
)
response_json = json.loads(response)
diff --git a/rdagent/scenarios/kaggle/experiment/scenario.py b/rdagent/scenarios/kaggle/experiment/scenario.py
index 2816d78a..2a1128b6 100644
--- a/rdagent/scenarios/kaggle/experiment/scenario.py
+++ b/rdagent/scenarios/kaggle/experiment/scenario.py
@@ -3,6 +3,7 @@ import json
import pickle
from datetime import datetime, timezone
from pathlib import Path
+from typing import Dict
import pandas as pd
from jinja2 import Environment, StrictUndefined
@@ -93,6 +94,7 @@ class KGScenario(Scenario):
user_prompt=user_prompt,
system_prompt=sys_prompt,
json_mode=True,
+ json_target_type=Dict[str, str | bool | int],
)
response_json_analysis = json.loads(response_analysis)
diff --git a/rdagent/scenarios/qlib/developer/feedback.py b/rdagent/scenarios/qlib/developer/feedback.py
index 1a80cfe6..73c6f84c 100644
--- a/rdagent/scenarios/qlib/developer/feedback.py
+++ b/rdagent/scenarios/qlib/developer/feedback.py
@@ -1,5 +1,6 @@
import json
from pathlib import Path
+from typing import Dict
import pandas as pd
from jinja2 import Environment, StrictUndefined
@@ -102,6 +103,7 @@ class QlibFactorExperiment2Feedback(Experiment2Feedback):
user_prompt=usr_prompt,
system_prompt=sys_prompt,
json_mode=True,
+ json_target_type=Dict[str, str | bool | int],
)
# Parse the JSON response to extract the feedback
@@ -159,6 +161,7 @@ class QlibModelExperiment2Feedback(Experiment2Feedback):
user_prompt=user_prompt,
system_prompt=system_prompt,
json_mode=True,
+ json_target_type=Dict[str, str | bool | int],
)
# Parse the JSON response to extract the feedback
diff --git a/rdagent/utils/__init__.py b/rdagent/utils/__init__.py
index 6b5acaee..75b37a70 100644
--- a/rdagent/utils/__init__.py
+++ b/rdagent/utils/__init__.py
@@ -120,7 +120,10 @@ def filter_progress_bar(stdout: str) -> str:
response = json.loads(
APIBackend().build_messages_and_create_chat_completion(
- user_prompt=user_prompt, system_prompt=system_prompt, json_mode=True
+ user_prompt=user_prompt,
+ system_prompt=system_prompt,
+ json_mode=True,
+ json_target_type=dict,
)
)
needs_sub = response.get("needs_sub", True)