Files
NexQuant/rdagent/oai/backend/litellm.py
T
炼金术师华华 5090c6153f feat(backend): integrate LiteLLM API Backend (#564)
* File structure for supporting litellm

* more litellm support

* feat: Add CachedAPIBackend class and dynamic API backend retrieval function

* fix: update benchmark folder path and add default values for architecture and hyperparameters

* feat: add LiteLLMAPIBackend and DeprecBackend ; changed structure of the project ; with bus

* fix : deprec_backend

* feat: Add LiteLLMAPIBackend class and related features; update configuration and test cases.

* feat: Enhance LiteLLMAPIBackend with encoder support and dynamic argument handling;Enhance log Colors

* lint

* fix lint...

* fix: Lint

* fix:make auto-lint

* fix:test oai

* fix:redundant _abckend.py

* fix: Optimize LiteLLMAPIBackend on token counting functiona, and clean up unused code;add test on this function

* feat: Add LiteLLMSettings class and update model settings usage

* fix: Update LiteLLMSettings environment variable prefix and model configurations

* fix : gitignore

* test: Consolidate and relocate test files for litellm backend and oai

* fix : lint

* fix: lint

* auto lint

* lint

* LINT

* lint

* chore: remove deprecated backend configuration comments

* refactor: Remove unused functions and imports from deprec.py and llm_utils.py

* refactor: Move md5_hash function from deprec.py to llm_utils.py

* chore: Remove extra newline and add missing import in deprec.py

* lint

* refactor: Move md5_hash function to utils module

* lint

* lint

* lint

---------

Co-authored-by: Young <afe.young@gmail.com>
Co-authored-by: Yihua Chen <v-yihuachen@microsoft.com>
2025-02-13 15:16:18 +08:00

150 lines
5.5 KiB
Python

import os
import uuid
from pathlib import Path
from typing import Any, Dict, List, Optional, Union
from litellm import acompletion, completion
from litellm import encode as encode_litellm
from litellm import 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
class LiteLLMSettings(ExtendedBaseSettings):
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"
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 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,
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",
)
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
response_list = []
for input_content_iter in input_content:
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")
if not isinstance(input_content_iter, str):
raise ValueError("Input content must be a string")
response = embedding(
model=model_name,
input=input_content_iter,
**kwargs,
)
response_list.append(response.data[0]["embedding"])
if single_input:
return response_list[0]
return response_list
def build_messages_and_calculate_token(
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 = []
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(
model=LITELLM_SETTINGS.chat_model,
messages=messages,
)
logger.info(f"{LogColors.CYAN}Token count: {LogColors.END} {num_tokens}", tag="debug_litellm_token")
return num_tokens