Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
40 changes: 39 additions & 1 deletion vectordb_bench/backend/clients/antfly/antfly.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@
import httpx

from ..api import DBCaseConfig, MetricType, VectorDB
from ...filter import Filter, FilterOp
from ...payload import PayloadProfile

log = logging.getLogger(__name__)
Expand Down Expand Up @@ -65,19 +66,28 @@ def _detect_api_root(host: str, port: int) -> str:


class Antfly(VectorDB):
supported_filter_types: list[FilterOp] = [
FilterOp.NonFilter,
FilterOp.NumGE,
FilterOp.StrEqual,
]

def __init__(
self,
dim: int,
db_config: dict,
db_case_config: DBCaseConfig,
collection_name: str = "vdbbench",
drop_old: bool = False,
with_scalar_labels: bool = False,
**kwargs,
):
self.db_config = db_config
self.case_config = db_case_config
self.collection_name = collection_name
self.dim = dim
self.with_scalar_labels = with_scalar_labels
self._filter_query: dict[str, Any] | None = None

# Antfly v0.1 used /api/v1; current antfly-zig serves the public DB API
# at /db/v1. Auto-detect unless ANTFLY_API_ROOT pins it explicitly.
Expand Down Expand Up @@ -462,12 +472,15 @@ def _serialize_insert_vector(self, vector: list[float]) -> str | list[float]:
return self._pack_vector(vector)

