mirror of
https://github.com/Canner/WrenAI.git
synced 2026-09-01 15:34:04 +08:00
feat(wren-ai-service): add litellm embedder (#1247)
This commit is contained in:
@@ -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
|
||||
|
||||
Generated
+600
-552
File diff suppressed because it is too large
Load Diff
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user