mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
fix(lowest_tpm_rpm_routing.py): fix base case where max tpm/rpm is 0
This commit is contained in:
parent
5d428ac94c
commit
47ca223d0b
2 changed files with 68 additions and 36 deletions
|
|
@ -11,6 +11,7 @@ from litellm import token_counter
|
|||
from litellm.caching import DualCache
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.utils import print_verbose
|
||||
|
||||
|
||||
class LowestTPMLoggingHandler(CustomLogger):
|
||||
|
|
@ -153,27 +154,21 @@ class LowestTPMLoggingHandler(CustomLogger):
|
|||
# Find lowest used model
|
||||
# ----------------------
|
||||
lowest_tpm = float("inf")
|
||||
deployment = None
|
||||
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
|
||||
if tpm_dict is None: # base case - none of the deployments have been used
|
||||
# initialize a tpm dict with {model_id: 0}
|
||||
tpm_dict = {}
|
||||
for deployment in healthy_deployments:
|
||||
tpm_dict[deployment["model_info"]["id"]] = 0
|
||||
else:
|
||||
for d in healthy_deployments:
|
||||
## if healthy deployment not yet used
|
||||
if d["model_info"]["id"] not in tpm_dict:
|
||||
tpm_dict[d["model_info"]["id"]] = 0
|
||||
|
||||
all_deployments = tpm_dict
|
||||
for d in healthy_deployments:
|
||||
## if healthy deployment not yet used
|
||||
if d["model_info"]["id"] not in all_deployments:
|
||||
all_deployments[d["model_info"]["id"]] = 0
|
||||
|
||||
deployment = None
|
||||
for item, item_tpm in all_deployments.items():
|
||||
## get the item from model list
|
||||
_deployment = None
|
||||
|
|
@ -184,24 +179,27 @@ class LowestTPMLoggingHandler(CustomLogger):
|
|||
if _deployment is None:
|
||||
continue # skip to next one
|
||||
|
||||
_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")
|
||||
)
|
||||
_deployment_tpm = None
|
||||
if _deployment_tpm is None:
|
||||
_deployment_tpm = _deployment.get("tpm")
|
||||
if _deployment_tpm is None:
|
||||
_deployment_tpm = _deployment.get("litellm_params", {}).get("tpm")
|
||||
if _deployment_tpm is None:
|
||||
_deployment_tpm = _deployment.get("model_info", {}).get("tpm")
|
||||
if _deployment_tpm is None:
|
||||
_deployment_tpm = float("inf")
|
||||
|
||||
_deployment_rpm = (
|
||||
_deployment.get("rpm", None)
|
||||
or _deployment.get("litellm_params", {}).get("rpm", None)
|
||||
or _deployment.get("model_info", {}).get("rpm", None)
|
||||
or float("inf")
|
||||
)
|
||||
_deployment_rpm = None
|
||||
if _deployment_rpm is None:
|
||||
_deployment_rpm = _deployment.get("rpm")
|
||||
if _deployment_rpm is None:
|
||||
_deployment_rpm = _deployment.get("litellm_params", {}).get("rpm")
|
||||
if _deployment_rpm is None:
|
||||
_deployment_rpm = _deployment.get("model_info", {}).get("rpm")
|
||||
if _deployment_rpm is None:
|
||||
_deployment_rpm = float("inf")
|
||||
|
||||
if item_tpm == 0:
|
||||
deployment = _deployment
|
||||
break
|
||||
elif item_tpm + input_tokens > _deployment_tpm:
|
||||
if item_tpm + input_tokens > _deployment_tpm:
|
||||
continue
|
||||
elif (rpm_dict is not None and item in rpm_dict) and (
|
||||
rpm_dict[item] + 1 > _deployment_rpm
|
||||
|
|
@ -210,5 +208,5 @@ class LowestTPMLoggingHandler(CustomLogger):
|
|||
elif item_tpm < lowest_tpm:
|
||||
lowest_tpm = item_tpm
|
||||
deployment = _deployment
|
||||
verbose_router_logger.info("returning picked lowest tpm/rpm deployment.")
|
||||
print_verbose("returning picked lowest tpm/rpm deployment.")
|
||||
return deployment
|
||||
|
|
|
|||
|
|
@ -264,7 +264,7 @@ def test_router_skip_rate_limited_deployments():
|
|||
end_time=end_time,
|
||||
)
|
||||
|
||||
## CHECK WHAT'S SELECTED ## - should skip 2, and pick 1
|
||||
## CHECK WHAT'S SELECTED ##
|
||||
# print(router.lowesttpm_logger.get_available_deployments(model_group="azure-model"))
|
||||
try:
|
||||
router.get_available_deployment(
|
||||
|
|
@ -273,7 +273,41 @@ def test_router_skip_rate_limited_deployments():
|
|||
)
|
||||
pytest.fail(f"Should have raised No Models Available error")
|
||||
except Exception as e:
|
||||
pass
|
||||
print(f"An exception occurred! {str(e)}")
|
||||
|
||||
|
||||
def test_single_deployment_tpm_zero():
|
||||
import litellm
|
||||
import os
|
||||
from datetime import datetime
|
||||
|
||||
model_list = [
|
||||
{
|
||||
"model_name": "gpt-3.5-turbo",
|
||||
"litellm_params": {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"api_key": os.getenv("OPENAI_API_KEY"),
|
||||
"tpm": 0,
|
||||
},
|
||||
}
|
||||
]
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=model_list,
|
||||
routing_strategy="usage-based-routing",
|
||||
cache_responses=True,
|
||||
)
|
||||
|
||||
model = "gpt-3.5-turbo"
|
||||
messages = [{"content": "Hello, how are you?", "role": "user"}]
|
||||
try:
|
||||
router.get_available_deployment(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": "Hey, how's it going?"}],
|
||||
)
|
||||
pytest.fail(f"Should have raised No Models Available error")
|
||||
except Exception as e:
|
||||
print(f"it worked - {str(e)}! \n{traceback.format_exc()}")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue