mirror of
https://github.com/Canner/WrenAI.git
synced 2026-09-24 23:29:49 +08:00
separate llm and embedder (#454)
This commit is contained in:
@@ -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
|
||||
|
||||
+28
-16
@@ -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
|
||||
QDRANT_HOST=qdrant
|
||||
|
||||
+2
-1
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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 ##
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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):
|
||||
...
|
||||
|
||||
@@ -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 (
|
||||
|
||||
@@ -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"))
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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": []}'
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
@@ -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,
|
||||
):
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
),
|
||||
}
|
||||
|
||||
@@ -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=(.*)`)
|
||||
|
||||
Reference in New Issue
Block a user