mirror of
https://github.com/Canner/WrenAI.git
synced 2026-09-24 23:29:49 +08:00
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:
+1
-1
@@ -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: []
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -72,7 +72,7 @@ services:
|
||||
- qdrant
|
||||
|
||||
qdrant:
|
||||
image: qdrant/qdrant:v1.7.4
|
||||
image: qdrant/qdrant:v1.10.1
|
||||
restart: on-failure
|
||||
expose:
|
||||
- 6333
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user