feat: add embeddings to CustomLLM (#10980)

* feat: add embeddings to CustomLLM

* feat: add aembedding to custom llm
This commit is contained in:
Tornike Gurgenidze 2025-05-23 09:55:46 +04:00 • committed by GitHub
parent e2d147102d
commit db4183715a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 130 additions and 2 deletions

View file

@ -14,7 +14,7 @@ import httpx
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.types.utils import GenericStreamingChunk
from litellm.utils import ImageResponse, ModelResponse
from litellm.utils import ImageResponse, ModelResponse, EmbeddingResponse
from .base import BaseLLM
@ -152,6 +152,30 @@ class CustomLLM(BaseLLM):
) -> ImageResponse:
raise CustomLLMError(status_code=500, message="Not implemented yet!")
def embedding(
self,
model: str,
input: list,
model_response: EmbeddingResponse,
print_verbose: Callable,
logging_obj: Any,
optional_params: dict,
litellm_params=None,
) -> EmbeddingResponse:
raise CustomLLMError(status_code=500, message="Not implemented yet!")
async def aembedding(
self,
model: str,
input: list,
model_response: EmbeddingResponse,
print_verbose: Callable,
logging_obj: Any,
optional_params: dict,
litellm_params=None,
) -> EmbeddingResponse:
raise CustomLLMError(status_code=500, message="Not implemented yet!")
def custom_chat_llm_router(
async_fn: bool, stream: Optional[bool], custom_llm: CustomLLM

View file

@ -4027,6 +4027,30 @@ def embedding( # noqa: PLR0915
client=client,
aembedding=aembedding,
)
elif (
custom_llm_provider in litellm._custom_providers
):
custom_handler: Optional[CustomLLM] = None
for item in litellm.custom_provider_map:
if item["provider"] == custom_llm_provider:
custom_handler = item["custom_handler"]
if custom_handler is None:
raise LiteLLMUnknownProvider(
model=model, custom_llm_provider=custom_llm_provider
)
handler_fn = custom_handler.embedding if not aembedding else custom_handler.aembedding
response = handler_fn(
model=model,
input=input,
logging_obj=logging,
optional_params=optional_params,
model_response=EmbeddingResponse(),
print_verbose=print_verbose,
litellm_params=litellm_params
)
else:
raise LiteLLMUnknownProvider(
model=model, custom_llm_provider=custom_llm_provider

View file

@ -44,7 +44,7 @@ from litellm import (
image_generation,
)
from litellm.utils import ModelResponseIterator
from litellm.types.utils import ImageResponse, ImageObject
from litellm.types.utils import ImageResponse, ImageObject, EmbeddingResponse
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
@ -257,6 +257,53 @@ class MyCustomLLM(CustomLLM):
response_ms=1000,
)
def embedding(
self,
model: str,
input: list,
model_response: EmbeddingResponse,
print_verbose: Callable,
logging_obj: Any,
optional_params: dict,
litellm_params=None,
aembedding=None,
) -> EmbeddingResponse:
model_response.model = model
model_response.data = [
{
"object": "embedding",
"embedding": [0.1, 0.2, 0.3],
"index": i,
}
for i, _ in enumerate(input)
]
return model_response
async def aembedding(
self,
model: str,
input: list,
model_response: EmbeddingResponse,
print_verbose: Callable,
logging_obj: Any,
optional_params: dict,
litellm_params=None,
) -> EmbeddingResponse:
model_response.model = model
model_response.data = [
{
"object": "embedding",
"embedding": [0.1, 0.2, 0.3],
"index": i,
}
for i, _ in enumerate(input)
]
return model_response
def test_get_llm_provider():
""""""
@ -452,3 +499,36 @@ def test_get_supported_openai_params():
response = get_supported_openai_params(model="my-custom-llm/my-fake-model")
assert response is not None
def test_simple_embedding():
my_custom_llm = MyCustomLLM()
litellm.custom_provider_map = [
{"provider": "custom_llm", "custom_handler": my_custom_llm}
]
resp = litellm.embedding(
model="custom_llm/my-fake-model",
input=["good morning from litellm", "good night from litellm"]
)
assert resp.data[1] == {
"object": "embedding",
"embedding": [0.1, 0.2, 0.3],
"index": 1,
}
@pytest.mark.asyncio
async def test_simple_aembedding():
my_custom_llm = MyCustomLLM()
litellm.custom_provider_map = [
{"provider": "custom_llm", "custom_handler": my_custom_llm}
]
resp = await litellm.aembedding(
model="custom_llm/my-fake-model",
input=["good morning from litellm", "good night from litellm"]
)
assert resp.data[1] == {
"object": "embedding",
"embedding": [0.1, 0.2, 0.3],
"index": 1,
}