From 993ca37125e62f9ce07466ced2e1e8bb6b98a362 Mon Sep 17 00:00:00 2001 From: Xu Yang Date: Fri, 28 Feb 2025 15:13:43 +0800 Subject: [PATCH] feat: add type checker to api backend & align litellm and old backend (#647) * move cache auto continue and retry to all api backend * add type checker to json mode output * fix CI * feat: Add json_mode handling and streaming support in chat completion function * lint * fix a bug when returning a dict which value could contain int or bool * remove litellm --------- Co-authored-by: Xu Yang Co-authored-by: Young --- .../app/qlib_rd_loop/factor_from_report.py | 3 +- .../coder/CoSTEER/knowledge_management.py | 3 +- .../coder/data_science/ensemble/__init__.py | 6 +- .../coder/data_science/feature/__init__.py | 6 +- .../coder/data_science/model/__init__.py | 3 + .../data_science/raw_data_loader/__init__.py | 42 +- .../coder/data_science/workflow/__init__.py | 6 +- .../coder/factor_coder/eva_utils.py | 8 +- .../coder/factor_coder/evolving_strategy.py | 6 +- .../components/coder/model_coder/eva_utils.py | 3 +- .../coder/model_coder/evolving_strategy.py | 2 + rdagent/core/experiment.py | 4 +- rdagent/oai/backend/base.py | 463 +++++++++++++++-- rdagent/oai/backend/deprec.py | 471 ++---------------- rdagent/oai/backend/litellm.py | 171 +++---- .../data_mining/developer/feedback.py | 2 + .../scenarios/data_science/dev/feedback.py | 8 +- .../data_science/dev/runner/__init__.py | 2 + .../data_science/proposal/exp_gen.py | 11 +- .../scenarios/data_science/scen/__init__.py | 2 + rdagent/scenarios/kaggle/developer/coder.py | 6 +- .../scenarios/kaggle/developer/feedback.py | 2 + .../scenarios/kaggle/experiment/scenario.py | 2 + rdagent/scenarios/qlib/developer/feedback.py | 3 + rdagent/utils/__init__.py | 5 +- 25 files changed, 634 insertions(+), 606 deletions(-) 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)