mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
add LiteLLMRouterEncoder
This commit is contained in:
parent
33e1b3b651
commit
328ab65888
2 changed files with 117 additions and 71 deletions
|
|
@ -1,71 +0,0 @@
|
|||
"""
|
||||
Auto-Routing Strategy that works with a Semantic Router Config
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
from typing import List, Literal, Optional, Union
|
||||
|
||||
from semantic_router.schema import RouteChoice
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy.proxy_server import DualCache, UserAPIKeyAuth
|
||||
|
||||
|
||||
class AutoRouter(CustomLogger):
|
||||
DEFAULT_AUTO_SYNC_VALUE = "local"
|
||||
def __init__(
|
||||
self,
|
||||
router_config_path: str,
|
||||
default_model: str,
|
||||
embedding_model: str,
|
||||
embedding_model_api_key: str,
|
||||
):
|
||||
from semantic_router import Route
|
||||
from semantic_router.encoders import LiteLLMEncoder
|
||||
from semantic_router.routers import SemanticRouter
|
||||
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=LiteLLMEncoder(
|
||||
name=embedding_model,
|
||||
api_key=embedding_model_api_key,
|
||||
),
|
||||
auto_sync=self.auto_sync_value,
|
||||
)
|
||||
self.default_model = default_model
|
||||
pass
|
||||
|
||||
|
||||
async def async_pre_routing_hook(
|
||||
self,
|
||||
data: dict, call_type: Literal[
|
||||
"completion",
|
||||
"text_completion",
|
||||
"embeddings",
|
||||
"image_generation",
|
||||
"moderation",
|
||||
"audio_transcription",
|
||||
]):
|
||||
from semantic_router.schema import RouteChoice
|
||||
|
||||
#self.routelayer.to_json("./config/router.json")
|
||||
# If the call type is embeddings, do not modify the model.
|
||||
if call_type in ["embeddings", "image_generation"]:
|
||||
#print("Call type is 'embeddings', using default behavior.")
|
||||
return data # Return without modifying the data
|
||||
|
||||
msg = data['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
|
||||
|
||||
117
litellm/router_strategy/auto_router/litellm_encoder.py
Normal file
117
litellm/router_strategy/auto_router/litellm_encoder.py
Normal file
|
|
@ -0,0 +1,117 @@
|
|||
from typing import TYPE_CHECKING, Any, Union
|
||||
|
||||
from semantic_router.encoders import DenseEncoder
|
||||
from semantic_router.encoders.base import AsymmetricDenseMixin
|
||||
|
||||
import litellm
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.router import Router
|
||||
else:
|
||||
Router = Any
|
||||
|
||||
|
||||
def litellm_to_list(embeds: litellm.EmbeddingResponse) -> list[list[float]]:
|
||||
"""Convert a LiteLLM embedding response to a list of embeddings.
|
||||
|
||||
:param embeds: The LiteLLM embedding response.
|
||||
:return: A list of embeddings.
|
||||
"""
|
||||
if (
|
||||
not embeds
|
||||
or not isinstance(embeds, litellm.EmbeddingResponse)
|
||||
or not embeds.data
|
||||
):
|
||||
raise ValueError("No embeddings found in LiteLLM embedding response.")
|
||||
return [x["embedding"] for x in embeds.data]
|
||||
|
||||
|
||||
class LiteLLMRouterEncoder(DenseEncoder, AsymmetricDenseMixin):
|
||||
"""LiteLLM encoder class for generating embeddings using LiteLLM.
|
||||
|
||||
The LiteLLMRouterEncoder class is a subclass of DenseEncoder and utilizes the LiteLLM Router SDK
|
||||
to generate embeddings for given documents. It supports all encoders supported by LiteLLM
|
||||
and supports customization of the score threshold for filtering or processing the embeddings.
|
||||
"""
|
||||
|
||||
type: str = "internal_litellm_router"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
litellm_router_instance: "Router",
|
||||
model_name: str,
|
||||
score_threshold: Union[float, None] = None,
|
||||
):
|
||||
"""Initialize the LiteLLMEncoder.
|
||||
|
||||
:param litellm_router_instance: The instance of the LiteLLM Router.
|
||||
:type litellm_router_instance: Router
|
||||
:param model_name: The name of the embedding model to use. Must use LiteLLM naming
|
||||
convention (e.g. "openai/text-embedding-3-small" or "mistral/mistral-embed").
|
||||
:type model_name: str
|
||||
:param score_threshold: The score threshold for the embeddings.
|
||||
:type score_threshold: float
|
||||
"""
|
||||
self.litellm_router_instance: "Router" = litellm_router_instance
|
||||
super().__init__(
|
||||
name=model_name,
|
||||
score_threshold=score_threshold if score_threshold is not None else 0.3,
|
||||
)
|
||||
|
||||
def __call__(self, docs: list[Any], **kwargs) -> list[list[float]]:
|
||||
"""Encode a list of text documents into embeddings using LiteLLM.
|
||||
|
||||
:param docs: List of text documents to encode.
|
||||
:return: List of embeddings for each document."""
|
||||
return self.encode_queries(docs, **kwargs)
|
||||
|
||||
async def acall(self, docs: list[Any], **kwargs) -> list[list[float]]:
|
||||
"""Encode a list of documents into embeddings using LiteLLM asynchronously.
|
||||
|
||||
:param docs: List of documents to encode.
|
||||
:return: List of embeddings for each document."""
|
||||
return await self.aencode_queries(docs, **kwargs)
|
||||
|
||||
def encode_queries(self, docs: list[str], **kwargs) -> list[list[float]]:
|
||||
try:
|
||||
embeds = self.litellm_router_instance.embedding(
|
||||
input=docs, model=f"{self.type}/{self.name}", **kwargs
|
||||
)
|
||||
return litellm_to_list(embeds)
|
||||
except Exception as e:
|
||||
raise ValueError(
|
||||
f"{self.type.capitalize()} API call failed. Error: {e}"
|
||||
) from e
|
||||
|
||||
def encode_documents(self, docs: list[str], **kwargs) -> list[list[float]]:
|
||||
try:
|
||||
embeds = self.litellm_router_instance.embedding(
|
||||
input=docs, model=f"{self.type}/{self.name}", **kwargs
|
||||
)
|
||||
return litellm_to_list(embeds)
|
||||
except Exception as e:
|
||||
raise ValueError(
|
||||
f"{self.type.capitalize()} API call failed. Error: {e}"
|
||||
) from e
|
||||
|
||||
async def aencode_queries(self, docs: list[str], **kwargs) -> list[list[float]]:
|
||||
try:
|
||||
embeds = await self.litellm_router_instance.aembedding(
|
||||
input=docs, model=f"{self.type}/{self.name}", **kwargs
|
||||
)
|
||||
return litellm_to_list(embeds)
|
||||
except Exception as e:
|
||||
raise ValueError(
|
||||
f"{self.type.capitalize()} API call failed. Error: {e}"
|
||||
) from e
|
||||
|
||||
async def aencode_documents(self, docs: list[str], **kwargs) -> list[list[float]]:
|
||||
try:
|
||||
embeds = await self.litellm_router_instance.aembedding(
|
||||
input=docs, model=f"{self.type}/{self.name}", **kwargs
|
||||
)
|
||||
return litellm_to_list(embeds)
|
||||
except Exception as e:
|
||||
raise ValueError(
|
||||
f"{self.type.capitalize()} API call failed. Error: {e}"
|
||||
) from e
|
||||
Loading…
Add table
Reference in a new issue