From 47ca223d0bef26982cc26ac8c16560043389d563 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Thu, 28 Mar 2024 14:51:31 -0700 Subject: [PATCH] fix(lowest_tpm_rpm_routing.py): fix base case where max tpm/rpm is 0 --- litellm/router_strategy/lowest_tpm_rpm.py | 66 +++++++++++------------ litellm/tests/test_tpm_rpm_routing.py | 38 ++++++++++++- 2 files changed, 68 insertions(+), 36 deletions(-) diff --git a/litellm/router_strategy/lowest_tpm_rpm.py b/litellm/router_strategy/lowest_tpm_rpm.py index d2bf7bdb5a2..1e1e6df98e4 100644 --- a/litellm/router_strategy/lowest_tpm_rpm.py +++ b/litellm/router_strategy/lowest_tpm_rpm.py @@ -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 diff --git a/litellm/tests/test_tpm_rpm_routing.py b/litellm/tests/test_tpm_rpm_routing.py index 588f3fab699..e7ada5eb840 100644 --- a/litellm/tests/test_tpm_rpm_routing.py +++ b/litellm/tests/test_tpm_rpm_routing.py @@ -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