mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
feat: add embeddings to CustomLLM (#10980)
* feat: add embeddings to CustomLLM * feat: add aembedding to custom llm
This commit is contained in:
parent
e2d147102d
commit
db4183715a
3 changed files with 130 additions and 2 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue