separate llm and embedder (#454)

This commit is contained in:
Chih-Yu Yeh
2024-07-03 10:52:38 +08:00
committed by GitHub
parent c433c1b257
commit a8dded673a
42 changed files with 889 additions and 814 deletions
+4 -2
View File
@@ -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
View File
@@ -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
View File
@@ -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
+4 -4
View File
@@ -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,
+4 -4
View File
@@ -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,
+29 -16
View File
@@ -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 -2
View File
@@ -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
-42
View File
@@ -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
-11
View File
@@ -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 ##
+2 -9
View File
@@ -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
+2
View File
@@ -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):
...
+1 -1
View File
@@ -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"))
+5 -4
View File
@@ -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,
)
+2 -2
View File
@@ -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
+27 -230
View File
@@ -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,
)
+13 -173
View File
@@ -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,
)
+24 -198
View File
@@ -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,
)
+3 -2
View File
@@ -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
+10 -7
View File
@@ -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,
),
}
+7 -3
View File
@@ -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=(.*)`)