feat(wren-ai-service): add litellm embedder (#1247)

This commit is contained in:
Chih-Yu Yeh
2025-02-03 14:10:37 +08:00
committed by GitHub
parent 957765d36b
commit b3543c9e48
12 changed files with 896 additions and 645 deletions
+1 -1
View File
@@ -42,7 +42,7 @@ jobs:
python-version-file: ./wren-ai-service/pyproject.toml
cache: "poetry"
- name: Install the project dependencies
run: poetry install
run: poetry install --without eval
- name: Install Just
uses: extractions/setup-just@v2
with:
@@ -1,12 +1,11 @@
type: llm
provider: litellm_llm # litellm supports Azure through its provider
provider: litellm_llm
timeout: 120
models:
# put AZURE_API_KEY=<your_api_key> in ~/.wrenai/.env
- model: azure/gpt-4 # Your Azure deployment name, put 'azure/' before deployment name
api_base: https://endpoint.openai.azure.com/ #Replace with your custom Azure endpoint
api_key_name: LLM_AZURE_OPENAI_API_KEY
api_base: https://endpoint.openai.azure.com # Replace with your custom Azure endpoint
api_version: 2024-02-15-preview
kwargs:
temperature: 0
n: 1
@@ -16,14 +15,13 @@ models:
---
type: embedder
provider: azure_openai_embedder
provider: litellm_embedder
models:
- model: text-embedding-ada-002 # Your Azure deployment name
# Must match model output check for your model
api_base: https://endpoint.openai.azure.com/ # Replace with your custom Azure endpoint
api_version: 2023-05-15 # Your Azure deployment name
timeout: 300
# put AZURE_API_KEY=<your_api_key> in ~/.wrenai/.env
- model: azure/text-embedding-ada-002 # Your Azure deployment name, put 'azure/' before deployment name
api_base: https://endpoint.openai.azure.com # Replace with your custom Azure endpoint
api_version: 2023-05-15
timeout: 300
---
type: engine
@@ -32,36 +30,34 @@ endpoint: http://wren-ui:3000
---
type: document_store
#name: qdrant
provider: qdrant
location: http://qdrant:6333 # Donot set the QDRANT_API_KEY if you are using the qdrant from docker
embedding_model_dim: 1536 # Must match model dimension from embedder
timeout: 120
recreate_index: true
# For each pipe line component
# Replace llm with Azure deployed LLM model
# Replace Embeddings with Azure deployed Embedding model
---
# please change the llm and embedder names to the ones you want to use
# the format of llm and embedder should be <provider>.<model_name> such as litellm_llm.gpt-4o-2024-08-06
# the pipes may be not the latest version, please refer to the latest version: https://raw.githubusercontent.com/canner/WrenAI/<WRENAI_VERSION_NUMBER>/docker/config.example.yaml
type: pipeline
pipes:
- name: db_schema_indexing
embedder: azure_openai_embedder.text-embedding-ada-002
embedder: litellm_embedder.azure/text-embedding-ada-002
document_store: qdrant # Match document_store name
llm: litellm_llm.azure/gpt-4
- name: historical_question_indexing
embedder: azure_openai_embedder.text-embedding-ada-002
embedder: litellm_embedder.azure/text-embedding-ada-002
document_store: qdrant
- name: table_description_indexing
embedder: azure_openai_embedder.text-embedding-ada-002
embedder: litellm_embedder.azure/text-embedding-ada-002
document_store: qdrant
- name: db_schema_retrieval
llm: litellm_llm.azure/gpt-4
embedder: azure_openai_embedder.text-embedding-ada-002
embedder: litellm_embedder.azure/text-embedding-ada-002
document_store: qdrant
- name: historical_question_retrieval
embedder: azure_openai_embedder.text-embedding-ada-002
embedder: litellm_embedder.azure/text-embedding-ada-002
document_store: qdrant
- name: sql_generation
llm: litellm_llm.azure/gpt-4
@@ -97,20 +93,20 @@ pipes:
llm: litellm_llm.azure/gpt-4
- name: intent_classification
llm: litellm_llm.azure/gpt-4
embedder: azure_openai_embedder.text-embedding-ada-002
embedder: litellm_embedder.azure/text-embedding-ada-002
document_store: qdrant
- name: data_assistance
llm: litellm_llm.azure/gpt-4
- name: sql_pairs_preparation
document_store: qdrant
embedder: azure_openai_embedder.text-embedding-ada-002
embedder: litellm_embedder.azure/text-embedding-ada-002
llm: litellm_llm.azure/gpt-4
- name: sql_pairs_deletion
document_store: qdrant
embedder: azure_openai_embedder.text-embedding-ada-002
embedder: litellm_embedder.azure/text-embedding-ada-002
- name: sql_pairs_retrieval
document_store: qdrant
embedder: azure_openai_embedder.text-embedding-ada-002
embedder: litellm_embedder.azure/text-embedding-ada-002
llm: litellm_llm.azure/gpt-4
- name: preprocess_sql_data
llm: litellm_llm.azure/gpt-4
@@ -122,12 +118,12 @@ pipes:
llm: litellm_llm.azure/gpt-4
- name: sql_pairs_indexing
document_store: qdrant
embedder: azure_openai_embedder.text-embedding-ada-002
embedder: litellm_embedder.azure/text-embedding-ada-002
- name: sql_generation_reasoning
llm: litellm_llm.azure/gpt-4
- name: question_recommendation_db_schema_retrieval
llm: litellm_llm.azure/gpt-4
embedder: azure_openai_embedder.text-embedding-ada-002
embedder: litellm_embedder.azure/text-embedding-ada-002
document_store: qdrant
- name: question_recommendation_sql_generation
llm: litellm_llm.azure/gpt-4
@@ -30,12 +30,13 @@ models:
---
type: embedder
provider: openai_embedder
provider: litellm_embedder
models:
# find EMBEDDER_OPENAI_API_KEY and fill in value of api key in ~/.wrenai/.env
- model: text-embedding-3-large # put your openai compatible embedder model name here
url: https://api.openai.com/v1 # change this according to your openai compatible embedder model
timeout: 120
# define OPENAI_API_KEY=<api_key> in ~/.wrenai/.env if you are using openai embedding model
# please refer to LiteLLM documentation for more details: https://docs.litellm.ai/docs/providers
- model: text-embedding-3-large # put your embedding model name here, if it is not openai embedding model, should be <provider>/<model_name>
api_base: https://api.openai.com/v1 # change this according to your embedding model
timeout: 120
---
type: engine
@@ -57,20 +58,20 @@ recreate_index: false
type: pipeline
pipes:
- name: db_schema_indexing
embedder: openai_embedder.text-embedding-3-large
embedder: litellm_embedder.text-embedding-3-large
document_store: qdrant
- name: historical_question_indexing
embedder: openai_embedder.text-embedding-3-large
embedder: litellm_embedder.text-embedding-3-large
document_store: qdrant
- name: table_description_indexing
embedder: openai_embedder.text-embedding-3-large
embedder: litellm_embedder.text-embedding-3-large
document_store: qdrant
- name: db_schema_retrieval
llm: litellm_llm.deepseek/deepseek-coder
embedder: openai_embedder.text-embedding-3-large
embedder: litellm_embedder.text-embedding-3-large
document_store: qdrant
- name: historical_question_retrieval
embedder: openai_embedder.text-embedding-3-large
embedder: litellm_embedder.text-embedding-3-large
document_store: qdrant
- name: sql_generation
llm: litellm_llm.deepseek/deepseek-coder
@@ -106,7 +107,7 @@ pipes:
llm: litellm_llm.deepseek/deepseek-coder
- name: question_recommendation_db_schema_retrieval
llm: litellm_llm.deepseek/deepseek-coder
embedder: openai_embedder.text-embedding-3-large
embedder: litellm_embedder.text-embedding-3-large
document_store: qdrant
- name: question_recommendation_sql_generation
llm: litellm_llm.deepseek/deepseek-coder
@@ -117,19 +118,19 @@ pipes:
llm: litellm_llm.deepseek/deepseek-coder
- name: intent_classification
llm: litellm_llm.deepseek/deepseek-coder
embedder: openai_embedder.text-embedding-3-large
embedder: litellm_embedder.text-embedding-3-large
document_store: qdrant
- name: data_assistance
llm: litellm_llm.deepseek/deepseek-chat
- name: sql_pairs_indexing
document_store: qdrant
embedder: openai_embedder.text-embedding-3-large
embedder: litellm_embedder.text-embedding-3-large
- name: sql_pairs_deletion
document_store: qdrant
embedder: openai_embedder.text-embedding-3-large
embedder: litellm_embedder.text-embedding-3-large
- name: sql_pairs_retrieval
document_store: qdrant
embedder: openai_embedder.text-embedding-3-large
embedder: litellm_embedder.text-embedding-3-large
llm: litellm_llm.deepseek/deepseek-coder
- name: preprocess_sql_data
llm: litellm_llm.deepseek/deepseek-coder
@@ -13,12 +13,12 @@ models:
---
type: embedder
provider: openai_embedder
provider: litellm_embedder
models:
# find EMBEDDER_OPENAI_API_KEY and fill in value of api key in ~/.wrenai/.env
- model: text-embedding-004 # put your openai compatible embedder model name here
url: https://generativelanguage.googleapis.com/v1beta/openai # change this according to your openai compatible embedder model
timeout: 120
# put GEMINI_API_KEY=<your_api_key> in ~/.wrenai/.env
- model: gemini/text-embedding-004 # gemini/<gemini_model_name>
api_base: https://generativelanguage.googleapis.com/v1beta/openai # change this according to your embedding model
timeout: 120
---
type: engine
@@ -40,20 +40,20 @@ recreate_index: false
type: pipeline
pipes:
- name: db_schema_indexing
embedder: openai_embedder.text-embedding-004
embedder: litellm_embedder.text-embedding-004
document_store: qdrant
- name: historical_question_indexing
embedder: openai_embedder.text-embedding-004
embedder: litellm_embedder.text-embedding-004
document_store: qdrant
- name: table_description_indexing
embedder: openai_embedder.text-embedding-004
embedder: litellm_embedder.text-embedding-004
document_store: qdrant
- name: db_schema_retrieval
llm: litellm_llm.gemini/gemini-2.0-flash-exp
embedder: openai_embedder.text-embedding-004
embedder: litellm_embedder.text-embedding-004
document_store: qdrant
- name: historical_question_retrieval
embedder: openai_embedder.text-embedding-004
embedder: litellm_embedder.text-embedding-004
document_store: qdrant
- name: sql_generation
llm: litellm_llm.gemini/gemini-2.0-flash-exp
@@ -89,7 +89,7 @@ pipes:
llm: litellm_llm.gemini/gemini-2.0-flash-exp
- name: question_recommendation_db_schema_retrieval
llm: litellm_llm.gemini/gemini-2.0-flash-exp
embedder: openai_embedder.text-embedding-004
embedder: litellm_embedder.text-embedding-004
document_store: qdrant
- name: question_recommendation_sql_generation
llm: litellm_llm.gemini/gemini-2.0-flash-exp
@@ -100,19 +100,19 @@ pipes:
llm: litellm_llm.gemini/gemini-2.0-flash-exp
- name: intent_classification
llm: litellm_llm.gemini/gemini-2.0-flash-exp
embedder: openai_embedder.text-embedding-004
embedder: litellm_embedder.text-embedding-004
document_store: qdrant
- name: data_assistance
llm: litellm_llm.gemini/gemini-2.0-flash-exp
- name: sql_pairs_indexing
document_store: qdrant
embedder: openai_embedder.text-embedding-004
embedder: litellm_embedder.text-embedding-004
- name: sql_pairs_deletion
document_store: qdrant
embedder: openai_embedder.text-embedding-004
embedder: litellm_embedder.text-embedding-004
- name: sql_pairs_retrieval
document_store: qdrant
embedder: openai_embedder.text-embedding-004
embedder: litellm_embedder.text-embedding-004
llm: litellm_llm.gemini/gemini-2.0-flash-exp
- name: preprocess_sql_data
llm: litellm_llm.gemini/gemini-2.0-flash-exp
@@ -14,12 +14,13 @@ models:
---
type: embedder
provider: openai_embedder
provider: litellm_embedder
models:
# find EMBEDDER_OPENAI_API_KEY and fill in value of api key in ~/.wrenai/.env
- model: text-embedding-3-large # put your openai compatible embedder model name here
url: https://api.openai.com/v1 # change this according to your openai compatible embedder model
timeout: 120
# define OPENAI_API_KEY=<api_key> in ~/.wrenai/.env if you are using openai embedding model
# please refer to LiteLLM documentation for more details: https://docs.litellm.ai/docs/providers
- model: text-embedding-3-large # put your embedding model name here, if it is not openai embedding model, should be <provider>/<model_name>
api_base: https://api.openai.com/v1 # change this according to your embedding model
timeout: 120
---
type: engine
@@ -41,20 +42,20 @@ recreate_index: false
type: pipeline
pipes:
- name: db_schema_indexing
embedder: openai_embedder.text-embedding-3-large
embedder: litellm_embedder.text-embedding-3-large
document_store: qdrant
- name: historical_question_indexing
embedder: openai_embedder.text-embedding-3-large
embedder: litellm_embedder.text-embedding-3-large
document_store: qdrant
- name: table_description_indexing
embedder: openai_embedder.text-embedding-3-large
embedder: litellm_embedder.text-embedding-3-large
document_store: qdrant
- name: db_schema_retrieval
llm: litellm_llm.groq/llama-3.3-70b-specdec
embedder: openai_embedder.text-embedding-3-large
embedder: litellm_embedder.text-embedding-3-large
document_store: qdrant
- name: historical_question_retrieval
embedder: openai_embedder.text-embedding-3-large
embedder: litellm_embedder.text-embedding-3-large
document_store: qdrant
- name: sql_generation
llm: litellm_llm.groq/llama-3.3-70b-specdec
@@ -90,7 +91,7 @@ pipes:
llm: litellm_llm.groq/llama-3.3-70b-specdec
- name: question_recommendation_db_schema_retrieval
llm: litellm_llm.groq/llama-3.3-70b-specdec
embedder: openai_embedder.text-embedding-3-large
embedder: litellm_embedder.text-embedding-3-large
document_store: qdrant
- name: question_recommendation_sql_generation
llm: litellm_llm.groq/llama-3.3-70b-specdec
@@ -101,19 +102,19 @@ pipes:
llm: litellm_llm.groq/llama-3.3-70b-specdec
- name: intent_classification
llm: litellm_llm.groq/llama-3.3-70b-specdec
embedder: openai_embedder.text-embedding-3-large
embedder: litellm_embedder.text-embedding-3-large
document_store: qdrant
- name: data_assistance
llm: litellm_llm.groq/llama-3.3-70b-specdec
- name: sql_pairs_indexing
document_store: qdrant
embedder: openai_embedder.text-embedding-3-large
embedder: litellm_embedder.text-embedding-3-large
- name: sql_pairs_deletion
document_store: qdrant
embedder: openai_embedder.text-embedding-3-large
embedder: litellm_embedder.text-embedding-3-large
- name: sql_pairs_retrieval
document_store: qdrant
embedder: openai_embedder.text-embedding-3-large
embedder: litellm_embedder.text-embedding-3-large
llm: litellm_llm.groq/llama-3.3-70b-specdec
- name: preprocess_sql_data
llm: litellm_llm.groq/llama-3.3-70b-specdec
@@ -14,11 +14,11 @@ models:
---
type: embedder
provider: ollama_embedder
provider: litellm_embedder
models:
- model: nomic-embed-text # put your ollama embedder model name here
url: http://host.docker.internal:11434 # change this to your ollama host, url should be <ollama_url>
timeout: 120
- model: openai/nomic-embed-text # put your ollama embedder model name here, openai/<ollama_model_name>
api_base: http://host.docker.internal:11434/v1 # change this to your ollama host, api_base should be <ollama_url>/v1
timeout: 120
---
type: engine
@@ -40,20 +40,20 @@ recreate_index: false
type: pipeline
pipes:
- name: db_schema_indexing
embedder: ollama_embedder.nomic-embed-text
embedder: litellm_embedder.openai/nomic-embed-text
document_store: qdrant
- name: historical_question_indexing
embedder: ollama_embedder.nomic-embed-text
embedder: litellm_embedder.openai/nomic-embed-text
document_store: qdrant
- name: table_description_indexing
embedder: ollama_embedder.nomic-embed-text
embedder: litellm_embedder.openai/nomic-embed-text
document_store: qdrant
- name: db_schema_retrieval
llm: litellm_llm.openai/phi4:14b
embedder: ollama_embedder.nomic-embed-text
embedder: litellm_embedder.openai/nomic-embed-text
document_store: qdrant
- name: historical_question_retrieval
embedder: ollama_embedder.nomic-embed-text
embedder: litellm_embedder.openai/nomic-embed-text
document_store: qdrant
- name: sql_generation
llm: litellm_llm.openai/phi4:14b
@@ -89,7 +89,7 @@ pipes:
llm: litellm_llm.openai/phi4:14b
- name: question_recommendation_db_schema_retrieval
llm: litellm_llm.openai/phi4:14b
embedder: ollama_embedder.nomic-embed-text
embedder: litellm_embedder.openai/nomic-embed-text
document_store: qdrant
- name: question_recommendation_sql_generation
llm: litellm_llm.openai/phi4:14b
@@ -100,19 +100,19 @@ pipes:
llm: litellm_llm.openai/phi4:14b
- name: intent_classification
llm: litellm_llm.openai/phi4:14b
embedder: ollama_embedder.nomic-embed-text
embedder: litellm_embedder.openai/nomic-embed-text
document_store: qdrant
- name: data_assistance
llm: litellm_llm.openai/phi4:14b
- name: sql_pairs_indexing
document_store: qdrant
embedder: ollama_embedder.nomic-embed-text
embedder: litellm_embedder.openai/nomic-embed-text
- name: sql_pairs_deletion
document_store: qdrant
embedder: ollama_embedder.nomic-embed-text
embedder: litellm_embedder.openai/nomic-embed-text
- name: sql_pairs_retrieval
document_store: qdrant
embedder: ollama_embedder.nomic-embed-text
embedder: litellm_embedder.openai/nomic-embed-text
llm: litellm_llm.openai/phi4:14b
- name: preprocess_sql_data
llm: litellm_llm.openai/phi4:14b
+600 -552
View File
File diff suppressed because it is too large Load Diff
+5 -5
View File
@@ -8,11 +8,11 @@ readme = "README.md"
package-mode = false
[tool.poetry.dependencies]
python = ">=3.12.*, <4.0"
python = ">=3.12.*, <3.13"
fastapi = "^0.115.2"
uvicorn = {extras = ["standard"], version = "^0.30.1"}
python-dotenv = "^1.0.1"
haystack-ai = "^2.4.0"
haystack-ai = "==2.7.0"
openai = "^1.40.0"
qdrant-haystack = "^7.0.0"
backoff = "^2.2.1"
@@ -46,17 +46,17 @@ sseclient-py = "^1.8.0"
dspy-ai = "^2.5.26"
requests = "^2.32.2"
extra-streamlit-components = "^0.1.71"
deepeval = "^1.0.6"
tomlkit = "^0.13.0"
nltk = "^3.9.1"
[tool.poetry.group.eval.dependencies]
tomlkit = "^0.13.0"
gitpython = "^3.1.43"
plotly = "^5.24.1"
nbformat = "^5.1.3"
ipykernel = "^6.29.5"
itables = "^2.2.1"
gdown = "^5.2.0"
nltk = "^3.9.1"
deepeval = "^1.0.6"
streamlit-tags = "^1.2.8"
[tool.poetry.group.test.dependencies]
@@ -123,9 +123,13 @@ def embedder_processor(entry: dict) -> dict:
returned = {}
for model in entry["models"]:
identifier = f"{entry['provider']}.{model['model']}"
model_additional_params = {
k: v for k, v in model.items() if k not in ["model", "kwargs"]
}
returned[identifier] = {
"provider": entry["provider"],
"model": model["model"],
**model_additional_params,
**others,
}
@@ -358,6 +358,7 @@ class QdrantProvider(DocumentStoreProvider):
self.get_store(recreate_index=recreate_index)
self.get_store(dataset_name="table_descriptions", recreate_index=recreate_index)
self.get_store(dataset_name="view_questions", recreate_index=recreate_index)
self.get_store(dataset_name="sql_pairs", recreate_index=recreate_index)
def get_store(
self,
@@ -0,0 +1,197 @@
import logging
import os
from typing import Any, Dict, List, Optional, Tuple
import backoff
import openai
from haystack import Document, component
from litellm import aembedding
from tqdm import tqdm
from src.core.provider import EmbedderProvider
from src.providers.loader import provider
from src.utils import remove_trailing_slash
logger = logging.getLogger("wren-ai-service")
def _prepare_texts_to_embed(documents: List[Document]) -> List[str]:
"""
Prepare the texts to embed by concatenating the Document text with the metadata fields to embed.
"""
texts_to_embed = []
for doc in documents:
text_to_embed = "\n".join([doc.content or ""])
# copied from OpenAI embedding_utils (https://github.com/openai/openai-python/blob/main/openai/embeddings_utils.py)
# replace newlines, which can negatively affect performance.
text_to_embed = text_to_embed.replace("\n", " ")
texts_to_embed.append(text_to_embed)
return texts_to_embed
@component
class AsyncTextEmbedder:
def __init__(
self,
model: str,
api_key: Optional[str] = None,
api_base_url: Optional[str] = None,
timeout: Optional[float] = None,
**kwargs,
):
self._api_key = api_key
self._model = model
self._api_base_url = api_base_url
self._timeout = timeout
self._kwargs = kwargs
@component.output_types(embedding=List[float], meta=Dict[str, Any])
@backoff.on_exception(backoff.expo, openai.APIError, max_time=60.0, max_tries=3)
async def run(self, text: str):
if not isinstance(text, str):
raise TypeError(
"AsyncTextEmbedder expects a string as an input."
"In case you want to embed a list of Documents, please use the AsyncDocumentEmbedder."
)
# copied from OpenAI embedding_utils (https://github.com/openai/openai-python/blob/main/openai/embeddings_utils.py)
# replace newlines, which can negatively affect performance.
text_to_embed = text.replace("\n", " ")
response = await aembedding(
model=self._model,
input=[text_to_embed],
api_key=self._api_key,
api_base=self._api_base_url,
timeout=self._timeout,
**self._kwargs,
)
meta = {
"model": response.model,
"usage": dict(response.usage) if response.usage else {},
}
return {"embedding": response.data[0]["embedding"], "meta": meta}
@component
class AsyncDocumentEmbedder:
def __init__(
self,
model: str,
api_key: Optional[str] = None,
api_base_url: Optional[str] = None,
timeout: Optional[float] = None,
**kwargs,
):
self._api_key = api_key
self._model = model
self._api_base_url = api_base_url
self._timeout = timeout
self._kwargs = kwargs
async def _embed_batch(
self, texts_to_embed: List[str], batch_size: int, progress_bar: bool = True
) -> Tuple[List[List[float]], Dict[str, Any]]:
all_embeddings = []
meta: Dict[str, Any] = {}
for i in tqdm(
range(0, len(texts_to_embed), batch_size),
disable=not progress_bar,
desc="Calculating embeddings",
):
batch = texts_to_embed[i : i + batch_size]
response = await aembedding(
model=self._model,
input=batch,
api_key=self._api_key,
api_base=self._api_base_url,
timeout=self._timeout,
**self._kwargs,
)
embeddings = [el["embedding"] for el in response.data]
all_embeddings.extend(embeddings)
if "model" not in meta:
meta["model"] = response.model
if "usage" not in meta:
meta["usage"] = dict(response.usage) if response.usage else {}
else:
meta["usage"]["prompt_tokens"] += response.usage.prompt_tokens
meta["usage"]["total_tokens"] += response.usage.total_tokens
return all_embeddings, meta
@component.output_types(documents=List[Document], meta=Dict[str, Any])
@backoff.on_exception(backoff.expo, openai.RateLimitError, max_time=60, max_tries=3)
async def run(
self, documents: List[Document], batch_size: int = 32, progress_bar: bool = True
):
if (
not isinstance(documents, list)
or documents
and not isinstance(documents[0], Document)
):
raise TypeError(
"AsyncDocumentEmbedder expects a list of Documents as input."
"In case you want to embed a string, please use the AsyncTextEmbedder."
)
texts_to_embed = _prepare_texts_to_embed(documents=documents)
embeddings, meta = await self._embed_batch(
texts_to_embed=texts_to_embed,
batch_size=batch_size,
progress_bar=progress_bar,
)
for doc, emb in zip(documents, embeddings):
doc.embedding = emb
return {"documents": documents, "meta": meta}
@provider("litellm_embedder")
class LitellmEmbedderProvider(EmbedderProvider):
def __init__(
self,
api_base: str,
model: str,
api_key_name: Optional[
str
] = None, # e.g. EMBEDDER_OPENAI_API_KEY, EMBEDDER_ANTHROPIC_API_KEY, etc.
timeout: Optional[float] = 120.0,
**kwargs,
):
self._api_key = os.getenv(api_key_name) if api_key_name else None
self._api_base = remove_trailing_slash(api_base)
self._embedding_model = model
self._timeout = timeout
if "provider" in kwargs:
del kwargs["provider"]
self._kwargs = kwargs
logger.info(
f"Initializing LitellmEmbedder provider with API base: {self._api_base}"
)
logger.info(f"Using Embedding Model: {self._embedding_model}")
def get_text_embedder(self):
return AsyncTextEmbedder(
api_key=self._api_key,
api_base_url=self._api_base,
model=self._embedding_model,
timeout=self._timeout,
**self._kwargs,
)
def get_document_embedder(self):
return AsyncDocumentEmbedder(
api_key=self._api_key,
api_base_url=self._api_base,
model=self._embedding_model,
timeout=self._timeout,
**self._kwargs,
)
@@ -3,7 +3,7 @@ from src.providers import loader
def test_import_mods():
loader.import_mods("src.providers")
assert len(loader.PROVIDERS) == 11
assert len(loader.PROVIDERS) == 12
def test_get_provider():
@@ -32,6 +32,9 @@ def test_get_provider():
provider = loader.get_provider("ollama_embedder")
assert provider.__name__ == "OllamaEmbedderProvider"
provider = loader.get_provider("litellm_embedder")
assert provider.__name__ == "LitellmEmbedderProvider"
# document store provider
provider = loader.get_provider("qdrant")
assert provider.__name__ == "QdrantProvider"