mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
add AutoRouter
This commit is contained in:
parent
48faad3bdd
commit
d5c7afd834
2 changed files with 131 additions and 1 deletions
|
|
@ -165,9 +165,12 @@ from .router_utils.pattern_match_deployments import PatternMatchRouter
|
|||
if TYPE_CHECKING:
|
||||
from opentelemetry.trace import Span as _Span
|
||||
|
||||
from litellm.router_strategy.auto_router.auto_router import AutoRouter
|
||||
|
||||
Span = Union[_Span, Any]
|
||||
else:
|
||||
Span = Any
|
||||
AutoRouter = Any
|
||||
|
||||
|
||||
class RoutingArgs(enum.Enum):
|
||||
|
|
@ -517,6 +520,7 @@ class Router:
|
|||
routing_strategy_args=routing_strategy_args,
|
||||
)
|
||||
self.access_groups = None
|
||||
self.auto_routers: Dict[str, "AutoRouter"] = {}
|
||||
## USAGE TRACKING ##
|
||||
if isinstance(litellm._async_success_callback, list):
|
||||
litellm.logging_callback_manager.add_litellm_async_success_callback(
|
||||
|
|
@ -4678,10 +4682,11 @@ class Router:
|
|||
- None: If the deployment is not active for the current environment (if 'supported_environments' is set in litellm_params)
|
||||
"""
|
||||
try:
|
||||
litellm_params: LiteLLM_Params = LiteLLM_Params(**_litellm_params)
|
||||
deployment = Deployment(
|
||||
**deployment_info,
|
||||
model_name=_model_name,
|
||||
litellm_params=LiteLLM_Params(**_litellm_params),
|
||||
litellm_params=litellm_params,
|
||||
model_info=_model_info,
|
||||
)
|
||||
for field in CustomPricingLiteLLMParams.model_fields.keys():
|
||||
|
|
@ -4696,6 +4701,13 @@ class Router:
|
|||
model_id: _model_info,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
#########################################################
|
||||
# Check if this is an auto-router deployment
|
||||
#########################################################
|
||||
if self._is_auto_router_deployment(litellm_params=litellm_params):
|
||||
self.init_auto_router_deployment(deployment=deployment)
|
||||
|
||||
## OLD MODEL REGISTRATION ## Kept to prevent breaking changes
|
||||
_model_name = deployment.litellm_params.model
|
||||
|
|
@ -4734,6 +4746,46 @@ class Router:
|
|||
return None
|
||||
else:
|
||||
raise e
|
||||
|
||||
def _is_auto_router_deployment(self, litellm_params: LiteLLM_Params) -> bool:
|
||||
"""
|
||||
Check if the deployment is an auto-router deployment.
|
||||
|
||||
Returns True if the litellm_params model starts with "auto_router/"
|
||||
"""
|
||||
if litellm_params.model.startswith("auto_router/"):
|
||||
return True
|
||||
return False
|
||||
|
||||
def init_auto_router_deployment(self, deployment: Deployment):
|
||||
"""
|
||||
Initialize the auto-router deployment.
|
||||
|
||||
This will initialize the auto-router and add it to the auto-routers dictionary.
|
||||
"""
|
||||
from litellm.router_strategy.auto_router.auto_router import AutoRouter
|
||||
router_config_path: Optional[str] = deployment.litellm_params.auto_router_config_path
|
||||
if router_config_path is None:
|
||||
raise ValueError("auto_router_config_path is required for auto-router deployments. Please set it in the litellm_params")
|
||||
|
||||
default_model: Optional[str] = deployment.litellm_params.auto_router_default_model
|
||||
if default_model is None:
|
||||
raise ValueError("auto_router_default_model is required for auto-router deployments. Please set it in the litellm_params")
|
||||
|
||||
embedding_model: Optional[str] = deployment.litellm_params.auto_router_embedding_model
|
||||
if embedding_model is None:
|
||||
raise ValueError("auto_router_embedding_model is required for auto-router deployments. Please set it in the litellm_params")
|
||||
|
||||
autor_router: AutoRouter = AutoRouter(
|
||||
model_name=deployment.model_name,
|
||||
router_config_path=router_config_path,
|
||||
default_model=default_model,
|
||||
embedding_model=embedding_model,
|
||||
litellm_router_instance=self,
|
||||
)
|
||||
if deployment.model_name in self.auto_routers:
|
||||
raise ValueError(f"Auto-router deployment {deployment.model_name} already exists. Please use a different model name.")
|
||||
self.auto_routers[deployment.model_name] = autor_router
|
||||
|
||||
def deployment_is_active_for_environment(self, deployment: Deployment) -> bool:
|
||||
"""
|
||||
|
|
|
|||
78
litellm/router_strategy/auto_router/auto_router.py
Normal file
78
litellm/router_strategy/auto_router/auto_router.py
Normal file
|
|
@ -0,0 +1,78 @@
|
|||
"""
|
||||
Auto-Routing Strategy that works with a Semantic Router Config
|
||||
"""
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
||||
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.router import Router
|
||||
else:
|
||||
Router = Any
|
||||
|
||||
class AutoRouter(CustomLogger):
|
||||
DEFAULT_AUTO_SYNC_VALUE = "local"
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
router_config_path: str,
|
||||
default_model: str,
|
||||
embedding_model: str,
|
||||
litellm_router_instance: "Router",
|
||||
):
|
||||
"""
|
||||
Auto-Router class that uses a semantic router to route requests to the appropriate model.
|
||||
|
||||
Args:
|
||||
model_name: The name of the model to use for the auto-router. eg. if model = "auto-router1" then us this router.
|
||||
router_config_path: The path to the router config file.
|
||||
default_model: The default model to use if no route is found.
|
||||
embedding_model: The embedding model to use for the auto-router.
|
||||
litellm_router_instance: The instance of the LiteLLM Router.
|
||||
"""
|
||||
from semantic_router import Route
|
||||
from semantic_router.routers import SemanticRouter
|
||||
|
||||
from litellm.router_strategy.auto_router.litellm_encoder import (
|
||||
LiteLLMRouterEncoder,
|
||||
)
|
||||
self.router_config_path = router_config_path
|
||||
self.auto_sync_value = self.DEFAULT_AUTO_SYNC_VALUE
|
||||
loaded_router: SemanticRouter = SemanticRouter.from_json(self.router_config_path)
|
||||
self.routelayer: SemanticRouter = SemanticRouter(
|
||||
routes=loaded_router.routes,
|
||||
encoder=LiteLLMRouterEncoder(
|
||||
litellm_router_instance=litellm_router_instance,
|
||||
model_name=embedding_model,
|
||||
),
|
||||
auto_sync=self.auto_sync_value,
|
||||
)
|
||||
self.default_model = default_model
|
||||
pass
|
||||
|
||||
|
||||
async def async_pre_routing_hook(
|
||||
self,
|
||||
model: str,
|
||||
request_kwargs: Dict,
|
||||
messages: Optional[List[Dict[str, str]]] = None,
|
||||
input: Optional[Union[str, List]] = None,
|
||||
specific_deployment: Optional[bool] = False,
|
||||
):
|
||||
pass
|
||||
# from semantic_router.schema import RouteChoice
|
||||
# if request_kwargs.get("model") == "embeddings":
|
||||
# return request_kwargs
|
||||
# if messages is None:
|
||||
|
||||
# msg = messages[-1]['content']
|
||||
# route_choice: Optional[Union[RouteChoice, List[RouteChoice]]] = self.routelayer(text=msg)
|
||||
# if isinstance(route_choice, RouteChoice):
|
||||
# data["model"] = route_choice.name
|
||||
# elif isinstance(route_choice, list):
|
||||
# data["model"] = route_choice[0].name
|
||||
# else:
|
||||
# data["model"] = self.default_model
|
||||
|
||||
# return data
|
||||
|
||||
Loading…
Add table
Reference in a new issue