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
This commit is contained in:
Chih-Yu Yeh
2024-08-05 13:43:01 +08:00
committed by GitHub
parent 882d71c37d
commit af556b0261
13 changed files with 154 additions and 36 deletions
@@ -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: []
+2 -2
View File
@@ -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
+1 -1
View File
@@ -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
+1 -1
View File
@@ -72,7 +72,7 @@ services:
- qdrant
qdrant:
image: qdrant/qdrant:v1.7.4
image: qdrant/qdrant:v1.10.1
restart: on-failure
expose:
- 6333
+1
View File
@@ -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
-1
View File
@@ -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
+20 -4
View File
@@ -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,
},
@@ -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,
@@ -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(
@@ -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", [])
@@ -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
@@ -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"""
@@ -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