mirror of
https://github.com/NicolasBohn/NexQuant.git
synced 2026-08-01 17:37:43 +00:00
9986b5f9ce
* refine prompt * refine the wording * add ratelimit retry to align with the suggested wait seconds * add max retry to 0 * don't delete hist --------- Co-authored-by: Xu <v-xuminrui@microsoft.com> Co-authored-by: Xu Yang <xuyang1@microsoft.com>
156 lines
5.8 KiB
Python
156 lines
5.8 KiB
Python
from typing import Any, Literal, cast
|
|
|
|
from litellm import (
|
|
completion,
|
|
completion_cost,
|
|
embedding,
|
|
supports_response_schema,
|
|
token_counter,
|
|
)
|
|
|
|
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 LLMSettings
|
|
|
|
|
|
class LiteLLMSettings(LLMSettings):
|
|
|
|
class Config:
|
|
env_prefix = "LITELLM_"
|
|
"""Use `LITELLM_` as prefix for environment variables"""
|
|
|
|
# Placeholder for LiteLLM specific settings, so far it's empty
|
|
|
|
|
|
LITELLM_SETTINGS = LiteLLMSettings()
|
|
logger.info(f"{LITELLM_SETTINGS}")
|
|
ACC_COST = 0.0
|
|
|
|
|
|
class LiteLLMAPIBackend(APIBackend):
|
|
"""LiteLLM implementation of APIBackend interface"""
|
|
|
|
def __init__(self, *args: Any, **kwargs: Any) -> None:
|
|
super().__init__(*args, **kwargs)
|
|
|
|
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,
|
|
)
|
|
logger.info(f"{LogColors.CYAN}Token count: {LogColors.END} {num_tokens}", tag="debug_litellm_token")
|
|
return num_tokens
|
|
|
|
def _create_embedding_inner_function(
|
|
self, input_content_list: list[str], *args: Any, **kwargs: Any
|
|
) -> list[list[float]]: # noqa: ARG002
|
|
"""
|
|
Call the embedding function
|
|
"""
|
|
model_name = LITELLM_SETTINGS.embedding_model
|
|
logger.info(f"{LogColors.GREEN}Using emb model{LogColors.END} {model_name}", tag="debug_litellm_emb")
|
|
logger.info(f"Creating embedding for: {input_content_list}", tag="debug_litellm_emb")
|
|
response = embedding(
|
|
model=model_name,
|
|
input=input_content_list,
|
|
*args,
|
|
**kwargs,
|
|
)
|
|
response_list = [data["embedding"] for data in response.data]
|
|
return response_list
|
|
|
|
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
|
|
"""
|
|
if json_mode and supports_response_schema(model=LITELLM_SETTINGS.chat_model):
|
|
kwargs["response_format"] = {"type": "json_object"}
|
|
|
|
logger.info(self._build_log_messages(messages), tag="llm_messages")
|
|
# Call LiteLLM completion
|
|
model = LITELLM_SETTINGS.chat_model
|
|
temperature = LITELLM_SETTINGS.chat_temperature
|
|
max_tokens = LITELLM_SETTINGS.chat_max_tokens
|
|
reasoning_effort = LITELLM_SETTINGS.reasoning_effort
|
|
|
|
if LITELLM_SETTINGS.chat_model_map:
|
|
for t, mc in LITELLM_SETTINGS.chat_model_map.items():
|
|
if t in logger._tag:
|
|
model = mc["model"]
|
|
if "temperature" in mc:
|
|
temperature = float(mc["temperature"])
|
|
if "max_tokens" in mc:
|
|
max_tokens = int(mc["max_tokens"])
|
|
if "reasoning_effort" in mc:
|
|
if mc["reasoning_effort"] in ["low", "medium", "high"]:
|
|
reasoning_effort = cast(Literal["low", "medium", "high"], mc["reasoning_effort"])
|
|
else:
|
|
reasoning_effort = None
|
|
break
|
|
response = completion(
|
|
model=model,
|
|
messages=messages,
|
|
stream=LITELLM_SETTINGS.chat_stream,
|
|
temperature=temperature,
|
|
max_tokens=max_tokens,
|
|
reasoning_effort=reasoning_effort,
|
|
max_retries=0,
|
|
**kwargs,
|
|
)
|
|
logger.info(f"{LogColors.GREEN}Using chat model{LogColors.END} {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
|
|
finish_reason_str = (
|
|
f"({LogColors.RED}Finish reason: {finish_reason}{LogColors.END})"
|
|
if finish_reason and finish_reason != "stop"
|
|
else ""
|
|
)
|
|
logger.info(f"{LogColors.BLUE}assistant:{LogColors.END} {finish_reason_str}\n{content}", tag="llm_messages")
|
|
|
|
global ACC_COST
|
|
cost = completion_cost(model=model, messages=messages, completion=content)
|
|
ACC_COST += cost
|
|
logger.info(
|
|
f"Current Cost: ${float(cost):.10f}; Accumulated Cost: ${float(ACC_COST):.10f}; {finish_reason=}",
|
|
)
|
|
prompt_tokens = token_counter(model=model, messages=messages)
|
|
completion_tokens = token_counter(model=model, text=content)
|
|
logger.log_object(
|
|
{
|
|
"model": model,
|
|
"prompt_tokens": prompt_tokens,
|
|
"completion_tokens": completion_tokens,
|
|
"cost": cost,
|
|
"accumulated_cost": ACC_COST,
|
|
},
|
|
tag="token_cost",
|
|
)
|
|
return content, finish_reason
|