From a8dded673aa15cf9874eb3b727edc2a462de3b86 Mon Sep 17 00:00:00 2001 From: Chih-Yu Yeh Date: Wed, 3 Jul 2024 10:52:38 +0800 Subject: [PATCH] separate llm and embedder (#454) --- .github/workflows/ai-service-ci.yaml | 6 +- docker/.env.ai.example | 44 +-- docker/.env.example | 3 +- docker/docker-compose-dev.yaml | 8 +- docker/docker-compose.yaml | 8 +- wren-ai-service/.env.dev.example | 45 +-- wren-ai-service/.env.example | 3 +- wren-ai-service/.env.prod.example | 42 --- wren-ai-service/Makefile | 11 - wren-ai-service/README.md | 11 +- wren-ai-service/docker/docker-compose.yaml | 34 --- wren-ai-service/src/core/provider.py | 2 + wren-ai-service/src/eval/ask/__main__.py | 2 +- .../src/eval/ask/eval_sampledata.py | 7 +- wren-ai-service/src/globals.py | 9 +- .../src/pipelines/ask/followup_generation.py | 2 +- .../src/pipelines/ask/generation.py | 2 +- .../src/pipelines/ask/historical_question.py | 12 +- .../src/pipelines/ask/retrieval.py | 10 +- .../src/pipelines/ask/sql_correction.py | 2 +- .../src/pipelines/ask_details/generation.py | 2 +- .../src/pipelines/indexing/indexing.py | 15 +- .../src/pipelines/semantics/description.py | 8 +- .../src/providers/document_store/qdrant.py | 4 +- .../src/providers/embedder/__init__.py | 0 .../src/providers/embedder/azure_openai.py | 226 +++++++++++++++ .../src/providers/embedder/ollama.py | 188 +++++++++++++ .../src/providers/embedder/openai.py | 226 +++++++++++++++ wren-ai-service/src/providers/engine/wren.py | 4 +- .../src/providers/llm/azure_openai.py | 257 ++---------------- wren-ai-service/src/providers/llm/ollama.py | 186 +------------ wren-ai-service/src/providers/llm/openai.py | 222 ++------------- wren-ai-service/src/providers/loader.py | 5 +- wren-ai-service/src/utils.py | 17 +- .../tests/pytest/pipelines/test_ask.py | 31 ++- .../pytest/pipelines/test_ask_details.py | 2 +- .../pytest/pipelines/test_document_cleaner.py | 2 +- .../tests/pytest/providers/test_loader.py | 22 +- .../tests/pytest/services/test_ask.py | 8 +- .../tests/pytest/services/test_ask_details.py | 2 +- .../tests/pytest/services/test_semantics.py | 3 +- wren-launcher/utils/docker.go | 10 +- 42 files changed, 889 insertions(+), 814 deletions(-) delete mode 100644 wren-ai-service/.env.prod.example delete mode 100644 wren-ai-service/docker/docker-compose.yaml create mode 100644 wren-ai-service/src/providers/embedder/__init__.py create mode 100644 wren-ai-service/src/providers/embedder/azure_openai.py create mode 100644 wren-ai-service/src/providers/embedder/ollama.py create mode 100644 wren-ai-service/src/providers/embedder/openai.py diff --git a/.github/workflows/ai-service-ci.yaml b/.github/workflows/ai-service-ci.yaml index 5397393c2..9023137c9 100644 --- a/.github/workflows/ai-service-ci.yaml +++ b/.github/workflows/ai-service-ci.yaml @@ -49,9 +49,11 @@ jobs: make test env: ENV: dev - LLM_PROVIDER: openai + LLM_PROVIDER: openai_llm + EMBEDDER_PROVIDER: openai_embedder DOCUMENT_STORE_PROVIDER: qdrant - OPENAI_API_KEY: ${{ secrets.OPENAI_API_KEY }} + LLM_OPENAI_API_KEY: ${{ secrets.OPENAI_API_KEY }} + EMBEDDER_OPENAI_API_KEY: ${{ secrets.OPENAI_API_KEY }} GENERATION_MODEL: gpt-3.5-turbo WREN_ENGINE_ENDPOINT: http://localhost:8080 WREN_UI_ENDPOINT: http://localhost:3000 diff --git a/docker/.env.ai.example b/docker/.env.ai.example index 642363721..51de50fac 100644 --- a/docker/.env.ai.example +++ b/docker/.env.ai.example @@ -1,29 +1,41 @@ ## LLM -LLM_PROVIDER= # openai, azure_openai, ollama +LLM_PROVIDER=openai_llm # openai_llm, azure_openai_llm, ollama_llm +GENERATION_MODEL=gpt-4o +GENERATION_MODEL_KWARGS={"temperature": 0, "n": 1, "max_tokens": 4096, "response_format": {"type": "json_object"}} -# openai or openai-api-compatible llm -OPENAI_API_KEY= -OPENAI_API_BASE= +# openai or openai-api-compatible +LLM_OPENAI_API_KEY=sk-xxxx +LLM_OPENAI_API_BASE=https://api.openai.com/v1 # azure_openai -AZURE_CHAT_BASE= -AZURE_CHAT_KEY= -AZURE_CHAT_VERSION= - -AZURE_EMBED_BASE= -AZURE_EMBED_KEY= -AZURE_EMBED_VERSION= +LLM_AZURE_OPENAI_API_KEY= +LLM_AZURE_OPENAI_API_BASE= +LLM_AZURE_OPENAI_VERSION= # ollama -OLLAMA_URL=http://host.docker.internal:11434 +LLM_OLLAMA_URL=http://host.docker.internal:11434 -GENERATION_MODEL= + +## EMBEDDER +EMBEDDER_PROVIDER=openai_embedder # openai_embedder, azure_openai_embedder, ollama_embedder # supported embedding models providers by qdrant: https://qdrant.tech/documentation/embeddings/ -EMBEDDING_MODEL= -EMBEDDING_MODEL_DIMENSION= +EMBEDDING_MODEL=text-embedding-3-large +EMBEDDING_MODEL_DIMENSION=3072 + +# openai or openai-api-compatible +EMBEDDER_OPENAI_API_KEY=sk-xxxx +EMBEDDER_OPENAI_API_BASE=https://api.openai.com/v1 + +# azure_openai +EMBEDDER_AZURE_OPENAI_API_KEY= +EMBEDDER_AZURE_OPENAI_API_BASE= +EMBEDDER_AZURE_OPENAI_VERSION= + +# ollama +EMBEDDER_OLLAMA_URL=http://host.docker.internal:11434 ## DOCUMENT_STORE DOCUMENT_STORE_PROVIDER=qdrant -QDRANT_HOST=qdrant \ No newline at end of file +QDRANT_HOST=qdrant diff --git a/docker/.env.example b/docker/.env.example index b2ea70e69..c9bb1ced4 100644 --- a/docker/.env.example +++ b/docker/.env.example @@ -14,7 +14,8 @@ IBIS_SERVER_PORT=8000 WREN_UI_ENDPOINT=http://docker.for.mac.localhost:3000 # LLM -OPENAI_API_KEY= +LLM_OPENAI_API_KEY= +EMBEDDER_OPENAI_API_KEY= GENERATION_MODEL=gpt-3.5-turbo # gpt-3.5-turbo, gpt-4o, gpt-4-turbo # version diff --git a/docker/docker-compose-dev.yaml b/docker/docker-compose-dev.yaml index c36c16337..6cbaa55f8 100644 --- a/docker/docker-compose-dev.yaml +++ b/docker/docker-compose-dev.yaml @@ -44,10 +44,10 @@ services: environment: WREN_AI_SERVICE_PORT: ${WREN_AI_SERVICE_PORT} WREN_UI_ENDPOINT: ${WREN_UI_ENDPOINT} - GENERATION_MODEL: ${GENERATION_MODEL} - OPENAI_API_KEY: ${OPENAI_API_KEY} - AZURE_CHAT_KEY: ${AZURE_CHAT_KEY} - AZURE_EMBED_KEY: ${AZURE_EMBED_KEY} + LLM_OPENAI_API_KEY: ${LLM_OPENAI_API_KEY} + EMBEDDER_OPENAI_API_KEY: ${EMBEDDER_OPENAI_API_KEY} + LLM_AZURE_OPENAI_API_KEY: ${LLM_AZURE_OPENAI_API_KEY} + EMBEDDER_AZURE_OPENAI_API_KEY: ${EMBEDDER_AZURE_OPENAI_API_KEY} ENABLE_TIMER: ${AI_SERVICE_ENABLE_TIMER} LOGGING_LEVEL: ${AI_SERVICE_LOGGING_LEVEL} # sometimes the console won't show print messages, diff --git a/docker/docker-compose.yaml b/docker/docker-compose.yaml index 027f7f0ae..cceab1013 100644 --- a/docker/docker-compose.yaml +++ b/docker/docker-compose.yaml @@ -56,10 +56,10 @@ services: environment: WREN_AI_SERVICE_PORT: ${WREN_AI_SERVICE_PORT} WREN_UI_ENDPOINT: http://wren-ui:${WREN_UI_PORT} - GENERATION_MODEL: ${GENERATION_MODEL} - OPENAI_API_KEY: ${OPENAI_API_KEY} - AZURE_CHAT_KEY: ${AZURE_CHAT_KEY} - AZURE_EMBED_KEY: ${AZURE_EMBED_KEY} + LLM_OPENAI_API_KEY: ${LLM_OPENAI_API_KEY} + EMBEDDER_OPENAI_API_KEY: ${EMBEDDER_OPENAI_API_KEY} + LLM_AZURE_OPENAI_API_KEY: ${LLM_AZURE_OPENAI_API_KEY} + EMBEDDER_AZURE_OPENAI_API_KEY: ${EMBEDDER_AZURE_OPENAI_API_KEY} ENABLE_TIMER: ${AI_SERVICE_ENABLE_TIMER} LOGGING_LEVEL: ${AI_SERVICE_LOGGING_LEVEL} # sometimes the console won't show print messages, diff --git a/wren-ai-service/.env.dev.example b/wren-ai-service/.env.dev.example index 411d52f35..a245ae5ae 100644 --- a/wren-ai-service/.env.dev.example +++ b/wren-ai-service/.env.dev.example @@ -5,39 +5,52 @@ WREN_ENGINE_ENDPOINT=http://localhost:8080 WREN_UI_ENDPOINT=http://localhost:3000 ## LLM -LLM_PROVIDER=openai # openai, azure_openai, ollama +LLM_PROVIDER=openai_llm # openai_llm, azure_openai_llm, ollama_llm +GENERATION_MODEL=gpt-3.5-turbo -# openai or openai-api-compatible llm -OPENAI_API_KEY=sk-1234567890 -OPENAI_API_BASE=https://api.openai.com/v1 +# openai or openai-api-compatible +LLM_OPENAI_API_KEY=sk-1234567890 +LLM_OPENAI_API_BASE=https://api.openai.com/v1 # azure_openai -AZURE_CHAT_BASE= -AZURE_CHAT_KEY= -AZURE_CHAT_VERSION= - -AZURE_EMBED_BASE= -AZURE_EMBED_KEY= -AZURE_EMBED_VERSION= +LLM_AZURE_OPENAI_API_KEY= +LLM_AZURE_OPENAI_API_BASE= +LLM_AZURE_OPENAI_VERSION= # ollama -OLLAMA_URL=http://localhost:11434 +LLM_OLLAMA_URL=http://localhost:11434 -GENERATION_MODEL=gpt-3.5-turbo + +## EMBEDDER +EMBEDDER_PROVIDER=openai_embedder # openai_embedder, azure_openai_embedder, ollama_embedder EMBEDDING_MODEL=text-embedding-3-large EMBEDDING_MODEL_DIMENSION=3072 +# openai or openai-api-compatible +EMBEDDER_OPENAI_API_KEY=sk-1234567890 +EMBEDDER_OPENAI_API_BASE=https://api.openai.com/v1 + +# azure_openai +EMBEDDER_AZURE_OPENAI_API_KEY= +EMBEDDER_AZURE_OPENAI_API_BASE= +EMBEDDER_AZURE_OPENAI_VERSION= + +# ollama +EMBEDDER_OLLAMA_URL=http://localhost:11434 + + ## DOCUMENT_STORE DOCUMENT_STORE_PROVIDER=qdrant QDRANT_HOST=http://localhost:6333 -ENGINE=wren-ui +## ENGINE +ENGINE=wren_ui -## when using wren-ui as the engine +## when using wren_ui as the engine WREN_UI_ENDPOINT=http://localhost:3000 -## when using ibis as the engine +## when using wren_ibis as the engine WREN_IBIS_ENDPOINT=http://localhost:8000 WREN_IBIS_SOURCE=bigquery WREN_IBIS_MANIFEST= # this is a base64 encoded string of the MDL diff --git a/wren-ai-service/.env.example b/wren-ai-service/.env.example index f71281257..dd5514dda 100644 --- a/wren-ai-service/.env.example +++ b/wren-ai-service/.env.example @@ -1,3 +1,2 @@ -# in the future, we might delete this file since -# we only need to use .env.dev or .env.prod +# if not specified, the default ENV is prod ENV=dev \ No newline at end of file diff --git a/wren-ai-service/.env.prod.example b/wren-ai-service/.env.prod.example deleted file mode 100644 index 371339503..000000000 --- a/wren-ai-service/.env.prod.example +++ /dev/null @@ -1,42 +0,0 @@ -# docker related -# This will determine the prefix of container name. -# If the service name in docker-compose.yaml is ai-service, -# then the container name will be wren-ai-service-1. -# If COMPOSE_PROJECT_NAME is not set, the default prefix name is docker. -COMPOSE_PROJECT_NAME=wren -PLATFORM=linux/amd64 - -# fastapi related -WREN_AI_SERVICE_PORT=5555 - -# app related -# LLM Provider name should be mapped to Haystack's supported LLM providers: https://docs.haystack.deepset.ai/v2.0/docs/generators -LLM_PROVIDER=openai # azure_openai as well - - -# Document Store Provider name should be mapped to Haystack's supported Document Store providers: see the haystack documentation's left sidebar -DOCUMENT_STORE_PROVIDER=qdrant - -# llm provider specific env variables names must be in the format of [LLM_PROVIDER]_[ENV_VARIABLE_NAME] and must be in uppercase -OPENAI_API_KEY= -OPENAI_API_BASE=https://api.openai.com/v1 -EMBEDDING_MODEL=text-embedding-3-large -EMBEDDING_MODEL_DIMENSION=3072 -GENERATION_MODEL=gpt-3.5-turbo # gpt-4o, gpt-4-turbo, gpt-3.5-turbo - -#Azure openai env -AZURE_CHAT_BASE= -AZURE_CHAT_KEY= -AZURE_CHAT_VERSION= - -AZURE_EMBED_BASE= -AZURE_EMBED_KEY= -AZURE_EMBED_VERSION= - -# document store provider specific env variables names must be in the format of [DOCUMENT_STORE_PROVIDER]_[ENV_VARIABLE_NAME] and must be in uppercase -QDRANT_HOST=qdrant - -WREN_UI_ENDPOINT= - -ENABLE_TIMER= -LOGGING_LEVEL=INFO diff --git a/wren-ai-service/Makefile b/wren-ai-service/Makefile index 30d23c045..7d2decfcd 100644 --- a/wren-ai-service/Makefile +++ b/wren-ai-service/Makefile @@ -13,17 +13,6 @@ dev-down: ## wren-ai-service related ## start: poetry run python -m src.__main__ - -build: - docker compose -f docker/docker-compose.yaml --env-file .env.prod build - -up: - make dev-up - docker compose -f docker/docker-compose.yaml --env-file .env.prod up -d - -down: - make dev-down - docker compose -f docker/docker-compose.yaml --env-file .env.prod down ## wren-ai-service related ## diff --git a/wren-ai-service/README.md b/wren-ai-service/README.md index 408f43054..a3c894800 100644 --- a/wren-ai-service/README.md +++ b/wren-ai-service/README.md @@ -24,13 +24,6 @@ The following commands can quickly start the service for development: - go to `http://WREN_UI_HOST:WREN_UI_PORT`(default is http://localhost:3000) to interact interact from the UI - `make dev-down` to stop the needed containers -## Production Environment Setup - -- copy `.env.prod.example` file to `.env.prod` and fill in the environment variables -- `make build` to build the docker image -- `make up` to run the wren-ai-service and other containers -- `make down` to stop the docker container - ## Pipeline Evaluation(Deprecated, will introduce new way to evaluate the speed in the future) - install `psql` @@ -74,9 +67,9 @@ The following commands can quickly start the service for development: - wren-ui: port should be 3000 - qdrant: ports should be 6333, 6334 -## Adding your preferred LLM or Document Store +## Adding your preferred LLM, Embedder or Document Store -Please read the [documentation](https://docs.getwren.ai/installation/custom_llm) here to check out how you can add your preferred LLM or Document Store. +Please read the [documentation](https://docs.getwren.ai/installation/custom_llm) here to check out how you can add your preferred LLM, Embedder or Document Store. ## Related Issues or PRs diff --git a/wren-ai-service/docker/docker-compose.yaml b/wren-ai-service/docker/docker-compose.yaml deleted file mode 100644 index d3af589be..000000000 --- a/wren-ai-service/docker/docker-compose.yaml +++ /dev/null @@ -1,34 +0,0 @@ -version: '3.8' - -networks: - wren: - driver: bridge - -services: - wren-ai-service: - image: wren-ai-service:latest - build: - context: .. - dockerfile: docker/Dockerfile - environment: - WREN_AI_SERVICE_PORT: ${WREN_AI_SERVICE_PORT} - OPENAI_API_KEY: ${OPENAI_API_KEY} - OPENAI_API_BASE: ${OPENAI_API_BASE} - GENERATION_MODEL: ${GENERATION_MODEL} - QDRANT_HOST: ${QDRANT_HOST} - WREN_UI_ENDPOINT: ${WREN_UI_ENDPOINT} - ENABLE_TIMER: ${ENABLE_TIMER} - LOGGING_LEVEL: ${LOGGING_LEVEL} - # sometimes the console won't show print messages, - # using PYTHONUNBUFFERED: 1 can fix this - PYTHONUNBUFFERED: 1 - ports: - - ${WREN_AI_SERVICE_PORT}:${WREN_AI_SERVICE_PORT} - depends_on: - - qdrant - - qdrant: - image: qdrant/qdrant:v1.7.4 - ports: - - 6333:6333 - - 6334:6334 diff --git a/wren-ai-service/src/core/provider.py b/wren-ai-service/src/core/provider.py index cd24432c5..46f5b73cb 100644 --- a/wren-ai-service/src/core/provider.py +++ b/wren-ai-service/src/core/provider.py @@ -8,6 +8,8 @@ class LLMProvider(metaclass=ABCMeta): def get_generator(self, *args, **kwargs): ... + +class EmbedderProvider(metaclass=ABCMeta): @abstractmethod def get_text_embedder(self, *args, **kwargs): ... diff --git a/wren-ai-service/src/eval/ask/__main__.py b/wren-ai-service/src/eval/ask/__main__.py index ecdd008af..ca5eb05bd 100644 --- a/wren-ai-service/src/eval/ask/__main__.py +++ b/wren-ai-service/src/eval/ask/__main__.py @@ -11,9 +11,9 @@ import orjson from tqdm import tqdm from src.pipelines.ask.generation import Generation -from src.pipelines.ask.indexing import Indexing from src.pipelines.ask.retrieval import Retrieval from src.pipelines.ask.sql_correction import SQLCorrection +from src.pipelines.indexing.indexing import Indexing from src.pipelines.semantics import description from src.utils import init_providers, load_env_vars from src.web.v1.services.semantics import ( diff --git a/wren-ai-service/src/eval/ask/eval_sampledata.py b/wren-ai-service/src/eval/ask/eval_sampledata.py index 5563f6ae9..1cc39e810 100644 --- a/wren-ai-service/src/eval/ask/eval_sampledata.py +++ b/wren-ai-service/src/eval/ask/eval_sampledata.py @@ -15,9 +15,9 @@ import requests from tqdm import tqdm from src.pipelines.ask.generation import Generation -from src.pipelines.ask.indexing import Indexing from src.pipelines.ask.retrieval import Retrieval from src.pipelines.ask.sql_correction import SQLCorrection +from src.pipelines.indexing.indexing import Indexing from src.utils import init_providers, load_env_vars load_env_vars() @@ -123,12 +123,12 @@ if __name__ == "__main__": Path("./outputs/ask/sampledata").mkdir(parents=True) # init ask pipeline - llm_provider, document_store_provider = init_providers() + llm_provider, embedder_provider, document_store_provider, _ = init_providers() document_store = document_store_provider.get_store( dataset_name=SAMPLE_DATASET_NAME, recreate_index=True, ) - embedder = llm_provider.get_text_embedder() + embedder = embedder_provider.get_text_embedder() retriever = document_store_provider.get_retriever( document_store=document_store, top_k=10, @@ -152,6 +152,7 @@ if __name__ == "__main__": mdl = get_mdl_from_wren_engine() indexing_pipeline = Indexing( llm_provider=llm_provider, + embedder_provider=embedder_provider, store_provider=document_store_provider, ) indexing_pipeline.run(orjson.dumps(mdl).decode("utf-8")) diff --git a/wren-ai-service/src/globals.py b/wren-ai-service/src/globals.py index c0d42efd2..3cd5e2ee2 100644 --- a/wren-ai-service/src/globals.py +++ b/wren-ai-service/src/globals.py @@ -33,7 +33,7 @@ ASK_DETAILS_SERVICE = None def init_globals(): global SEMANTIC_SERVICE, ASK_SERVICE, ASK_DETAILS_SERVICE - llm_provider, document_store_provider, engine = init_providers() + llm_provider, embedder_provider, document_store_provider, engine = init_providers() # Recreate the document store to ensure a clean slate # TODO: for SaaS, we need to use a flag to prevent this collection_recreation @@ -46,6 +46,7 @@ def init_globals(): pipelines={ "generate_description": description.Generation( llm_provider=llm_provider, + embedder_provider=embedder_provider, document_store_provider=document_store_provider, ), }, @@ -54,15 +55,15 @@ def init_globals(): ASK_SERVICE = AskService( pipelines={ "indexing": indexing.Indexing( - llm_provider=llm_provider, + embedder_provider=embedder_provider, document_store_provider=document_store_provider, ), "retrieval": ask_retrieval.Retrieval( - llm_provider=llm_provider, + embedder_provider=embedder_provider, document_store_provider=document_store_provider, ), "historical_question": historical_question.HistoricalQuestion( - llm_provider=llm_provider, + embedder_provider=embedder_provider, store_provider=document_store_provider, ), "generation": ask_generation.Generation( diff --git a/wren-ai-service/src/pipelines/ask/followup_generation.py b/wren-ai-service/src/pipelines/ask/followup_generation.py index 6e80b19d4..1d6b8aff4 100644 --- a/wren-ai-service/src/pipelines/ask/followup_generation.py +++ b/wren-ai-service/src/pipelines/ask/followup_generation.py @@ -231,7 +231,7 @@ if __name__ == "__main__": load_env_vars() - llm_provider, _, engine = init_providers() + llm_provider, _, _, engine = init_providers() pipeline = FollowUpGeneration(llm_provider=llm_provider, engine=engine) pipeline.visualize( diff --git a/wren-ai-service/src/pipelines/ask/generation.py b/wren-ai-service/src/pipelines/ask/generation.py index b5664bbb3..018af68a6 100644 --- a/wren-ai-service/src/pipelines/ask/generation.py +++ b/wren-ai-service/src/pipelines/ask/generation.py @@ -189,7 +189,7 @@ if __name__ == "__main__": load_env_vars() - llm_provider, _, engine = init_providers() + llm_provider, _, _, engine = init_providers() pipeline = Generation( llm_provider=llm_provider, engine=engine, diff --git a/wren-ai-service/src/pipelines/ask/historical_question.py b/wren-ai-service/src/pipelines/ask/historical_question.py index bac80ae40..556a7e482 100644 --- a/wren-ai-service/src/pipelines/ask/historical_question.py +++ b/wren-ai-service/src/pipelines/ask/historical_question.py @@ -9,7 +9,7 @@ from hamilton.experimental.h_async import AsyncDriver from haystack import Document, component from src.core.pipeline import BasicPipeline, async_validate -from src.core.provider import DocumentStoreProvider, LLMProvider +from src.core.provider import DocumentStoreProvider, EmbedderProvider from src.utils import ( async_timer, init_providers, @@ -87,9 +87,11 @@ def formatted_output( class HistoricalQuestion(BasicPipeline): def __init__( - self, llm_provider: LLMProvider, store_provider: DocumentStoreProvider + self, + embedder_provider: EmbedderProvider, + store_provider: DocumentStoreProvider, ) -> None: - self._embedder = llm_provider.get_text_embedder() + self._embedder = embedder_provider.get_text_embedder() self._retriever = store_provider.get_retriever( document_store=store_provider.get_store(dataset_name="view_questions"), ) @@ -143,10 +145,10 @@ if __name__ == "__main__": load_env_vars() - llm_provider, document_store_provider, _ = init_providers() + _, embedder_provider, document_store_provider, _ = init_providers() pipeline = HistoricalQuestion( - llm_provider=llm_provider, store_provider=document_store_provider + embedder_provider=embedder_provider, store_provider=document_store_provider ) pipeline.visualize("this is a query") diff --git a/wren-ai-service/src/pipelines/ask/retrieval.py b/wren-ai-service/src/pipelines/ask/retrieval.py index 27b9709f6..b01e146a5 100644 --- a/wren-ai-service/src/pipelines/ask/retrieval.py +++ b/wren-ai-service/src/pipelines/ask/retrieval.py @@ -7,7 +7,7 @@ from hamilton import base from hamilton.experimental.h_async import AsyncDriver from src.core.pipeline import BasicPipeline, async_validate -from src.core.provider import DocumentStoreProvider, LLMProvider +from src.core.provider import DocumentStoreProvider, EmbedderProvider from src.utils import async_timer, init_providers logger = logging.getLogger("wren-ai-service") @@ -31,10 +31,10 @@ async def retrieval(embedding: dict, retriever: Any) -> dict: class Retrieval(BasicPipeline): def __init__( self, - llm_provider: LLMProvider, + embedder_provider: EmbedderProvider, document_store_provider: DocumentStoreProvider, ): - self._embedder = llm_provider.get_text_embedder() + self._embedder = embedder_provider.get_text_embedder() self._retriever = document_store_provider.get_retriever( document_store_provider.get_store() ) @@ -81,9 +81,9 @@ if __name__ == "__main__": load_env_vars() - llm_provider, document_store_provider, _ = init_providers() + _, embedder_provider, document_store_provider, _ = init_providers() pipeline = Retrieval( - llm_provider=llm_provider, + embedder_provider=embedder_provider, document_store_provider=document_store_provider, ) diff --git a/wren-ai-service/src/pipelines/ask/sql_correction.py b/wren-ai-service/src/pipelines/ask/sql_correction.py index 2aba540e1..b29f7c93f 100644 --- a/wren-ai-service/src/pipelines/ask/sql_correction.py +++ b/wren-ai-service/src/pipelines/ask/sql_correction.py @@ -155,7 +155,7 @@ if __name__ == "__main__": load_env_vars() - llm_provider, _, engine = init_providers() + llm_provider, _, _, engine = init_providers() pipeline = SQLCorrection( llm_provider=llm_provider, engine=engine, diff --git a/wren-ai-service/src/pipelines/ask_details/generation.py b/wren-ai-service/src/pipelines/ask_details/generation.py index b815ea54a..6b1846fc7 100644 --- a/wren-ai-service/src/pipelines/ask_details/generation.py +++ b/wren-ai-service/src/pipelines/ask_details/generation.py @@ -194,7 +194,7 @@ if __name__ == "__main__": load_env_vars() - llm_provider, _, engine = init_providers() + llm_provider, _, _, engine = init_providers() pipeline = Generation( llm_provider=llm_provider, engine=engine, diff --git a/wren-ai-service/src/pipelines/indexing/indexing.py b/wren-ai-service/src/pipelines/indexing/indexing.py index bf4eedf84..613b8fb3d 100644 --- a/wren-ai-service/src/pipelines/indexing/indexing.py +++ b/wren-ai-service/src/pipelines/indexing/indexing.py @@ -15,7 +15,7 @@ from haystack.document_stores.types import DocumentStore, DuplicatePolicy from tqdm import tqdm from src.core.pipeline import BasicPipeline, async_validate -from src.core.provider import DocumentStoreProvider, LLMProvider +from src.core.provider import DocumentStoreProvider, EmbedderProvider from src.utils import async_timer, init_providers, timer logger = logging.getLogger("wren-ai-service") @@ -376,7 +376,9 @@ def write_view(embed_view: Dict[str, Any], view_writer: DocumentWriter) -> None: class Indexing(BasicPipeline): def __init__( - self, llm_provider: LLMProvider, document_store_provider: DocumentStoreProvider + self, + embedder_provider: EmbedderProvider, + document_store_provider: DocumentStoreProvider, ) -> None: ddl_store = document_store_provider.get_store() view_store = document_store_provider.get_store(dataset_name="view_questions") @@ -385,13 +387,13 @@ class Indexing(BasicPipeline): self.validator = MDLValidator() self.ddl_converter = DDLConverter() - self.ddl_embedder = llm_provider.get_document_embedder() + self.ddl_embedder = embedder_provider.get_document_embedder() self.ddl_writer = DocumentWriter( document_store=ddl_store, policy=DuplicatePolicy.OVERWRITE, ) self.view_converter = ViewConverter() - self.view_embedder = llm_provider.get_document_embedder() + self.view_embedder = embedder_provider.get_document_embedder() self.view_writer = DocumentWriter( document_store=view_store, policy=DuplicatePolicy.OVERWRITE, @@ -447,10 +449,11 @@ if __name__ == "__main__": from src.utils import load_env_vars load_env_vars() - llm_provider, document_store_provider, _ = init_providers() + _, embedder_provider, document_store_provider, _ = init_providers() pipeline = Indexing( - llm_provider=llm_provider, document_store_provider=document_store_provider + embedder_provider=embedder_provider, + document_store_provider=document_store_provider, ) input = '{"models": [], "views": [], "relationships": [], "metrics": []}' diff --git a/wren-ai-service/src/pipelines/semantics/description.py b/wren-ai-service/src/pipelines/semantics/description.py index d7719b2c3..3a59b37e2 100644 --- a/wren-ai-service/src/pipelines/semantics/description.py +++ b/wren-ai-service/src/pipelines/semantics/description.py @@ -6,7 +6,7 @@ from haystack import Pipeline from haystack.components.builders import PromptBuilder from src.core.pipeline import BasicPipeline -from src.core.provider import DocumentStoreProvider, LLMProvider +from src.core.provider import DocumentStoreProvider, EmbedderProvider, LLMProvider from src.utils import init_providers _TEMPLATE = """ @@ -54,11 +54,12 @@ class Generation(BasicPipeline): def __init__( self, llm_provider: LLMProvider, + embedder_provider: EmbedderProvider, document_store_provider: DocumentStoreProvider, ): self._prompt_builder = PromptBuilder(template=_TEMPLATE) self._pipe = Pipeline() - self._pipe.add_component("text_embedder", llm_provider.get_text_embedder()) + self._pipe.add_component("text_embedder", embedder_provider.get_text_embedder()) self._pipe.add_component( "retriever", document_store_provider.get_retriever(document_store_provider.get_store()), @@ -102,9 +103,10 @@ if __name__ == "__main__": load_env_vars() - llm_provider, document_store_provider, _ = init_providers() + llm_provider, embedder_provider, document_store_provider, _ = init_providers() pipe = Generation( llm_provider=llm_provider, + embedder_provider=embedder_provider, document_store_provider=document_store_provider, ) diff --git a/wren-ai-service/src/providers/document_store/qdrant.py b/wren-ai-service/src/providers/document_store/qdrant.py index d94f928f1..7a1a2af7b 100644 --- a/wren-ai-service/src/providers/document_store/qdrant.py +++ b/wren-ai-service/src/providers/document_store/qdrant.py @@ -203,7 +203,9 @@ class QdrantProvider(DocumentStoreProvider): if os.getenv("EMBEDDING_MODEL_DIMENSION") else 0 ) - or get_default_embedding_model_dim(os.getenv("LLM_PROVIDER", "openai")), + or get_default_embedding_model_dim( + os.getenv("EMBEDDER_PROVIDER", "openai_embedder") + ), dataset_name: Optional[str] = None, recreate_index: bool = False, ): diff --git a/wren-ai-service/src/providers/embedder/__init__.py b/wren-ai-service/src/providers/embedder/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/wren-ai-service/src/providers/embedder/azure_openai.py b/wren-ai-service/src/providers/embedder/azure_openai.py new file mode 100644 index 000000000..cb1ddb20f --- /dev/null +++ b/wren-ai-service/src/providers/embedder/azure_openai.py @@ -0,0 +1,226 @@ +import logging +import os +from typing import Any, Dict, List, Optional, Tuple + +import backoff +import openai +from haystack import Document, component +from haystack.components.embedders import ( + AzureOpenAIDocumentEmbedder, + AzureOpenAITextEmbedder, +) +from haystack.utils import Secret +from openai import AsyncAzureOpenAI +from tqdm import tqdm + +from src.core.provider import EmbedderProvider +from src.providers.loader import provider + +logger = logging.getLogger("wren-ai-service") + +EMBEDDING_MODEL = "text-embedding-3-small" +EMBEDDING_MODEL_DIMENSION = 1536 + + +@component +class AsyncTextEmbedder(AzureOpenAITextEmbedder): + def __init__( + self, + api_key: Secret = Secret.from_env_var("EMBEDDER_AZURE_OPENAI_API_KEY"), + model: str = "text-embedding-3-small", + dimensions: Optional[int] = None, + api_base_url: Optional[str] = None, + api_version: Optional[str] = None, + organization: Optional[str] = None, + prefix: str = "", + suffix: str = "", + ): + super(AsyncTextEmbedder, self).__init__( + azure_endpoint=api_base_url, + api_version=api_version, + azure_deployment=model, + dimensions=dimensions, + api_key=api_key, + organization=organization, + prefix=prefix, + suffix=suffix, + ) + + self.client = AsyncAzureOpenAI( + azure_endpoint=api_base_url, + api_version=api_version, + api_key=api_key.resolve_value(), + ) + + @component.output_types(embedding=List[float], meta=Dict[str, Any]) + @backoff.on_exception(backoff.expo, openai.RateLimitError, max_time=60, max_tries=3) + async def run(self, text: str): + if not isinstance(text, str): + raise TypeError( + "AzureOpenAITextEmbedder expects a string as an input." + "In case you want to embed a list of Documents, please use the AzureOpenAIDocumentEmbedder." + ) + + logger.info(f"Running Async Azure OpenAI text embedder with text: {text}") + + text_to_embed = self.prefix + text + self.suffix + + # 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", " ") + + if self.dimensions is not None: + response = await self.client.embeddings.create( + model=self.azure_deployment, + dimensions=self.dimensions, + input=text_to_embed, + ) + else: + response = await self.client.embeddings.create( + model=self.azure_deployment, input=text_to_embed + ) + + meta = {"model": response.model, "usage": dict(response.usage)} + + return {"embedding": response.data[0].embedding, "meta": meta} + + +@component +class AsyncDocumentEmbedder(AzureOpenAIDocumentEmbedder): + def __init__( + self, + api_key: Secret = Secret.from_env_var("EMBEDDER_AZURE_OPENAI_API_KEY"), + model: str = "text-embedding-3-small", + dimensions: Optional[int] = None, + api_base_url: Optional[str] = None, + api_version: Optional[str] = None, + organization: Optional[str] = None, + prefix: str = "", + suffix: str = "", + batch_size: int = 32, + progress_bar: bool = True, + meta_fields_to_embed: Optional[List[str]] = None, + embedding_separator: str = "\n", + ): + super(AsyncDocumentEmbedder, self).__init__( + azure_endpoint=api_base_url, + api_version=api_version, + azure_deployment=model, + dimensions=dimensions, + api_key=api_key, + organization=organization, + prefix=prefix, + suffix=suffix, + batch_size=batch_size, + progress_bar=progress_bar, + meta_fields_to_embed=meta_fields_to_embed, + embedding_separator=embedding_separator, + ) + + self.client = AsyncAzureOpenAI( + azure_endpoint=api_base_url, + api_version=api_version, + api_key=api_key.resolve_value(), + ) + + async def _embed_batch( + self, texts_to_embed: List[str], batch_size: int + ) -> 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 self.progress_bar, + desc="Calculating embeddings", + ): + batch = texts_to_embed[i : i + batch_size] + if self.dimensions is not None: + response = await self.client.embeddings.create( + model=self.azure_deployment, dimensions=self.dimensions, input=batch + ) + else: + response = await self.client.embeddings.create( + model=self.azure_deployment, input=batch + ) + 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) + 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]): + if ( + not isinstance(documents, list) + or documents + and not isinstance(documents[0], Document) + ): + raise TypeError( + "AzureOpenAIDocumentEmbedder expects a list of Documents as input." + "In case you want to embed a string, please use the AzureOpenAITextEmbedder." + ) + + logger.info( + f"Running Async OpenAI document embedder with documents: {documents}" + ) + + texts_to_embed = self._prepare_texts_to_embed(documents=documents) + + embeddings, meta = await self._embed_batch( + texts_to_embed=texts_to_embed, batch_size=self.batch_size + ) + + for doc, emb in zip(documents, embeddings): + doc.embedding = emb + + return {"documents": documents, "meta": meta} + + +@provider("azure_openai_embedder") +class AzureOpenAIEmbedderProvider(EmbedderProvider): + def __init__( + self, + embed_api_key: Secret = Secret.from_env_var("EMBEDDER_AZURE_OPENAI_API_KEY"), + embed_api_base: str = os.getenv("EMBEDDER_AZURE_OPENAI_API_BASE"), + embed_api_version: str = os.getenv("EMBEDDER_AZURE_OPENAI_VERSION"), + embedding_model: str = os.getenv("EMBEDDING_MODEL") or EMBEDDING_MODEL, + embedding_model_dim: int = ( + int(os.getenv("EMBEDDING_MODEL_DIMENSION")) + if os.getenv("EMBEDDING_MODEL_DIMENSION") + else 0 + ) + or EMBEDDING_MODEL_DIMENSION, + ): + logger.info(f"Using Azure OpenAI Embedding Model: {embedding_model}") + + self._embedding_api_base = embed_api_base + self._embedding_api_key = embed_api_key + self._embedding_api_version = embed_api_version + self._embedding_model = embedding_model + self._embedding_model_dim = embedding_model_dim + + def get_text_embedder(self): + return AsyncTextEmbedder( + api_key=self._embedding_api_key, + model=self._embedding_model, + dimensions=self._embedding_model_dim, + api_base_url=self._embedding_api_base, + api_version=self._embedding_api_version, + ) + + def get_document_embedder(self): + return AsyncDocumentEmbedder( + api_key=self._embedding_api_key, + model=self._embedding_model, + dimensions=self._embedding_model_dim, + api_base_url=self._embedding_api_base, + api_version=self._embedding_api_version, + ) diff --git a/wren-ai-service/src/providers/embedder/ollama.py b/wren-ai-service/src/providers/embedder/ollama.py new file mode 100644 index 000000000..11b3f9e4a --- /dev/null +++ b/wren-ai-service/src/providers/embedder/ollama.py @@ -0,0 +1,188 @@ +import logging +import os +import time +from typing import Any, Dict, List, Optional + +import aiohttp +from haystack import Document, component +from haystack_integrations.components.embedders.ollama import ( + OllamaDocumentEmbedder, + OllamaTextEmbedder, +) +from tqdm import tqdm + +from src.core.provider import EmbedderProvider +from src.providers.loader import provider + +logger = logging.getLogger("wren-ai-service") + +EMBEDDER_OLLAMA_URL = "http://localhost:11434" +EMBEDDING_MODEL = "nomic-embed-text" +EMBEDDING_MODEL_DIMENSION = 768 # https://huggingface.co/nomic-ai/nomic-embed-text-v1.5 + + +@component +class AsyncTextEmbedder(OllamaTextEmbedder): + def __init__( + self, + model: str = "nomic-embed-text", + url: str = "http://localhost:11434/api/embeddings", + generation_kwargs: Optional[Dict[str, Any]] = None, + timeout: int = 120, + ): + super(AsyncTextEmbedder, self).__init__( + model=model, + url=url, + generation_kwargs=generation_kwargs, + timeout=timeout, + ) + + @component.output_types(embedding=List[float], meta=Dict[str, Any]) + async def run( + self, + text: str, + generation_kwargs: Optional[Dict[str, Any]] = None, + ): + logger.debug(f"Running Ollama text embedder with text: {text}") + + payload = self._create_json_payload(text, generation_kwargs) + + start = time.perf_counter() + async with aiohttp.ClientSession( + timeout=aiohttp.ClientTimeout(self.timeout) + ) as session: + async with session.post( + self.url, + json=payload, + ) as response: + elapsed = time.perf_counter() - start + result = await response.json() + + result["meta"] = {"model": self.model, "duration": elapsed} + + return result + + +@component +class AsyncDocumentEmbedder(OllamaDocumentEmbedder): + def __init__( + self, + model: str = "nomic-embed-text", + url: str = "http://localhost:11434/api/embeddings", + generation_kwargs: Optional[Dict[str, Any]] = None, + timeout: int = 120, + prefix: str = "", + suffix: str = "", + progress_bar: bool = True, + meta_fields_to_embed: Optional[List[str]] = None, + embedding_separator: str = "\n", + ): + super(AsyncDocumentEmbedder, self).__init__( + model=model, + url=url, + generation_kwargs=generation_kwargs, + timeout=timeout, + prefix=prefix, + suffix=suffix, + progress_bar=progress_bar, + meta_fields_to_embed=meta_fields_to_embed, + embedding_separator=embedding_separator, + ) + + async def _embed_batch( + self, + texts_to_embed: List[str], + batch_size: int, + generation_kwargs: Optional[Dict[str, Any]] = None, + ): + """ + Ollama Embedding only allows single uploads, not batching. Currently the batch size is set to 1. + If this changes in the future, line 86 (the first line within the for loop), can contain: + batch = texts_to_embed[i + i + batch_size] + """ + + all_embeddings = [] + meta: Dict[str, Any] = {"model": self.model} + + async with aiohttp.ClientSession( + timeout=aiohttp.ClientTimeout(self.timeout) + ) as session: + for i in tqdm( + range(0, len(texts_to_embed), batch_size), + disable=not self.progress_bar, + desc="Calculating embeddings", + ): + batch = texts_to_embed[i] # Single batch only + payload = self._create_json_payload(batch, generation_kwargs) + + async with session.post( + self.url, + json=payload, + ) as response: + result = await response.json() + all_embeddings.append(result["embedding"]) + + return all_embeddings, meta + + @component.output_types(embedding=List[float], meta=Dict[str, Any]) + async def run( + self, + documents: List[str], + generation_kwargs: Optional[Dict[str, Any]] = None, + ): + logger.debug(f"Running Ollama document embedder with documents: {documents}") + + if ( + not isinstance(documents, list) + or documents + and not isinstance(documents[0], Document) + ): + msg = ( + "OllamaDocumentEmbedder expects a list of Documents as input." + "In case you want to embed a list of strings, please use the OllamaTextEmbedder." + ) + raise TypeError(msg) + + texts_to_embed = self._prepare_texts_to_embed(documents=documents) + embeddings, meta = await self._embed_batch( + texts_to_embed=texts_to_embed, + batch_size=self.batch_size, + generation_kwargs=generation_kwargs, + ) + + for doc, emb in zip(documents, embeddings): + doc.embedding = emb + + return {"documents": documents, "meta": meta} + + +@provider("ollama_embedder") +class OllamaEmbedderProvider(EmbedderProvider): + def __init__( + self, + url: str = os.getenv("EMBEDDER_OLLAMA_URL") or EMBEDDER_OLLAMA_URL, + embedding_model: str = os.getenv("EMBEDDING_MODEL") or EMBEDDING_MODEL, + ): + logger.info(f"Using Ollama Embedding Model: {embedding_model}") + self._url = url + self._embedding_model = embedding_model + + def get_text_embedder( + self, + model_kwargs: Optional[Dict[str, Any]] = None, + ): + return AsyncTextEmbedder( + model=self._embedding_model, + url=f"{self._url}/api/embeddings", + generation_kwargs=model_kwargs, + ) + + def get_document_embedder( + self, + model_kwargs: Optional[Dict[str, Any]] = None, + ): + return AsyncDocumentEmbedder( + model=self._embedding_model, + url=f"{self._url}/api/embeddings", + generation_kwargs=model_kwargs, + ) diff --git a/wren-ai-service/src/providers/embedder/openai.py b/wren-ai-service/src/providers/embedder/openai.py new file mode 100644 index 000000000..936e50628 --- /dev/null +++ b/wren-ai-service/src/providers/embedder/openai.py @@ -0,0 +1,226 @@ +import logging +import os +from typing import Any, Dict, List, Optional, Tuple + +import backoff +import openai +from haystack import Document, component +from haystack.components.embedders import OpenAIDocumentEmbedder, OpenAITextEmbedder +from haystack.utils import Secret +from openai import AsyncOpenAI, OpenAI +from tqdm import tqdm + +from src.core.provider import EmbedderProvider +from src.providers.loader import provider + +logger = logging.getLogger("wren-ai-service") + +EMBEDDER_OPENAI_API_BASE = "https://api.openai.com/v1" +EMBEDDING_MODEL = "text-embedding-3-large" +EMBEDDING_MODEL_DIMENSION = 3072 + + +@component +class AsyncTextEmbedder(OpenAITextEmbedder): + def __init__( + self, + api_key: Secret = Secret.from_env_var("EMBEDDER_OPENAI_API_KEY"), + model: str = "text-embedding-ada-002", + dimensions: Optional[int] = None, + api_base_url: Optional[str] = None, + organization: Optional[str] = None, + prefix: str = "", + suffix: str = "", + ): + super(AsyncTextEmbedder, self).__init__( + api_key, + model, + dimensions, + api_base_url, + organization, + prefix, + suffix, + ) + self.client = AsyncOpenAI( + api_key=api_key.resolve_value(), + organization=organization, + base_url=api_base_url, + ) + + @component.output_types(embedding=List[float], meta=Dict[str, Any]) + @backoff.on_exception(backoff.expo, openai.RateLimitError, max_time=60, max_tries=3) + async def run(self, text: str): + if not isinstance(text, str): + raise TypeError( + "OpenAITextEmbedder expects a string as an input." + "In case you want to embed a list of Documents, please use the OpenAIDocumentEmbedder." + ) + + logger.debug(f"Running Async OpenAI text embedder with text: {text}") + + text_to_embed = self.prefix + text + self.suffix + + # 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", " ") + + if self.dimensions is not None: + response = await self.client.embeddings.create( + model=self.model, dimensions=self.dimensions, input=text_to_embed + ) + else: + response = await self.client.embeddings.create( + model=self.model, input=text_to_embed + ) + + meta = {"model": response.model, "usage": dict(response.usage)} + + return {"embedding": response.data[0].embedding, "meta": meta} + + +@component +class AsyncDocumentEmbedder(OpenAIDocumentEmbedder): + def __init__( + self, + api_key: Secret = Secret.from_env_var("EMBEDDER_OPENAI_API_KEY"), + model: str = "text-embedding-ada-002", + dimensions: Optional[int] = None, + api_base_url: Optional[str] = None, + organization: Optional[str] = None, + prefix: str = "", + suffix: str = "", + batch_size: int = 32, + progress_bar: bool = True, + meta_fields_to_embed: Optional[List[str]] = None, + embedding_separator: str = "\n", + ): + super(AsyncDocumentEmbedder, self).__init__( + api_key, + model, + dimensions, + api_base_url, + organization, + prefix, + suffix, + batch_size, + progress_bar, + meta_fields_to_embed, + embedding_separator, + ) + self.client = AsyncOpenAI( + api_key=api_key.resolve_value(), + organization=organization, + base_url=api_base_url, + ) + + async def _embed_batch( + self, texts_to_embed: List[str], batch_size: int + ) -> 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 self.progress_bar, + desc="Calculating embeddings", + ): + batch = texts_to_embed[i : i + batch_size] + if self.dimensions is not None: + response = await self.client.embeddings.create( + model=self.model, dimensions=self.dimensions, input=batch + ) + else: + response = await self.client.embeddings.create( + model=self.model, input=batch + ) + 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) + 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]): + if ( + not isinstance(documents, list) + or documents + and not isinstance(documents[0], Document) + ): + raise TypeError( + "OpenAIDocumentEmbedder expects a list of Documents as input." + "In case you want to embed a string, please use the OpenAITextEmbedder." + ) + + logger.debug( + f"Running Async OpenAI document embedder with documents: {documents}" + ) + + texts_to_embed = self._prepare_texts_to_embed(documents=documents) + + embeddings, meta = await self._embed_batch( + texts_to_embed=texts_to_embed, batch_size=self.batch_size + ) + + for doc, emb in zip(documents, embeddings): + doc.embedding = emb + + return {"documents": documents, "meta": meta} + + +@provider("openai_embedder") +class OpenAIEmbedderProvider(EmbedderProvider): + def __init__( + self, + api_key: Secret = Secret.from_env_var("EMBEDDER_OPENAI_API_KEY"), + api_base: str = os.getenv("EMBEDDER_OPENAI_API_BASE") + or EMBEDDER_OPENAI_API_BASE, + embedding_model: str = os.getenv("EMBEDDING_MODEL") or EMBEDDING_MODEL, + embedding_model_dim: int = ( + int(os.getenv("EMBEDDING_MODEL_DIMENSION")) + if os.getenv("EMBEDDING_MODEL_DIMENSION") + else 0 + ) + or EMBEDDING_MODEL_DIMENSION, + ): + def _verify_api_key(api_key: str, api_base: str) -> None: + """ + this is a temporary solution to verify that the required environment variables are set + """ + OpenAI(api_key=api_key, base_url=api_base).models.list() + + logger.info(f"Initializing OpenAIEmbedder provider with API base: {api_base}") + # TODO: currently only OpenAI api key can be verified + if api_base == EMBEDDER_OPENAI_API_BASE: + _verify_api_key(api_key.resolve_value(), api_base) + logger.info(f"Using OpenAI Embedding Model: {embedding_model}") + else: + logger.info( + f"Using OpenAI API-compatible Embedding Model: {embedding_model}" + ) + self._api_key = api_key + self._api_base = api_base + self._embedding_model = embedding_model + self._embedding_model_dim = embedding_model_dim + + def get_text_embedder(self): + return AsyncTextEmbedder( + api_key=self._api_key, + api_base_url=self._api_base, + model=self._embedding_model, + dimensions=self._embedding_model_dim, + ) + + def get_document_embedder(self): + return AsyncDocumentEmbedder( + api_key=self._api_key, + api_base_url=self._api_base, + model=self._embedding_model, + dimensions=self._embedding_model_dim, + ) diff --git a/wren-ai-service/src/providers/engine/wren.py b/wren-ai-service/src/providers/engine/wren.py index 499eae67d..c83f4ace5 100644 --- a/wren-ai-service/src/providers/engine/wren.py +++ b/wren-ai-service/src/providers/engine/wren.py @@ -11,7 +11,7 @@ from src.providers.loader import provider logger = logging.getLogger("wren-ai-service") -@provider("wren-ui") +@provider("wren_ui") class WrenUI(Engine): def __init__(self, endpoint: str = os.getenv("WREN_UI_ENDPOINT")): self._endpoint = endpoint @@ -41,7 +41,7 @@ class WrenUI(Engine): return False, res.get("errors", [{}])[0].get("message", "Unknown error") -@provider("wren-ibis") +@provider("wren_ibis") class WrenIbis(Engine): def __init__(self, endpoint: str = os.getenv("WREN_IBIS_ENDPOINT")): self._endpoint = endpoint diff --git a/wren-ai-service/src/providers/llm/azure_openai.py b/wren-ai-service/src/providers/llm/azure_openai.py index de1c8d803..fa68418d0 100644 --- a/wren-ai-service/src/providers/llm/azure_openai.py +++ b/wren-ai-service/src/providers/llm/azure_openai.py @@ -1,57 +1,44 @@ import logging import os -from typing import Any, Callable, Dict, List, Optional, Tuple, Union +from typing import Any, Callable, Dict, List, Optional, Union import backoff import openai -from haystack import Document, component -from haystack.components.embedders import ( - AzureOpenAIDocumentEmbedder, - AzureOpenAITextEmbedder, -) +import orjson +from haystack import component from haystack.components.generators import AzureOpenAIGenerator from haystack.dataclasses import ChatMessage, StreamingChunk from haystack.utils import Secret from openai import AsyncAzureOpenAI, Stream from openai.types.chat import ChatCompletion, ChatCompletionChunk -from tqdm import tqdm from src.core.provider import LLMProvider from src.providers.loader import provider -EMBEDDING_MODEL_NAME = "text-embedding-3-small" -EMBEDDING_MODEL_DIMENSION = 1536 logger = logging.getLogger("wren-ai-service") -AZURE_GENERATION_MODEL = "gpt-4-turbo" -AZURE_GENERATION_MODEL_KWARGS = { + +GENERATION_MODEL = "gpt-4-turbo" +GENERATION_MODEL_KWARGS = { "temperature": 0, "n": 1, "max_tokens": 1000, "response_format": {"type": "json_object"}, } -chat_base = os.getenv("AZURE_CHAT_BASE") -chat_token = Secret.from_env_var("AZURE_CHAT_KEY") -chat_version = os.getenv("AZURE_CHAT_VERSION") - -embed_base = os.getenv("AZURE_EMBED_BASE") -embed_token = Secret.from_env_var("AZURE_EMBED_KEY") -embed_version = os.getenv("AZURE_EMBED_VERSION") - @component -class AsyncAzureGenerator(AzureOpenAIGenerator): +class AsyncGenerator(AzureOpenAIGenerator): def __init__( self, - api_key: Secret = chat_token, + api_key: Secret = Secret.from_env_var("LLM_AZURE_OPENAI_API_KEY"), model: str = "gpt-4-turbo", - api_base: str = chat_base, - api_version: str = chat_version, + api_base: str = os.getenv("LLM_AZURE_OPENAI_API_BASE"), + api_version: str = os.getenv("LLM_AZURE_OPENAI_VERSION"), streaming_callback: Optional[Callable[[StreamingChunk], None]] = None, system_prompt: Optional[str] = None, generation_kwargs: Optional[Dict[str, Any]] = None, ): - super(AsyncAzureGenerator, self).__init__( + super(AsyncGenerator, self).__init__( azure_endpoint=api_base, api_version=api_version, azure_deployment=model, @@ -62,9 +49,9 @@ class AsyncAzureGenerator(AzureOpenAIGenerator): ) self.client = AsyncAzureOpenAI( - azure_endpoint=chat_base, + azure_endpoint=api_base, api_version=api_version, - api_key=chat_token.resolve_value(), + api_key=api_key.resolve_value(), ) @component.output_types(replies=List[str], meta=List[Dict[str, Any]]) @@ -125,206 +112,34 @@ class AsyncAzureGenerator(AzureOpenAIGenerator): } -@component -class AsyncAzureTextEmbedder(AzureOpenAITextEmbedder): - def __init__( - self, - api_key: Secret = embed_token, - model: str = "text-embedding-3-small", - dimensions: Optional[int] = None, - api_base_url: Optional[str] = None, - api_version: Optional[str] = None, - organization: Optional[str] = None, - prefix: str = "", - suffix: str = "", - ): - super(AsyncAzureTextEmbedder, self).__init__( - azure_endpoint=api_base_url, - api_version=api_version, - azure_deployment=model, - dimensions=dimensions, - api_key=api_key, - organization=organization, - prefix=prefix, - suffix=suffix, - ) - - self.client = AsyncAzureOpenAI( - azure_endpoint=api_base_url, - api_version=api_version, - api_key=api_key.resolve_value(), - ) - - @component.output_types(embedding=List[float], meta=Dict[str, Any]) - @backoff.on_exception(backoff.expo, openai.RateLimitError, max_time=60, max_tries=3) - async def run(self, text: str): - if not isinstance(text, str): - raise TypeError( - "AzureOpenAITextEmbedder expects a string as an input." - "In case you want to embed a list of Documents, please use the AzureOpenAIDocumentEmbedder." - ) - - logger.info(f"Running Async Azure OpenAI text embedder with text: {text}") - - text_to_embed = self.prefix + text + self.suffix - - # 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", " ") - - if self.dimensions is not None: - response = await self.client.embeddings.create( - model=self.azure_deployment, - dimensions=self.dimensions, - input=text_to_embed, - ) - else: - response = await self.client.embeddings.create( - model=self.azure_deployment, input=text_to_embed - ) - - meta = {"model": response.model, "usage": dict(response.usage)} - - return {"embedding": response.data[0].embedding, "meta": meta} - - -@component -class AsyncAzureDocumentEmbedder(AzureOpenAIDocumentEmbedder): - def __init__( - self, - api_key: Secret = embed_token, - model: str = "text-embedding-3-small", - dimensions: Optional[int] = None, - api_base_url: Optional[str] = None, - api_version: Optional[str] = None, - organization: Optional[str] = None, - prefix: str = "", - suffix: str = "", - batch_size: int = 32, - progress_bar: bool = True, - meta_fields_to_embed: Optional[List[str]] = None, - embedding_separator: str = "\n", - ): - super(AsyncAzureDocumentEmbedder, self).__init__( - azure_endpoint=api_base_url, - api_version=api_version, - azure_deployment=model, - dimensions=dimensions, - api_key=api_key, - organization=organization, - prefix=prefix, - suffix=suffix, - batch_size=batch_size, - progress_bar=progress_bar, - meta_fields_to_embed=meta_fields_to_embed, - embedding_separator=embedding_separator, - ) - - self.client = AsyncAzureOpenAI( - azure_endpoint=api_base_url, - api_version=api_version, - api_key=api_key.resolve_value(), - ) - - async def _embed_batch( - self, texts_to_embed: List[str], batch_size: int - ) -> 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 self.progress_bar, - desc="Calculating embeddings", - ): - batch = texts_to_embed[i : i + batch_size] - if self.dimensions is not None: - response = await self.client.embeddings.create( - model=self.azure_deployment, dimensions=self.dimensions, input=batch - ) - else: - response = await self.client.embeddings.create( - model=self.azure_deployment, input=batch - ) - 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) - 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]): - if ( - not isinstance(documents, list) - or documents - and not isinstance(documents[0], Document) - ): - raise TypeError( - "AzureOpenAIDocumentEmbedder expects a list of Documents as input." - "In case you want to embed a string, please use the AzureOpenAITextEmbedder." - ) - - logger.info( - f"Running Async OpenAI document embedder with documents: {documents}" - ) - - texts_to_embed = self._prepare_texts_to_embed(documents=documents) - - embeddings, meta = await self._embed_batch( - texts_to_embed=texts_to_embed, batch_size=self.batch_size - ) - - for doc, emb in zip(documents, embeddings): - doc.embedding = emb - - return {"documents": documents, "meta": meta} - - -@provider("azure_openai") +@provider("azure_openai_llm") class AzureOpenAILLMProvider(LLMProvider): def __init__( self, - chat_api_key: Secret = Secret.from_env_var("AZURE_CHAT_KEY"), - chat_api_base: str = os.getenv("AZURE_CHAT_BASE"), - embed_api_key: Secret = Secret.from_env_var("AZURE_EMBED_KEY"), - embed_api_base: str = os.getenv("AZURE_EMBED_BASE"), - chat_api_version: str = os.getenv("AZURE_CHAT_VERSION"), - embed_api_version: str = os.getenv("AZURE_EMBED_VERSION"), - generation_model: str = os.getenv("AZURE_GENERATION_MODEL") - or AZURE_GENERATION_MODEL, - embedding_model: str = os.getenv("EMBEDDING_MODEL") or EMBEDDING_MODEL_NAME, - embedding_model_dim: int = ( - int(os.getenv("EMBEDDING_MODEL_DIMENSION")) - if os.getenv("EMBEDDING_MODEL_DIMENSION") - else 0 - ) - or EMBEDDING_MODEL_DIMENSION, + chat_api_key: Secret = Secret.from_env_var("LLM_AZURE_OPENAI_API_KEY"), + chat_api_base: str = os.getenv("LLM_AZURE_OPENAI_API_BASE"), + chat_api_version: str = os.getenv("LLM_AZURE_OPENAI_VERSION"), + generation_model: str = os.getenv("GENERATION_MODEL") or GENERATION_MODEL, ): - logger.info(f"Using Azure OpenAI Generation Model: {generation_model}") + logger.info(f"Using Azure OpenAI LLM: {generation_model}") self._generation_api_key = chat_api_key self._generation_api_base = chat_api_base self._generation_api_version = chat_api_version self._generation_model = generation_model - self._embedding_api_base = embed_api_base - self._embedding_api_key = embed_api_key - self._embedding_api_version = embed_api_version - self._embedding_model = embedding_model - self._embedding_model_dim = embedding_model_dim def get_generator( self, - model_kwargs: Optional[Dict[str, Any]] = AZURE_GENERATION_MODEL_KWARGS, + model_kwargs: Dict[str, Any] = orjson.loads( + os.getenv("GENERATION_MODEL_KWARGS", "{}") + ) + or GENERATION_MODEL_KWARGS, system_prompt: Optional[str] = None, ): - return AsyncAzureGenerator( + logger.info( + f"Creating Azure OpenAI generator with model kwargs: {model_kwargs}" + ) + return AsyncGenerator( api_key=self._generation_api_key, model=self._generation_model, api_base=self._generation_api_base, @@ -332,21 +147,3 @@ class AzureOpenAILLMProvider(LLMProvider): system_prompt=system_prompt, generation_kwargs=model_kwargs, ) - - def get_text_embedder(self): - return AsyncAzureTextEmbedder( - api_key=self._embedding_api_key, - model=self._embedding_model, - dimensions=self._embedding_model_dim, - api_base_url=self._embedding_api_base, - api_version=self._embedding_api_version, - ) - - def get_document_embedder(self): - return AsyncAzureDocumentEmbedder( - api_key=self._embedding_api_key, - model=self._embedding_model, - dimensions=self._embedding_model_dim, - api_base_url=self._embedding_api_base, - api_version=self._embedding_api_version, - ) diff --git a/wren-ai-service/src/providers/llm/ollama.py b/wren-ai-service/src/providers/llm/ollama.py index 42418d245..bacaf5706 100644 --- a/wren-ai-service/src/providers/llm/ollama.py +++ b/wren-ai-service/src/providers/llm/ollama.py @@ -1,30 +1,23 @@ import logging import os -import time from typing import Any, Callable, Dict, List, Optional import aiohttp -from haystack import Document, component +import orjson +from haystack import component from haystack.dataclasses import StreamingChunk -from haystack_integrations.components.embedders.ollama import ( - OllamaDocumentEmbedder, - OllamaTextEmbedder, -) from haystack_integrations.components.generators.ollama import OllamaGenerator -from tqdm import tqdm from src.core.provider import LLMProvider from src.providers.loader import provider logger = logging.getLogger("wren-ai-service") -OLLAMA_URL = "http://localhost:11434" -GENERATION_MODEL_NAME = "llama3:70b" +LLM_OLLAMA_URL = "http://localhost:11434" +GENERATION_MODEL = "llama3:70b" GENERATION_MODEL_KWARGS = { "temperature": 0, } -EMBEDDING_MODEL_NAME = "nomic-embed-text" -EMBEDDING_MODEL_DIMENSION = 768 # https://huggingface.co/nomic-ai/nomic-embed-text-v1.5 @component @@ -126,182 +119,29 @@ class AsyncGenerator(OllamaGenerator): return await self._convert_to_response(response) -@component -class AsyncTextEmbedder(OllamaTextEmbedder): - def __init__( - self, - model: str = "nomic-embed-text", - url: str = "http://localhost:11434/api/embeddings", - generation_kwargs: Optional[Dict[str, Any]] = None, - timeout: int = 120, - ): - super(AsyncTextEmbedder, self).__init__( - model=model, - url=url, - generation_kwargs=generation_kwargs, - timeout=timeout, - ) - - @component.output_types(embedding=List[float], meta=Dict[str, Any]) - async def run( - self, - text: str, - generation_kwargs: Optional[Dict[str, Any]] = None, - ): - logger.debug(f"Running Ollama text embedder with text: {text}") - - payload = self._create_json_payload(text, generation_kwargs) - - start = time.perf_counter() - async with aiohttp.ClientSession( - timeout=aiohttp.ClientTimeout(self.timeout) - ) as session: - async with session.post( - self.url, - json=payload, - ) as response: - elapsed = time.perf_counter() - start - result = await response.json() - - result["meta"] = {"model": self.model, "duration": elapsed} - - return result - - -@component -class AsyncDocumentEmbedder(OllamaDocumentEmbedder): - def __init__( - self, - model: str = "nomic-embed-text", - url: str = "http://localhost:11434/api/embeddings", - generation_kwargs: Optional[Dict[str, Any]] = None, - timeout: int = 120, - prefix: str = "", - suffix: str = "", - progress_bar: bool = True, - meta_fields_to_embed: Optional[List[str]] = None, - embedding_separator: str = "\n", - ): - super(AsyncDocumentEmbedder, self).__init__( - model=model, - url=url, - generation_kwargs=generation_kwargs, - timeout=timeout, - prefix=prefix, - suffix=suffix, - progress_bar=progress_bar, - meta_fields_to_embed=meta_fields_to_embed, - embedding_separator=embedding_separator, - ) - - async def _embed_batch( - self, - texts_to_embed: List[str], - batch_size: int, - generation_kwargs: Optional[Dict[str, Any]] = None, - ): - """ - Ollama Embedding only allows single uploads, not batching. Currently the batch size is set to 1. - If this changes in the future, line 86 (the first line within the for loop), can contain: - batch = texts_to_embed[i + i + batch_size] - """ - - all_embeddings = [] - meta: Dict[str, Any] = {"model": self.model} - - async with aiohttp.ClientSession( - timeout=aiohttp.ClientTimeout(self.timeout) - ) as session: - for i in tqdm( - range(0, len(texts_to_embed), batch_size), - disable=not self.progress_bar, - desc="Calculating embeddings", - ): - batch = texts_to_embed[i] # Single batch only - payload = self._create_json_payload(batch, generation_kwargs) - - async with session.post( - self.url, - json=payload, - ) as response: - result = await response.json() - all_embeddings.append(result["embedding"]) - - return all_embeddings, meta - - @component.output_types(embedding=List[float], meta=Dict[str, Any]) - async def run( - self, - documents: List[str], - generation_kwargs: Optional[Dict[str, Any]] = None, - ): - logger.debug(f"Running Ollama document embedder with documents: {documents}") - - if ( - not isinstance(documents, list) - or documents - and not isinstance(documents[0], Document) - ): - msg = ( - "OllamaDocumentEmbedder expects a list of Documents as input." - "In case you want to embed a list of strings, please use the OllamaTextEmbedder." - ) - raise TypeError(msg) - - texts_to_embed = self._prepare_texts_to_embed(documents=documents) - embeddings, meta = await self._embed_batch( - texts_to_embed=texts_to_embed, - batch_size=self.batch_size, - generation_kwargs=generation_kwargs, - ) - - for doc, emb in zip(documents, embeddings): - doc.embedding = emb - - return {"documents": documents, "meta": meta} - - -@provider("ollama") +@provider("ollama_llm") class OllamaLLMProvider(LLMProvider): def __init__( self, - url: str = os.getenv("OLLAMA_URL") or OLLAMA_URL, - generation_model: str = os.getenv("GENERATION_MODEL") or GENERATION_MODEL_NAME, - embedding_model: str = os.getenv("EMBEDDING_MODEL") or EMBEDDING_MODEL_NAME, + url: str = os.getenv("LLM_OLLAMA_URL") or LLM_OLLAMA_URL, + generation_model: str = os.getenv("GENERATION_MODEL") or GENERATION_MODEL, ): - logger.info(f"Using Ollama Generation Model: {generation_model}") + logger.info(f"Using Ollama LLM: {generation_model}") self._url = url self._generation_model = generation_model - self._embedding_model = embedding_model def get_generator( self, - model_kwargs: Optional[Dict[str, Any]] = GENERATION_MODEL_KWARGS, + model_kwargs: Dict[str, Any] = orjson.loads( + os.getenv("GENERATION_MODEL_KWARGS", "{}") + ) + or GENERATION_MODEL_KWARGS, system_prompt: Optional[str] = None, ): + logger.info(f"Creating Ollama generator with model kwargs: {model_kwargs}") return AsyncGenerator( model=self._generation_model, url=f"{self._url}/api/generate", generation_kwargs=model_kwargs, system_prompt=system_prompt, ) - - def get_text_embedder( - self, - model_kwargs: Optional[Dict[str, Any]] = None, - ): - return AsyncTextEmbedder( - model=self._embedding_model, - url=f"{self._url}/api/embeddings", - generation_kwargs=model_kwargs, - ) - - def get_document_embedder( - self, - model_kwargs: Optional[Dict[str, Any]] = None, - ): - return AsyncDocumentEmbedder( - model=self._embedding_model, - url=f"{self._url}/api/embeddings", - generation_kwargs=model_kwargs, - ) diff --git a/wren-ai-service/src/providers/llm/openai.py b/wren-ai-service/src/providers/llm/openai.py index 163b01652..c63405936 100644 --- a/wren-ai-service/src/providers/llm/openai.py +++ b/wren-ai-service/src/providers/llm/openai.py @@ -1,40 +1,37 @@ import logging import os -from typing import Any, Callable, Dict, List, Optional, Tuple, Union +from typing import Any, Callable, Dict, List, Optional, Union import backoff import openai -from haystack import Document, component -from haystack.components.embedders import OpenAIDocumentEmbedder, OpenAITextEmbedder +import orjson +from haystack import component from haystack.components.generators import OpenAIGenerator from haystack.dataclasses import ChatMessage, StreamingChunk from haystack.utils import Secret from openai import AsyncOpenAI, OpenAI, Stream from openai.types.chat import ChatCompletion, ChatCompletionChunk -from tqdm import tqdm from src.core.provider import LLMProvider from src.providers.loader import provider logger = logging.getLogger("wren-ai-service") -OPENAI_API_BASE = "https://api.openai.com/v1" -GENERATION_MODEL_NAME = "gpt-3.5-turbo" +LLM_OPENAI_API_BASE = "https://api.openai.com/v1" +GENERATION_MODEL = "gpt-3.5-turbo" GENERATION_MODEL_KWARGS = { "temperature": 0, "n": 1, "max_tokens": 4096, "response_format": {"type": "json_object"}, } -EMBEDDING_MODEL_NAME = "text-embedding-3-large" -EMBEDDING_MODEL_DIMENSION = 3072 @component class AsyncGenerator(OpenAIGenerator): def __init__( self, - api_key: Secret = Secret.from_env_var("OPENAI_API_KEY"), + api_key: Secret = Secret.from_env_var("LLM_OPENAI_API_KEY"), model: str = "gpt-3.5-turbo", streaming_callback: Optional[Callable[[StreamingChunk], None]] = None, api_base_url: Optional[str] = None, @@ -116,174 +113,13 @@ class AsyncGenerator(OpenAIGenerator): } -@component -class AsyncTextEmbedder(OpenAITextEmbedder): - def __init__( - self, - api_key: Secret = Secret.from_env_var("OPENAI_API_KEY"), - model: str = "text-embedding-ada-002", - dimensions: Optional[int] = None, - api_base_url: Optional[str] = None, - organization: Optional[str] = None, - prefix: str = "", - suffix: str = "", - ): - super(AsyncTextEmbedder, self).__init__( - api_key, - model, - dimensions, - api_base_url, - organization, - prefix, - suffix, - ) - self.client = AsyncOpenAI( - api_key=api_key.resolve_value(), - organization=organization, - base_url=api_base_url, - ) - - @component.output_types(embedding=List[float], meta=Dict[str, Any]) - @backoff.on_exception(backoff.expo, openai.RateLimitError, max_time=60, max_tries=3) - async def run(self, text: str): - if not isinstance(text, str): - raise TypeError( - "OpenAITextEmbedder expects a string as an input." - "In case you want to embed a list of Documents, please use the OpenAIDocumentEmbedder." - ) - - logger.debug(f"Running Async OpenAI text embedder with text: {text}") - - text_to_embed = self.prefix + text + self.suffix - - # 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", " ") - - if self.dimensions is not None: - response = await self.client.embeddings.create( - model=self.model, dimensions=self.dimensions, input=text_to_embed - ) - else: - response = await self.client.embeddings.create( - model=self.model, input=text_to_embed - ) - - meta = {"model": response.model, "usage": dict(response.usage)} - - return {"embedding": response.data[0].embedding, "meta": meta} - - -@component -class AsyncDocumentEmbedder(OpenAIDocumentEmbedder): - def __init__( - self, - api_key: Secret = Secret.from_env_var("OPENAI_API_KEY"), - model: str = "text-embedding-ada-002", - dimensions: Optional[int] = None, - api_base_url: Optional[str] = None, - organization: Optional[str] = None, - prefix: str = "", - suffix: str = "", - batch_size: int = 32, - progress_bar: bool = True, - meta_fields_to_embed: Optional[List[str]] = None, - embedding_separator: str = "\n", - ): - super(AsyncDocumentEmbedder, self).__init__( - api_key, - model, - dimensions, - api_base_url, - organization, - prefix, - suffix, - batch_size, - progress_bar, - meta_fields_to_embed, - embedding_separator, - ) - self.client = AsyncOpenAI( - api_key=api_key.resolve_value(), - organization=organization, - base_url=api_base_url, - ) - - async def _embed_batch( - self, texts_to_embed: List[str], batch_size: int - ) -> 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 self.progress_bar, - desc="Calculating embeddings", - ): - batch = texts_to_embed[i : i + batch_size] - if self.dimensions is not None: - response = await self.client.embeddings.create( - model=self.model, dimensions=self.dimensions, input=batch - ) - else: - response = await self.client.embeddings.create( - model=self.model, input=batch - ) - 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) - 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]): - if ( - not isinstance(documents, list) - or documents - and not isinstance(documents[0], Document) - ): - raise TypeError( - "OpenAIDocumentEmbedder expects a list of Documents as input." - "In case you want to embed a string, please use the OpenAITextEmbedder." - ) - - logger.debug( - f"Running Async OpenAI document embedder with documents: {documents}" - ) - - texts_to_embed = self._prepare_texts_to_embed(documents=documents) - - embeddings, meta = await self._embed_batch( - texts_to_embed=texts_to_embed, batch_size=self.batch_size - ) - - for doc, emb in zip(documents, embeddings): - doc.embedding = emb - - return {"documents": documents, "meta": meta} - - -@provider("openai") +@provider("openai_llm") class OpenAILLMProvider(LLMProvider): def __init__( self, - api_key: Secret = Secret.from_env_var("OPENAI_API_KEY"), - api_base: str = os.getenv("OPENAI_API_BASE") or OPENAI_API_BASE, - embedding_model: str = os.getenv("EMBEDDING_MODEL") or EMBEDDING_MODEL_NAME, - embedding_model_dim: int = ( - int(os.getenv("EMBEDDING_MODEL_DIMENSION")) - if os.getenv("EMBEDDING_MODEL_DIMENSION") - else 0 - ) - or EMBEDDING_MODEL_DIMENSION, - generation_model: str = os.getenv("GENERATION_MODEL") or GENERATION_MODEL_NAME, + api_key: Secret = Secret.from_env_var("LLM_OPENAI_API_KEY"), + api_base: str = os.getenv("LLM_OPENAI_API_BASE") or LLM_OPENAI_API_BASE, + generation_model: str = os.getenv("GENERATION_MODEL") or GENERATION_MODEL, ): def _verify_api_key(api_key: str, api_base: str) -> None: """ @@ -293,24 +129,30 @@ class OpenAILLMProvider(LLMProvider): logger.info(f"Initializing OpenAILLM provider with API base: {api_base}") # TODO: currently only OpenAI api key can be verified - if api_base == OPENAI_API_BASE: + if api_base == LLM_OPENAI_API_BASE: _verify_api_key(api_key.resolve_value(), api_base) - logger.info(f"Using OpenAI Generation Model: {generation_model}") + logger.info(f"Using OpenAI LLM: {generation_model}") else: - logger.info( - f"Using OpenAI API-compatible Generation Model: {generation_model}" - ) + logger.info(f"Using OpenAI API-compatible LLM: {generation_model}") self._api_key = api_key self._api_base = api_base - self._embedding_model = embedding_model - self._embedding_model_dim = embedding_model_dim self._generation_model = generation_model def get_generator( self, - model_kwargs: Optional[Dict[str, Any]] = GENERATION_MODEL_KWARGS, + model_kwargs: Dict[str, Any] = orjson.loads( + os.getenv("GENERATION_MODEL_KWARGS", "{}") + ) + or GENERATION_MODEL_KWARGS, system_prompt: Optional[str] = None, ): + if self._api_base == LLM_OPENAI_API_BASE: + logger.info(f"Creating OpenAI generator with model kwargs: {model_kwargs}") + else: + logger.info( + f"Creating OpenAI API-compatible generator with model kwargs: {model_kwargs}" + ) + return AsyncGenerator( api_key=self._api_key, api_base_url=self._api_base, @@ -318,19 +160,3 @@ class OpenAILLMProvider(LLMProvider): system_prompt=system_prompt, generation_kwargs=model_kwargs, ) - - def get_text_embedder(self): - return AsyncTextEmbedder( - api_key=self._api_key, - api_base_url=self._api_base, - model=self._embedding_model, - dimensions=self._embedding_model_dim, - ) - - def get_document_embedder(self): - return AsyncDocumentEmbedder( - api_key=self._api_key, - api_base_url=self._api_base, - model=self._embedding_model, - dimensions=self._embedding_model_dim, - ) diff --git a/wren-ai-service/src/providers/loader.py b/wren-ai-service/src/providers/loader.py index 3483519d9..c01b5a8a4 100644 --- a/wren-ai-service/src/providers/loader.py +++ b/wren-ai-service/src/providers/loader.py @@ -93,7 +93,8 @@ def get_provider(name: str): return PROVIDERS[name] -def get_default_embedding_model_dim(llm_provider: str): +def get_default_embedding_model_dim(embedder_provider: str): + file_name = embedder_provider.split("_embedder")[0] return importlib.import_module( - f"src.providers.llm.{llm_provider}" + f"src.providers.embedder.{file_name}" ).EMBEDDING_MODEL_DIMENSION diff --git a/wren-ai-service/src/utils.py b/wren-ai-service/src/utils.py index 240402070..0836088a2 100644 --- a/wren-ai-service/src/utils.py +++ b/wren-ai-service/src/utils.py @@ -8,7 +8,7 @@ from typing import Tuple from dotenv import load_dotenv from src.core.engine import Engine -from src.core.provider import DocumentStoreProvider, LLMProvider +from src.core.provider import DocumentStoreProvider, EmbedderProvider, LLMProvider from src.providers import loader logger = logging.getLogger("wren-ai-service") @@ -53,23 +53,26 @@ def load_env_vars() -> str: if is_dev_env := os.getenv("ENV") and os.getenv("ENV").lower() == "dev": load_dotenv(".env.dev", override=True) - else: - load_dotenv(".env.prod", override=True) return "dev" if is_dev_env else "prod" -def init_providers() -> Tuple[LLMProvider, DocumentStoreProvider, Engine]: +def init_providers() -> ( + Tuple[LLMProvider, EmbedderProvider, DocumentStoreProvider, Engine] +): logger.info("Initializing providers...") loader.import_mods() - llm_provider = loader.get_provider(os.getenv("LLM_PROVIDER", "openai"))() + llm_provider = loader.get_provider(os.getenv("LLM_PROVIDER", "openai_llm"))() + embedder_provider = loader.get_provider( + os.getenv("EMBEDDER_PROVIDER", "openai_embedder") + )() document_store_provider = loader.get_provider( os.getenv("DOCUMENT_STORE_PROVIDER", "qdrant") )() - engine = loader.get_provider(os.getenv("ENGINE", "wren-ui"))() + engine = loader.get_provider(os.getenv("ENGINE", "wren_ui"))() - return llm_provider, document_store_provider, engine + return llm_provider, embedder_provider, document_store_provider, engine def timer(func): diff --git a/wren-ai-service/tests/pytest/pipelines/test_ask.py b/wren-ai-service/tests/pytest/pipelines/test_ask.py index 75ab9451b..f88f40c6b 100644 --- a/wren-ai-service/tests/pytest/pipelines/test_ask.py +++ b/wren-ai-service/tests/pytest/pipelines/test_ask.py @@ -4,7 +4,7 @@ import orjson import pytest from src.core.pipeline import async_validate -from src.core.provider import DocumentStoreProvider, LLMProvider +from src.core.provider import DocumentStoreProvider, EmbedderProvider from src.pipelines.ask.followup_generation import FollowUpGeneration from src.pipelines.ask.generation import Generation from src.pipelines.ask.retrieval import Retrieval @@ -26,24 +26,31 @@ def mdl_str(): @pytest.fixture def llm_provider(): - llm_provider, _, _ = init_providers() + llm_provider, _, _, _ = init_providers() return llm_provider +@pytest.fixture +def embedder_provider(): + _, embedder_provider, _, _ = init_providers() + + return embedder_provider + + @pytest.fixture def document_store_provider(): - _, document_store_provider, _ = init_providers() + _, _, document_store_provider, _ = init_providers() return document_store_provider def test_clear_documents(mdl_str: str): - llm_provider, document_store_provider, _ = init_providers() + _, embedder_provider, document_store_provider, _ = init_providers() store = document_store_provider.get_store() indexing_pipeline = Indexing( - llm_provider=llm_provider, + embedder_provider=embedder_provider, document_store_provider=document_store_provider, ) @@ -78,11 +85,11 @@ def test_clear_documents(mdl_str: str): def test_indexing_pipeline( mdl_str: str, - llm_provider: LLMProvider, + embedder_provider: EmbedderProvider, document_store_provider: DocumentStoreProvider, ): indexing_pipeline = Indexing( - llm_provider=llm_provider, + embedder_provider=embedder_provider, document_store_provider=document_store_provider, ) @@ -98,11 +105,11 @@ def test_indexing_pipeline( def test_retrieval_pipeline( - llm_provider: LLMProvider, + embedder_provider: EmbedderProvider, document_store_provider: DocumentStoreProvider, ): retrieval_pipeline = Retrieval( - llm_provider=llm_provider, + embedder_provider=embedder_provider, document_store_provider=document_store_provider, ) @@ -119,7 +126,7 @@ def test_retrieval_pipeline( def test_generation_pipeline(): - llm_provider, _, engine = init_providers() + llm_provider, _, _, engine = init_providers() generation_pipeline = Generation(llm_provider=llm_provider, engine=engine) generation_result = async_validate( lambda: generation_pipeline.run( @@ -146,7 +153,7 @@ def test_generation_pipeline(): def test_followup_generation_pipeline(): - llm_provider, _, engine = init_providers() + llm_provider, _, _, engine = init_providers() generation_pipeline = FollowUpGeneration(llm_provider=llm_provider, engine=engine) generation_result = async_validate( lambda: generation_pipeline.run( @@ -172,7 +179,7 @@ def test_followup_generation_pipeline(): def test_sql_correction_pipeline(): - llm_provider, _, engine = init_providers() + llm_provider, _, _, engine = init_providers() sql_correction_pipeline = SQLCorrection(llm_provider=llm_provider, engine=engine) sql_correction_result = async_validate( diff --git a/wren-ai-service/tests/pytest/pipelines/test_ask_details.py b/wren-ai-service/tests/pytest/pipelines/test_ask_details.py index c9e3a489a..baf39ea99 100644 --- a/wren-ai-service/tests/pytest/pipelines/test_ask_details.py +++ b/wren-ai-service/tests/pytest/pipelines/test_ask_details.py @@ -4,7 +4,7 @@ from src.utils import init_providers def test_generation_pipeline_producing_executable_sqls(): - llm_provider, _, engine = init_providers() + llm_provider, _, _, engine = init_providers() generation_pipeline = Generation( llm_provider=llm_provider, engine=engine, diff --git a/wren-ai-service/tests/pytest/pipelines/test_document_cleaner.py b/wren-ai-service/tests/pytest/pipelines/test_document_cleaner.py index 37b9f05bc..afdb85412 100644 --- a/wren-ai-service/tests/pytest/pipelines/test_document_cleaner.py +++ b/wren-ai-service/tests/pytest/pipelines/test_document_cleaner.py @@ -6,7 +6,7 @@ from src.utils import init_providers def _mock_store(name: str = "default") -> DocumentStore: - _, document_store_provider, _ = init_providers() + _, _, document_store_provider, _ = init_providers() store = document_store_provider.get_store( embedding_model_dim=5, dataset_name=name, diff --git a/wren-ai-service/tests/pytest/providers/test_loader.py b/wren-ai-service/tests/pytest/providers/test_loader.py index 7598a5c69..50deb557b 100644 --- a/wren-ai-service/tests/pytest/providers/test_loader.py +++ b/wren-ai-service/tests/pytest/providers/test_loader.py @@ -3,29 +3,39 @@ from src.providers import loader def test_import_mods(): loader.import_mods("src.providers") - assert len(loader.PROVIDERS) == 6 + assert len(loader.PROVIDERS) == 9 def test_get_provider(): loader.import_mods("src.providers") # llm provider - provider = loader.get_provider("openai") + provider = loader.get_provider("openai_llm") assert provider.__name__ == "OpenAILLMProvider" - provider = loader.get_provider("azure_openai") + provider = loader.get_provider("azure_openai_llm") assert provider.__name__ == "AzureOpenAILLMProvider" - provider = loader.get_provider("ollama") + provider = loader.get_provider("ollama_llm") assert provider.__name__ == "OllamaLLMProvider" + # embedder provider + provider = loader.get_provider("openai_embedder") + assert provider.__name__ == "OpenAIEmbedderProvider" + + provider = loader.get_provider("azure_openai_embedder") + assert provider.__name__ == "AzureOpenAIEmbedderProvider" + + provider = loader.get_provider("ollama_embedder") + assert provider.__name__ == "OllamaEmbedderProvider" + # document store provider provider = loader.get_provider("qdrant") assert provider.__name__ == "QdrantProvider" # engine provider - provider = loader.get_provider("wren-ui") + provider = loader.get_provider("wren_ui") assert provider.__name__ == "WrenUI" - provider = loader.get_provider("wren-ibis") + provider = loader.get_provider("wren_ibis") assert provider.__name__ == "WrenIbis" diff --git a/wren-ai-service/tests/pytest/services/test_ask.py b/wren-ai-service/tests/pytest/services/test_ask.py index ea38c196e..b488c793a 100644 --- a/wren-ai-service/tests/pytest/services/test_ask.py +++ b/wren-ai-service/tests/pytest/services/test_ask.py @@ -23,20 +23,20 @@ from src.web.v1.services.ask import ( @pytest.fixture def ask_service(): - llm_provider, document_store_provider, engine = init_providers() + llm_provider, embedder_provider, document_store_provider, engine = init_providers() return AskService( { "indexing": indexing.Indexing( - llm_provider=llm_provider, + embedder_provider=embedder_provider, document_store_provider=document_store_provider, ), "retrieval": retrieval.Retrieval( - llm_provider=llm_provider, + embedder_provider=embedder_provider, document_store_provider=document_store_provider, ), "historical_question": historical_question.HistoricalQuestion( - llm_provider=llm_provider, + embedder_provider=embedder_provider, store_provider=document_store_provider, ), "generation": generation.Generation( diff --git a/wren-ai-service/tests/pytest/services/test_ask_details.py b/wren-ai-service/tests/pytest/services/test_ask_details.py index db2b444a3..3f1a04317 100644 --- a/wren-ai-service/tests/pytest/services/test_ask_details.py +++ b/wren-ai-service/tests/pytest/services/test_ask_details.py @@ -14,7 +14,7 @@ from src.web.v1.services.ask_details import ( @pytest.fixture def ask_details_service(): - llm_provider, _, engine = init_providers() + llm_provider, _, _, engine = init_providers() return AskDetailsService( { "generation": generation.Generation( diff --git a/wren-ai-service/tests/pytest/services/test_semantics.py b/wren-ai-service/tests/pytest/services/test_semantics.py index 0481b18b1..2455c5a71 100644 --- a/wren-ai-service/tests/pytest/services/test_semantics.py +++ b/wren-ai-service/tests/pytest/services/test_semantics.py @@ -9,12 +9,13 @@ from src.web.v1.services.semantics import ( @pytest.fixture def semantics_service(): - llm_provider, document_store_provider, _ = init_providers() + llm_provider, embedder_provider, document_store_provider, _ = init_providers() return SemanticsService( pipelines={ "generate_description": description.Generation( llm_provider=llm_provider, + embedder_provider=embedder_provider, document_store_provider=document_store_provider, ), } diff --git a/wren-launcher/utils/docker.go b/wren-launcher/utils/docker.go index 2f71df269..72dbb88b5 100644 --- a/wren-launcher/utils/docker.go +++ b/wren-launcher/utils/docker.go @@ -38,9 +38,13 @@ func replaceEnvFileContent(content string, projectDir string, openaiApiKey strin reg := regexp.MustCompile(`PROJECT_DIR=(.*)`) str := reg.ReplaceAllString(content, "PROJECT_DIR="+projectDir) - // replace OPENAI_API_KEY - reg = regexp.MustCompile(`OPENAI_API_KEY=(.*)`) - str = reg.ReplaceAllString(str, "OPENAI_API_KEY="+openaiApiKey) + // replace LLM_OPENAI_API_KEY + reg = regexp.MustCompile(`LLM_OPENAI_API_KEY=(.*)`) + str = reg.ReplaceAllString(str, "LLM_OPENAI_API_KEY="+openaiApiKey) + + // replace EMBEDDER_OPENAI_API_KEY + reg = regexp.MustCompile(`EMBEDDER_OPENAI_API_KEY=(.*)`) + str = reg.ReplaceAllString(str, "EMBEDDER_OPENAI_API_KEY="+openaiApiKey) // replace GENERATION_MODEL reg = regexp.MustCompile(`GENERATION_MODEL=(.*)`)