Merge pull request #1525 from BerriAI/litellm_router_improvements

[Feat] Router improvements
This commit is contained in:
Ishaan Jaff 2024-01-19 15:02:05 -08:00 committed by GitHub
commit 73684bc93f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 227 additions and 11 deletions

View file

@ -94,6 +94,7 @@ class Router:
timeout: Optional[float] = None,
default_litellm_params={}, # default params for Router.chat.completion.create
set_verbose: bool = False,
debug_level: Literal["DEBUG", "INFO"] = "INFO",
fallbacks: List = [],
allowed_fails: Optional[int] = None,
context_window_fallbacks: List = [],
@ -108,6 +109,11 @@ class Router:
routing_strategy_args: dict = {}, # just for latency-based routing
) -> None:
self.set_verbose = set_verbose
if self.set_verbose:
if debug_level == "INFO":
verbose_router_logger.setLevel(logging.INFO)
elif debug_level == "DEBUG":
verbose_router_logger.setLevel(logging.DEBUG)
self.deployment_names: List = (
[]
) # names of models under litellm_params. ex. azure/chatgpt-v-2
@ -259,6 +265,7 @@ class Router:
raise e
def _completion(self, model: str, messages: List[Dict[str, str]], **kwargs):
model_name = None
try:
# pick the one that is available (lowest TPM/RPM)
deployment = self.get_available_deployment(
@ -271,6 +278,7 @@ class Router:
)
data = deployment["litellm_params"].copy()
kwargs["model_info"] = deployment.get("model_info", {})
model_name = data["model"]
for k, v in self.default_litellm_params.items():
if (
k not in kwargs
@ -292,7 +300,7 @@ class Router:
else:
model_client = potential_model_client
return litellm.completion(
response = litellm.completion(
**{
**data,
"messages": messages,
@ -301,7 +309,14 @@ class Router:
**kwargs,
}
)
verbose_router_logger.info(
f"litellm.completion(model={model_name})\033[32m 200 OK\033[0m"
)
return response
except Exception as e:
verbose_router_logger.info(
f"litellm.completion(model={model_name})\033[31m Exception {str(e)}\033[0m"
)
raise e
async def acompletion(self, model: str, messages: List[Dict[str, str]], **kwargs):
@ -1828,6 +1843,9 @@ class Router:
selected_index = random.choices(range(len(rpms)), weights=weights)[0]
verbose_router_logger.debug(f"\n selected index, {selected_index}")
deployment = healthy_deployments[selected_index]
verbose_router_logger.info(
f"get_available_deployment for model: {model}, Selected deployment: {deployment or deployment[0]} for model: {model}"
)
return deployment or deployment[0]
############## Check if we can do a RPM/TPM based weighted pick #################
tpm = healthy_deployments[0].get("litellm_params").get("tpm", None)
@ -1842,6 +1860,9 @@ class Router:
selected_index = random.choices(range(len(tpms)), weights=weights)[0]
verbose_router_logger.debug(f"\n selected index, {selected_index}")
deployment = healthy_deployments[selected_index]
verbose_router_logger.info(
f"get_available_deployment for model: {model}, Selected deployment: {deployment or deployment[0]} for model: {model}"
)
return deployment or deployment[0]
############## No RPM/TPM passed, we do a random pick #################
@ -1866,8 +1887,13 @@ class Router:
)
if deployment is None:
verbose_router_logger.info(
f"get_available_deployment for model: {model}, No deployment available"
)
raise ValueError("No models available.")
verbose_router_logger.info(
f"get_available_deployment for model: {model}, Selected deployment: {deployment} for model: {model}"
)
return deployment
def flush_cache(self):

View file

@ -10,6 +10,7 @@ import traceback
from litellm import token_counter
from litellm.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm._logging import verbose_router_logger
class LowestTPMLoggingHandler(CustomLogger):
@ -130,6 +131,9 @@ class LowestTPMLoggingHandler(CustomLogger):
Returns a deployment with the lowest TPM/RPM usage.
"""
# get list of potential deployments
verbose_router_logger.debug(
f"get_available_deployments - Usage Based. model_group: {model_group}, healthy_deployments: {healthy_deployments}"
)
current_minute = datetime.now().strftime("%H-%M")
tpm_key = f"{model_group}:tpm:{current_minute}"
rpm_key = f"{model_group}:rpm:{current_minute}"
@ -137,14 +141,31 @@ class LowestTPMLoggingHandler(CustomLogger):
tpm_dict = self.router_cache.get_cache(key=tpm_key)
rpm_dict = self.router_cache.get_cache(key=rpm_key)
verbose_router_logger.debug(
f"tpm_key={tpm_key}, tpm_dict: {tpm_dict}, rpm_dict: {rpm_dict}"
)
try:
input_tokens = token_counter(messages=messages, text=input)
except:
input_tokens = 0
# -----------------------
# Find lowest used model
# ----------------------
lowest_tpm = float("inf")
deployment = None
if tpm_dict is None: # base case
item = random.choice(healthy_deployments)
return item
if tpm_dict is None: # base case - none of the deployments have been used
# Return the 1st deployment where deployment["tpm"] >= input_tokens
for deployment in healthy_deployments:
_deployment_tpm = (
deployment.get("tpm", None)
or deployment.get("litellm_params", {}).get("tpm", None)
or deployment.get("model_info", {}).get("tpm", None)
or float("inf")
)
if _deployment_tpm >= input_tokens:
return deployment
return None
all_deployments = tpm_dict
for d in healthy_deployments:
@ -152,11 +173,6 @@ class LowestTPMLoggingHandler(CustomLogger):
if d["model_info"]["id"] not in all_deployments:
all_deployments[d["model_info"]["id"]] = 0
try:
input_tokens = token_counter(messages=messages, text=input)
except:
input_tokens = 0
for item, item_tpm in all_deployments.items():
## get the item from model list
_deployment = None

View file

@ -69,7 +69,10 @@ def test_async_fallbacks(caplog):
# on circle ci the captured logs get some async task exception logs - filter them out
"Task exception was never retrieved"
captured_logs = [
log for log in captured_logs if "Task exception was never retrieved" not in log
log
for log in captured_logs
if "Task exception was never retrieved" not in log
and "get_available_deployment" not in log
]
print("\n Captured caplog records - ", captured_logs)

View file

@ -698,3 +698,101 @@ async def test_async_fallbacks_max_retries_per_request():
pytest.fail(f"An exception occurred: {e}")
finally:
router.reset()
def test_usage_based_routing_fallbacks():
try:
# [Prod Test]
# IT tests Usage Based Routing with fallbacks
# The Request should fail azure/gpt-4-fast. Then fallback -> "azure/gpt-4-basic" -> "openai-gpt-4"
# It should work with "openai-gpt-4"
import os
import litellm
from litellm import Router
from dotenv import load_dotenv
load_dotenv()
# Constants for TPM and RPM allocation
AZURE_FAST_TPM = 3
AZURE_BASIC_TPM = 4
OPENAI_TPM = 2000
ANTHROPIC_TPM = 100000
def get_azure_params(deployment_name: str):
params = {
"model": f"azure/{deployment_name}",
"api_key": os.environ["AZURE_API_KEY"],
"api_version": os.environ["AZURE_API_VERSION"],
"api_base": os.environ["AZURE_API_BASE"],
}
return params
def get_openai_params(model: str):
params = {
"model": model,
"api_key": os.environ["OPENAI_API_KEY"],
}
return params
def get_anthropic_params(model: str):
params = {
"model": model,
"api_key": os.environ["ANTHROPIC_API_KEY"],
}
return params
model_list = [
{
"model_name": "azure/gpt-4-fast",
"litellm_params": get_azure_params("chatgpt-v-2"),
"tpm": AZURE_FAST_TPM,
},
{
"model_name": "azure/gpt-4-basic",
"litellm_params": get_azure_params("chatgpt-v-2"),
"tpm": AZURE_BASIC_TPM,
},
{
"model_name": "openai-gpt-4",
"litellm_params": get_openai_params("gpt-3.5-turbo"),
"tpm": OPENAI_TPM,
},
{
"model_name": "anthropic-claude-instant-1.2",
"litellm_params": get_anthropic_params("claude-instant-1.2"),
"tpm": ANTHROPIC_TPM,
},
]
# litellm.set_verbose=True
fallbacks_list = [
{"azure/gpt-4-fast": ["azure/gpt-4-basic"]},
{"azure/gpt-4-basic": ["openai-gpt-4"]},
{"openai-gpt-4": ["anthropic-claude-instant-1.2"]},
]
router = Router(
model_list=model_list,
fallbacks=fallbacks_list,
set_verbose=True,
routing_strategy="usage-based-routing",
redis_host=os.environ["REDIS_HOST"],
redis_port=os.environ["REDIS_PORT"],
)
messages = [
{"content": "Tell me a joke.", "role": "user"},
]
response = router.completion(
model="azure/gpt-4-fast", messages=messages, timeout=5
)
print("response: ", response)
print("response._hidden_params: ", response._hidden_params)
# in this test, we expect azure/gpt-4 fast to fail, then azure-gpt-4 basic to fail and then openai-gpt-4 to pass
# the token count of this message is > AZURE_FAST_TPM, > AZURE_BASIC_TPM
assert response._hidden_params["custom_llm_provider"] == "openai"
except Exception as e:
pytest.fail(f"An exception occurred {e}")

View file

@ -375,3 +375,76 @@ def test_model_group_aliases():
# test_model_group_aliases()
def test_usage_based_routing():
"""
in this test we, have a model group with two models in it, model-a and model-b.
Then at some point, we exceed the TPM limit (set in the litellm_params)
for model-a only; but for model-b we are still under the limit
"""
try:
def get_azure_params(deployment_name: str):
params = {
"model": f"azure/{deployment_name}",
"api_key": os.environ["AZURE_API_KEY"],
"api_version": os.environ["AZURE_API_VERSION"],
"api_base": os.environ["AZURE_API_BASE"],
}
return params
model_list = [
{
"model_name": "azure/gpt-4",
"litellm_params": get_azure_params("chatgpt-low-tpm"),
"tpm": 100,
},
{
"model_name": "azure/gpt-4",
"litellm_params": get_azure_params("chatgpt-high-tpm"),
"tpm": 1000,
},
]
router = Router(
model_list=model_list,
set_verbose=True,
debug_level="DEBUG",
routing_strategy="usage-based-routing",
redis_host=os.environ["REDIS_HOST"],
redis_port=os.environ["REDIS_PORT"],
)
messages = [
{"content": "Tell me a joke.", "role": "user"},
]
selection_counts = defaultdict(int)
for _ in range(25):
response = router.completion(
model="azure/gpt-4",
messages=messages,
timeout=5,
mock_response="good morning",
)
# print(response)
selection_counts[response["model"]] += 1
print(selection_counts)
total_requests = sum(selection_counts.values())
# Assert that 'chatgpt-low-tpm' has more than 2 requests
assert (
selection_counts["chatgpt-low-tpm"] > 2
), f"Assertion failed: 'chatgpt-low-tpm' does not have more than 2 request in the weighted load balancer. Selection counts {selection_counts}"
# Assert that 'chatgpt-high-tpm' has about 80% of the total requests
assert (
selection_counts["chatgpt-high-tpm"] / total_requests > 0.8
), f"Assertion failed: 'chatgpt-high-tpm' does not have about 80% of the total requests in the weighted load balancer. Selection counts {selection_counts}"
except Exception as e:
pytest.fail(f"Error occurred: {e}")