def _metadata_query_body(self, query: list[float], k: int) -> dict[str, Any]:
return {
body = {
"embeddings": {"vec": self._serialize_query_vector(query)},
"limit": k,
"fields": [],
**self.case_config.search_param(),
}
if getattr(self, "_filter_query", None) is not None:
body["filter_query"] = self._filter_query
return body

def _store_query_body(self, query: list[float], k: int) -> dict[str, Any]:
search_params = self.case_config.search_param()
Expand Down Expand Up @@ -533,6 +546,24 @@ def optimize(self, data_size: int | None = None):
self._wait_for_index_ready(client, expected_total=data_size)
self._maybe_log_bench_status(client, "optimize_end", force=True)

def prepare_filter(self, filters: Filter):
if filters.type == FilterOp.NonFilter:
self._filter_query = None
elif filters.type == FilterOp.NumGE:
self._filter_query = {
"numeric_range": {
"field": filters.int_field,
"min": filters.int_value,
"inclusive_min": True,
}
}
elif filters.type == FilterOp.StrEqual:
self._filter_query = {
"term": {filters.label_field: filters.label_value}
}
else:
raise ValueError(f"Unsupported Antfly filter: {filters}")

def _catch_up_lag_sequences(self) -> int | None:
try:
payload = self._get_index_status(self.client)
Expand Down Expand Up @@ -564,6 +595,7 @@ def insert_embeddings(
self,
embeddings: list[list[float]],
metadata: list[int],
labels_data: list[str] | None = None,
**kwargs: Any,
) -> tuple[int, Exception]:
total = len(embeddings)
Expand All @@ -584,6 +616,10 @@ def insert_embeddings(
SOURCE_FIELD: str(metadata[i]),
"_embeddings": {"vec": serialized_embedding},
}
if self.with_scalar_labels:
if labels_data is None:
raise ValueError("Antfly label-filter load requires labels_data")
inserts[key]["labels"] = labels_data[i]
payload = {"inserts": inserts, "sync_level": self._write_sync_level}
r = self.client.post(
f"/tables/{self.collection_name}/batch", json=payload
Expand Down Expand Up @@ -613,6 +649,8 @@ def search_embedding(
query = self._normalize_vector(query)

if self._use_direct_store_search:
if self._filter_query is not None:
raise ValueError("Antfly filtered ANN requires the public metadata query API")
if self._direct_shard_id is None:
self._refresh_direct_search_routing(self.client)
r = self.store_client.post(
Expand Down
5 changes: 4 additions & 1 deletion vectordb_bench/backend/clients/antfly/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
CommonTypedDict,
cli,
click_parameter_decorators_from_typed_dict,
get_custom_case_config,
run,
)
from ...cases import CaseType
Expand Down Expand Up @@ -81,7 +82,9 @@ def AntflyAKNN(**parameters: Unpack[AntflyAKNNTypedDict]):
metric_type = (
MetricType(parameters["metric_type"])
if parameters["metric_type"]
else CaseType[parameters["case_type"]].case_cls(parameters.get("custom_case")).dataset.data.metric_type
else CaseType[parameters["case_type"]]
.case_cls(get_custom_case_config(parameters))
.dataset.data.metric_type
)

run(
Expand Down
20 changes: 18 additions & 2 deletions vectordb_bench/backend/clients/chroma/chroma.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,8 @@

import chromadb

from vectordb_bench.backend.filter import Filter, FilterOp

from ..api import VectorDB
from .config import ChromaIndexConfig

Expand All @@ -17,6 +19,11 @@ class ChromaClient(VectorDB):
To change to running in process, modify the HttpClient() in __init__() and init().
"""

supported_filter_types: list[FilterOp] = [
FilterOp.NonFilter,
FilterOp.NumGE,
]

def __init__(
self,
dim: int,
Expand All @@ -42,6 +49,7 @@ def __init__(

self.client = None
self.collection = None
self._where_filter: dict | None = None

@contextmanager
def init(self):
Expand Down Expand Up @@ -85,10 +93,18 @@ def search_embedding(
self, query: list[float], k: int = 100, filters: dict | None = None, timeout: int | None = None
) -> list[int]:
assert self.client is not None, "Please call self.init() before"
if filters:
if self._where_filter is not None:
results = self.collection.query(
query_embeddings=[query], n_results=k, where={"id": {"$gt": filters.get("id")}}
query_embeddings=[query], n_results=k, where=self._where_filter
)
else:
results = self.collection.query(query_embeddings=[query], n_results=k)
return [int(idx) for idx in results["ids"][0]]

def prepare_filter(self, filters: Filter):
if filters.type == FilterOp.NonFilter:
self._where_filter = None
elif filters.type == FilterOp.NumGE:
self._where_filter = {"index": {"$gte": filters.int_value}}
else:
raise ValueError(f"Unsupported Chroma filter: {filters}")
1 change: 1 addition & 0 deletions vectordb_bench/backend/clients/chroma/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,7 @@ def Chroma(**parameters: Unpack[ChromaTypeDict]):
run(
db=DBTYPE,
db_config=ChromaConfig(
db_label=parameters["db_label"],
user=parameters["user"],
password=SecretStr(parameters["password"]) if parameters["password"] else None,
host=SecretStr(parameters["host"]),
Expand Down
4 changes: 2 additions & 2 deletions vectordb_bench/backend/clients/elastic_cloud/elastic_cloud.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,7 +50,7 @@ def __init__(

from elasticsearch import Elasticsearch

client = Elasticsearch(**self.db_config)
client = Elasticsearch(**self.db_config, request_timeout=180)

if drop_old:
log.info(f"Elasticsearch client drop_old indices: {self.indice}")
Expand Down Expand Up @@ -154,7 +154,7 @@ def prepare_filter(self, filters: Filter):
if filters.type == FilterOp.NonFilter:
self.filter = []
elif filters.type == FilterOp.NumGE:
self.filter = {"range": {self.id_col_name: {"gt": filters.int_value}}}
self.filter = {"range": {self.id_col_name: {"gte": filters.int_value}}}
elif filters.type == FilterOp.StrEqual:
self.filter = {"term": {self.label_col_name: filters.label_value}}
if self.case_config.use_routing:
Expand Down
5 changes: 4 additions & 1 deletion vectordb_bench/backend/clients/qdrant_local/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,7 +49,10 @@ def QdrantLocal(**parameters: Unpack[QdrantLocalTypedDict]):

run(
db=DBTYPE,
db_config=QdrantLocalConfig(url=SecretStr(parameters["url"])),
db_config=QdrantLocalConfig(
db_label=parameters["db_label"],
url=SecretStr(parameters["url"]),
),
db_case_config=QdrantLocalIndexConfig(
on_disk=parameters["on_disk"],
m=parameters["m"],
Expand Down
37 changes: 24 additions & 13 deletions vectordb_bench/backend/clients/qdrant_local/qdrant_local.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,8 @@
VectorParams,
)

from vectordb_bench.backend.filter import Filter as BenchFilter, FilterOp

from ..api import VectorDB
from .config import QdrantLocalIndexConfig

Expand All @@ -40,6 +42,11 @@ def qdrant_collection_exists(client: QdrantClient, collection_name: str) -> bool


class QdrantLocal(VectorDB):
supported_filter_types: list[FilterOp] = [
FilterOp.NonFilter,
FilterOp.NumGE,
]

def __init__(
self,
dim: int,
Expand All @@ -60,6 +67,7 @@ def __init__(

self._primary_field = "pk"
self._vector_field = "vector"
self._query_filter: Filter | None = None

client = QdrantClient(**self.db_config)

Expand Down Expand Up @@ -209,25 +217,28 @@ def search_embedding(
"""
assert self.client is not None

f = None
if filters:
f = Filter(
must=[
FieldCondition(
key=self._primary_field,
range=Range(
gt=filters.get("id"),
),
),
],
)
res = self.client.query_points(
collection_name=self.collection_name,
query=query,
limit=k,
query_filter=f,
query_filter=self._query_filter,
search_params=SearchParams(**self.search_parameter),
timeout=timeout,
).points

return [result.id for result in res]

def prepare_filter(self, filters: BenchFilter):
if filters.type == FilterOp.NonFilter:
self._query_filter = None
elif filters.type == FilterOp.NumGE:
self._query_filter = Filter(
must=[
FieldCondition(
key=self._primary_field,
range=Range(gte=filters.int_value),
)
]
)
else:
raise ValueError(f"Unsupported QdrantLocal filter: {filters}")
28 changes: 21 additions & 7 deletions vectordb_bench/backend/clients/weaviate_cloud/weaviate_cloud.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,13 +7,19 @@
import weaviate
from weaviate.exceptions import WeaviateBaseError

from vectordb_bench.backend.filter import Filter, FilterOp

from ..api import DBCaseConfig, VectorDB

log = logging.getLogger(__name__)


class WeaviateCloud(VectorDB):
thread_safe: bool = False
supported_filter_types: list[FilterOp] = [
FilterOp.NonFilter,
FilterOp.NumGE,
]

def __init__(
self,
Expand All @@ -40,6 +46,7 @@ def __init__(
self._scalar_field = "key"
self._vector_field = "vector"
self._index_name = "vector_idx"
self._where_filter: dict | None = None

# If local setup is used, we
if db_config["no_auth"]:
Expand Down Expand Up @@ -148,16 +155,23 @@ def search_embedding(
.with_near_vector({"vector": query})
.with_limit(k)
)
if filters:
where_filter = {
"path": "key",
"operator": "GreaterThanEqual",
"valueInt": filters.get("id"),
}
query_obj = query_obj.with_where(where_filter)
if self._where_filter is not None:
query_obj = query_obj.with_where(self._where_filter)

# Perform the search.
res = query_obj.do()

# Organize results.
return [result[self._scalar_field] for result in res["data"]["Get"][self.collection_name]]

def prepare_filter(self, filters: Filter):
if filters.type == FilterOp.NonFilter:
self._where_filter = None
elif filters.type == FilterOp.NumGE:
self._where_filter = {
"path": [self._scalar_field],
"operator": "GreaterThanEqual",
"valueInt": filters.int_value,
}
else:
raise ValueError(f"Unsupported Weaviate filter: {filters}")
7 changes: 7 additions & 0 deletions vectordb_bench/backend/dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
"""

import logging
import os
import pathlib
from enum import Enum
from typing import Any, ClassVar, NamedTuple
Expand Down Expand Up @@ -374,6 +375,12 @@ def prepare(
if self.data.with_scalar_labels and self.data.scalar_labels_file_separated:
download_files.append(self.data.scalar_labels_file)
download_files = [file for file in download_files if file is not None]
if (
os.environ.get("VDBB_USE_LOCAL_FILTER_GT") == "1"
and gt_file is not None
and self.data_dir.joinpath(gt_file).exists()
):
download_files = [file for file in download_files if file != gt_file]
source.reader().read(
dataset=self.data.dir_name.lower(),
files=download_files,
Expand Down
Loading