From af556b0261bbbb0a0cb2f1e16e4a0f1fd2a2d39a Mon Sep 17 00:00:00 2001 From: Chih-Yu Yeh Date: Mon, 5 Aug 2024 13:43:01 +0800 Subject: [PATCH] feat(wren-ai-service): optimize qdrant for multi users setting (#558) * add binary quantization * allow multi users setting * user_id backward compatible * fix test * fix conflict * add user_id to indexing * remove print * rename user_id to id * upgrade qdrant to v1.10.1 * update * allow using qdrant cloud * update * enable binary quantization only if embedding model dim >= 1024 * fix tests * fix test * revert * fix mkdir * simple refactor * fix --- ....0.yaml => helm-values-qdrant_0.10.1.yaml} | 2 +- deployment/kustomizations/kustomization.yaml | 4 +- docker/docker-compose-dev.yaml | 2 +- docker/docker-compose.yaml | 2 +- wren-ai-service/.env.dev.example | 1 + wren-ai-service/README.md | 1 - .../src/pipelines/ask/retrieval.py | 24 ++++- .../src/pipelines/indexing/indexing.py | 48 ++++++---- .../src/providers/document_store/qdrant.py | 96 +++++++++++++++++-- wren-ai-service/src/web/v1/services/ask.py | 1 + .../src/web/v1/services/indexing.py | 5 +- wren-ai-service/tests/locust/locust_script.py | 2 +- .../tools/dev/docker-compose-dev.yaml | 2 +- 13 files changed, 154 insertions(+), 36 deletions(-) rename deployment/kustomizations/{helm-values-qdrant_0.9.0.yaml => helm-values-qdrant_0.10.1.yaml} (99%) diff --git a/deployment/kustomizations/helm-values-qdrant_0.9.0.yaml b/deployment/kustomizations/helm-values-qdrant_0.10.1.yaml similarity index 99% rename from deployment/kustomizations/helm-values-qdrant_0.9.0.yaml rename to deployment/kustomizations/helm-values-qdrant_0.10.1.yaml index 3a6f3f729..d9f0bd6ea 100644 --- a/deployment/kustomizations/helm-values-qdrant_0.9.0.yaml +++ b/deployment/kustomizations/helm-values-qdrant_0.10.1.yaml @@ -3,7 +3,7 @@ replicaCount: 1 image: repository: docker.io/qdrant/qdrant pullPolicy: IfNotPresent - tag: "v1.7.4" + tag: "v1.10.1" useUnprivilegedImage: false imagePullSecrets: [] diff --git a/deployment/kustomizations/kustomization.yaml b/deployment/kustomizations/kustomization.yaml index 66078bdaa..fe74811c4 100644 --- a/deployment/kustomizations/kustomization.yaml +++ b/deployment/kustomizations/kustomization.yaml @@ -27,9 +27,9 @@ helmCharts: - name: qdrant repo: https://qdrant.github.io/qdrant-helm - version: 0.9.0 + version: 0.10.1 releaseName: wren-qdrant - valuesFile: helm-values-qdrant_0.9.0.yaml + valuesFile: helm-values-qdrant_0.10.1.yaml includeCRDs: true # The Same Namespace namespace: wren diff --git a/docker/docker-compose-dev.yaml b/docker/docker-compose-dev.yaml index 9f8330089..a42c034fb 100644 --- a/docker/docker-compose-dev.yaml +++ b/docker/docker-compose-dev.yaml @@ -74,7 +74,7 @@ services: - wren qdrant: - image: qdrant/qdrant:v1.7.4 + image: qdrant/qdrant:v1.10.1 pull_policy: always ports: - 6333:6333 diff --git a/docker/docker-compose.yaml b/docker/docker-compose.yaml index 1b7c591c0..a85f24317 100644 --- a/docker/docker-compose.yaml +++ b/docker/docker-compose.yaml @@ -72,7 +72,7 @@ services: - qdrant qdrant: - image: qdrant/qdrant:v1.7.4 + image: qdrant/qdrant:v1.10.1 restart: on-failure expose: - 6333 diff --git a/wren-ai-service/.env.dev.example b/wren-ai-service/.env.dev.example index e5e7a15e6..ab1b3d314 100644 --- a/wren-ai-service/.env.dev.example +++ b/wren-ai-service/.env.dev.example @@ -41,6 +41,7 @@ EMBEDDER_OLLAMA_URL=http://localhost:11434 DOCUMENT_STORE_PROVIDER=qdrant QDRANT_HOST=http://localhost:6333 +QDRANT_API_KEY= ENGINE=wren_ui # wren_ui, wren_ibis, wren_engine diff --git a/wren-ai-service/README.md b/wren-ai-service/README.md index bc60b0a64..bc2bad846 100644 --- a/wren-ai-service/README.md +++ b/wren-ai-service/README.md @@ -42,7 +42,6 @@ For a comprehensive understanding of how to evaluate the pipelines, please refer - to run the load test - setup `DATASET_NAME` in `.env.dev` - adjust test config if needed - - adjust test config in pyproject.toml `tool.locust` section - adjust user count in `tests/locust/config_users.json` - in wren-ai-service folder, run `just up` to start the docker containers - in wren-ai-service folder, run `just start` to start the ai service diff --git a/wren-ai-service/src/pipelines/ask/retrieval.py b/wren-ai-service/src/pipelines/ask/retrieval.py index a7547a30f..cf0cf8a04 100644 --- a/wren-ai-service/src/pipelines/ask/retrieval.py +++ b/wren-ai-service/src/pipelines/ask/retrieval.py @@ -1,7 +1,7 @@ import logging import sys from pathlib import Path -from typing import Any +from typing import Any, Optional from hamilton import base from hamilton.experimental.h_async import AsyncDriver @@ -24,8 +24,21 @@ async def embedding(query: str, embedder: Any) -> dict: @async_timer @observe(capture_input=False) -async def retrieval(embedding: dict, retriever: Any) -> dict: - return await retriever.run(query_embedding=embedding.get("embedding")) +async def retrieval(embedding: dict, id: str, retriever: Any) -> dict: + filters = ( + { + "operator": "AND", + "conditions": [ + {"field": "id", "operator": "==", "value": id}, + ], + } + if id + else None + ) + + return await retriever.run( + query_embedding=embedding.get("embedding"), filters=filters + ) ## End of Pipeline @@ -49,6 +62,7 @@ class Retrieval(BasicPipeline): def visualize( self, query: str, + id: Optional[str] = None, ) -> None: destination = "outputs/pipelines/ask" if not Path(destination).exists(): @@ -59,6 +73,7 @@ class Retrieval(BasicPipeline): output_file_path=f"{destination}/retrieval.dot", inputs={ "query": query, + "id": id or "", "embedder": self._embedder, "retriever": self._retriever, }, @@ -68,12 +83,13 @@ class Retrieval(BasicPipeline): @async_timer @observe(name="Ask Retrieval") - async def run(self, query: str): + async def run(self, query: str, id: Optional[str] = None): logger.info("Ask Retrieval pipeline is running...") return await self._pipe.execute( ["retrieval"], inputs={ "query": query, + "id": id or "", "embedder": self._embedder, "retriever": self._retriever, }, diff --git a/wren-ai-service/src/pipelines/indexing/indexing.py b/wren-ai-service/src/pipelines/indexing/indexing.py index 46c64c335..e0ef16ffd 100644 --- a/wren-ai-service/src/pipelines/indexing/indexing.py +++ b/wren-ai-service/src/pipelines/indexing/indexing.py @@ -36,15 +36,27 @@ class DocumentCleaner: self._stores = stores @component.output_types(mdl=str) - async def run(self, mdl: str) -> str: - async def _clear_documents(store: DocumentStore) -> None: - document_count = await store.count_documents() + async def run(self, mdl: str, id: Optional[str] = None) -> str: + async def _clear_documents( + store: DocumentStore, id: Optional[str] = None + ) -> None: + filters = ( + { + "operator": "AND", + "conditions": [ + {"field": "id", "operator": "==", "value": id}, + ], + } + if id + else None + ) + document_count = await store.count_documents(filters=filters) ids = [str(i) for i in range(document_count)] if ids: await store.delete_documents(ids) logger.info("Ask Indexing pipeline is clearing old documents...") - await asyncio.gather(*[_clear_documents(store) for store in self._stores]) + await asyncio.gather(*[_clear_documents(store, id) for store in self._stores]) return {"mdl": mdl} @@ -87,7 +99,7 @@ class ViewConverter: """ @component.output_types(documents=List[Document]) - def run(self, mdl: Dict[str, Any]) -> None: + def run(self, mdl: Dict[str, Any], id: Optional[str] = None) -> None: def _format(view: Dict[str, Any]) -> List[str]: properties = view.get("properties", {}) return str( @@ -105,7 +117,7 @@ class ViewConverter: "documents": [ Document( id=str(i), - meta={"id": str(i)}, + meta={"id": id} if id else {}, content=converted_view, ) for i, converted_view in enumerate( @@ -121,7 +133,7 @@ class ViewConverter: @component class DDLConverter: @component.output_types(documents=List[Document]) - def run(self, mdl: Dict[str, Any]): + def run(self, mdl: Dict[str, Any], id: Optional[str] = None): logger.info("Ask Indexing pipeline is writing new documents...") logger.debug(f"original mdl_json: {mdl}") @@ -132,7 +144,7 @@ class DDLConverter: "documents": [ Document( id=str(i), - meta={"id": str(i)}, + meta={"id": id} if id else {}, content=ddl_command, ) for i, ddl_command in enumerate( @@ -336,10 +348,10 @@ class AsyncDocumentWriter(DocumentWriter): @async_timer @observe(capture_input=False, capture_output=False) async def clean_document_store( - mdl_str: str, cleaner: DocumentCleaner + mdl_str: str, cleaner: DocumentCleaner, id: Optional[str] = None ) -> Dict[str, Any]: logger.debug(f"input in clean_document_store: {mdl_str}") - return await cleaner.run(mdl=mdl_str) + return await cleaner.run(mdl=mdl_str, id=id) @timer @@ -358,11 +370,13 @@ def validate_mdl( @timer @observe(capture_input=False) -def convert_to_ddl(mdl: Dict[str, Any], ddl_converter: DDLConverter) -> Dict[str, Any]: +def convert_to_ddl( + mdl: Dict[str, Any], ddl_converter: DDLConverter, id: Optional[str] = None +) -> Dict[str, Any]: logger.debug( f"input in convert_to_ddl: {orjson.dumps(mdl, option=orjson.OPT_INDENT_2).decode()}" ) - return ddl_converter.run(mdl=mdl) + return ddl_converter.run(mdl=mdl, id=id) @async_timer @@ -385,12 +399,12 @@ async def write_ddl(embed_ddl: Dict[str, Any], ddl_writer: DocumentWriter) -> No @timer @observe(capture_input=False) def convert_to_view( - mdl: Dict[str, Any], view_converter: ViewConverter + mdl: Dict[str, Any], view_converter: ViewConverter, id: Optional[str] = None ) -> Dict[str, Any]: logger.debug( f"input in convert_to_view: {orjson.dumps(mdl, option=orjson.OPT_INDENT_2).decode()}" ) - return view_converter.run(mdl=mdl) + return view_converter.run(mdl=mdl, id=id) @async_timer @@ -442,7 +456,7 @@ class Indexing(BasicPipeline): AsyncDriver({}, sys.modules[__name__], result_builder=base.DictResult()) ) - def visualize(self, mdl_str: str) -> None: + def visualize(self, mdl_str: str, id: Optional[str] = None) -> None: destination = "outputs/pipelines/indexing" if not Path(destination).exists(): Path(destination).mkdir(parents=True, exist_ok=True) @@ -452,6 +466,7 @@ class Indexing(BasicPipeline): output_file_path=f"{destination}/indexing.dot", inputs={ "mdl_str": mdl_str, + "id": id, "cleaner": self.cleaner, "validator": self.validator, "ddl_converter": self.ddl_converter, @@ -467,12 +482,13 @@ class Indexing(BasicPipeline): @async_timer @observe(name="Ask Indexing") - async def run(self, mdl_str: str) -> Dict[str, Any]: + async def run(self, mdl_str: str, id: Optional[str] = None) -> Dict[str, Any]: logger.info("Ask Indexing pipeline is running...") return await self._pipe.execute( ["write_ddl", "write_view"], inputs={ "mdl_str": mdl_str, + "id": id, "cleaner": self.cleaner, "validator": self.validator, "ddl_converter": self.ddl_converter, diff --git a/wren-ai-service/src/providers/document_store/qdrant.py b/wren-ai-service/src/providers/document_store/qdrant.py index 64bf6ad69..f717262b3 100644 --- a/wren-ai-service/src/providers/document_store/qdrant.py +++ b/wren-ai-service/src/providers/document_store/qdrant.py @@ -14,7 +14,6 @@ from haystack_integrations.document_stores.qdrant import ( ) from haystack_integrations.document_stores.qdrant.converters import ( DENSE_VECTORS_NAME, - convert_haystack_documents_to_qdrant_points, convert_id, convert_qdrant_point_to_haystack_document, ) @@ -30,6 +29,43 @@ from src.providers.loader import get_default_embedding_model_dim, provider logger = logging.getLogger("wren-ai-service") +def convert_haystack_documents_to_qdrant_points( + documents: List[Document], + *, + embedding_field: str, + use_sparse_embeddings: bool, +) -> List[rest.PointStruct]: + DENSE_VECTORS_NAME = "text-dense" + SPARSE_VECTORS_NAME = "text-sparse" + + points = [] + for document in documents: + payload = document.to_dict(flatten=True) + if use_sparse_embeddings: + vector = {} + + dense_vector = payload.pop(embedding_field, None) + if dense_vector is not None: + vector[DENSE_VECTORS_NAME] = dense_vector + + sparse_vector = payload.pop("sparse_embedding", None) + if sparse_vector is not None: + sparse_vector_instance = rest.SparseVector(**sparse_vector) + vector[SPARSE_VECTORS_NAME] = sparse_vector_instance + + else: + vector = payload.pop(embedding_field) or {} + _id = convert_id(payload.get("id")) + + point = rest.PointStruct( + payload=payload, + vector=vector, + id=_id, + ) + points.append(point) + return points + + class AsyncQdrantDocumentStore(QdrantDocumentStore): def __init__( self, @@ -78,7 +114,7 @@ class AsyncQdrantDocumentStore(QdrantDocumentStore): grpc_port=grpc_port, prefer_grpc=prefer_grpc, https=https, - api_key=api_key.resolve_value() if api_key else None, + api_key=api_key, prefix=prefix, timeout=timeout, host=host, @@ -111,7 +147,6 @@ class AsyncQdrantDocumentStore(QdrantDocumentStore): payload_fields_to_index=payload_fields_to_index, ) - metadata = metadata or {} self.async_client = qdrant_client.AsyncQdrantClient( location=location, url=url, @@ -124,7 +159,13 @@ class AsyncQdrantDocumentStore(QdrantDocumentStore): timeout=timeout, host=host, path=path, - metadata=metadata, + metadata=metadata or {}, + ) + + # to improve the indexing performance + # see https://qdrant.tech/documentation/guides/multiple-partitions/?q=mul#calibrate-performance + self.client.create_payload_index( + collection_name=index, field_name="id", field_schema="keyword" ) async def _query_by_embedding( @@ -143,6 +184,17 @@ class AsyncQdrantDocumentStore(QdrantDocumentStore): name=DENSE_VECTORS_NAME if self.use_sparse_embeddings else "", vector=query_embedding, ), + search_params=( + rest.SearchParams( + quantization=rest.QuantizationSearchParams( + rescore=True, + oversampling=3.0, + ), + ) + if len(query_embedding) + >= 1024 # reference: https://qdrant.tech/articles/binary-quantization/#when-should-you-not-use-bq + else None + ), query_filter=qdrant_filters, limit=top_k, with_vectors=return_embedding, @@ -176,8 +228,14 @@ class AsyncQdrantDocumentStore(QdrantDocumentStore): "Called QdrantDocumentStore.delete_documents() on a non-existing ID", ) - async def count_documents(self) -> int: - return (await self.async_client.count(collection_name=self.index)).count + async def count_documents(self, filters: Optional[Dict[str, Any]] = None) -> int: + qdrant_filters = convert_filters_to_qdrant(filters) + + return ( + await self.async_client.count( + collection_name=self.index, count_filter=qdrant_filters + ) + ).count async def write_documents( self, documents: List[Document], policy: DuplicatePolicy = DuplicatePolicy.FAIL @@ -269,8 +327,15 @@ class AsyncQdrantEmbeddingRetriever(QdrantEmbeddingRetriever): @provider("qdrant") class QdrantProvider(DocumentStoreProvider): - def __init__(self, location: str = os.getenv("QDRANT_HOST", "qdrant")): + def __init__( + self, + location: str = os.getenv("QDRANT_HOST", "qdrant"), + api_key: Optional[Secret] = Secret.from_env_var("QDRANT_API_KEY") + if os.getenv("QDRANT_API_KEY") + else None, + ): self._location = location + self._api_key = api_key def get_store( self, @@ -291,9 +356,26 @@ class QdrantProvider(DocumentStoreProvider): return AsyncQdrantDocumentStore( location=self._location, + api_key=self._api_key, embedding_dim=embedding_model_dim, index=dataset_name or "Document", recreate_index=recreate_index, + on_disk=True, + quantization_config=( + rest.BinaryQuantization( + binary=rest.BinaryQuantizationConfig( + always_ram=True, + ) + ) + if embedding_model_dim >= 1024 + else None + ), + # to improve the indexing performance, we disable building global index for the whole collection + # see https://qdrant.tech/documentation/guides/multiple-partitions/?q=mul#calibrate-performance + hnsw_config=rest.HnswConfigDiff( + payload_m=16, + m=0, + ), ) def get_retriever( diff --git a/wren-ai-service/src/web/v1/services/ask.py b/wren-ai-service/src/web/v1/services/ask.py index 2adc8f441..27716bb3a 100644 --- a/wren-ai-service/src/web/v1/services/ask.py +++ b/wren-ai-service/src/web/v1/services/ask.py @@ -130,6 +130,7 @@ class AskService: retrieval_result = await self._pipelines["retrieval"].run( query=ask_request.query, + id=ask_request.project_id, ) documents = retrieval_result.get("retrieval", {}).get("documents", []) diff --git a/wren-ai-service/src/web/v1/services/indexing.py b/wren-ai-service/src/web/v1/services/indexing.py index dc3227bdb..976999ede 100644 --- a/wren-ai-service/src/web/v1/services/indexing.py +++ b/wren-ai-service/src/web/v1/services/indexing.py @@ -57,7 +57,10 @@ class IndexingService: ): try: logger.info(f"MDL: {prepare_semantics_request.mdl}") - await self._pipelines["indexing"].run(prepare_semantics_request.mdl) + await self._pipelines["indexing"].run( + mdl_str=prepare_semantics_request.mdl, + id=prepare_semantics_request.project_id, + ) self._prepare_semantics_statuses[ prepare_semantics_request.mdl_hash diff --git a/wren-ai-service/tests/locust/locust_script.py b/wren-ai-service/tests/locust/locust_script.py index 80b79e0f9..b68afdfaf 100644 --- a/wren-ai-service/tests/locust/locust_script.py +++ b/wren-ai-service/tests/locust/locust_script.py @@ -9,7 +9,7 @@ load_env_vars() filename = f"locust_report_{time.strftime("%Y%m%d_%H%M%S")}" if not Path("./outputs/locust").exists(): - Path("./outputs/locust").mkdir() + Path("./outputs/locust").mkdir(parents=True, exist_ok=True) os.system( f""" diff --git a/wren-ai-service/tools/dev/docker-compose-dev.yaml b/wren-ai-service/tools/dev/docker-compose-dev.yaml index 1cbfea987..965562aee 100644 --- a/wren-ai-service/tools/dev/docker-compose-dev.yaml +++ b/wren-ai-service/tools/dev/docker-compose-dev.yaml @@ -33,7 +33,7 @@ services: - wren qdrant: - image: qdrant/qdrant:v1.7.4 + image: qdrant/qdrant:v1.10.1 pull_policy: always ports: - 6333:6333