refactor: change rerank interface from map-based to vector-based (#458)

* refactor: change rerank interface from map-based to vector-based (#452)

- Define QueryResult = list[Doc] type alias in doc.py
- Change C++ Reranker::rerank() signature from map<string, DocPtrList> to vector<DocPtrList>
- Extend bind_schema() to accept field_names for index-based field lookup
- Update ScoreBasedReranker/WeightedReranker/CallbackReranker implementations
- Adapt collection.cc MultiQuery path to use vector<DocPtrList>
- Update Python binding to expose rerank() and use vector<double> weights
- Refactor Python RerankFunction interface to list[QueryResult] -> QueryResult
- Remove Python-layer rerank logic from RrfReRanker/WeightedReRanker (delegate to C++)
- Update query_executor to return list[list[Doc]] instead of dict
- Update all related unit tests (C++ and Python)

* refactor: replace list[Doc] with QueryResult type alias in executor and rerank functions

* refactor: replace list[list[Doc]] with list[QueryResult] in query_executor

* fix: remove unused Doc import in rerank_function.py (ruff F401)

* refactor(query_executor): merge duplicate rerank return paths

* refactor: RrfReRanker/WeightedReRanker.rerank() directly call C++ reranker

* refactor: simplify QueryExecutor into unified class, remove Factory/subclasses/validation/concurrency

* refactor: rename _VectorQuery to _SearchQuery, from_vector_query to from_search_query

* refactor(query_executor): split execute into single/multi paths, rename core_vector to search_query, drop unused core_vectors

* style: apply ruff formatter to test_reranker.py and query_executor.py

* refactor: make rescore() private in ScoreBasedReranker hierarchy

* style: apply clang-format to reranker.h

* style: apply clang-format to all modified C++ files

* refactor: rename private methods in QueryExecutor for clearer semantics

* refactor: rename mvq to multi_query for clarity

* fix: make BasicRRF test order-independent for equal scores

* fix: update collection_test to use vector-based reranker interface

* fix: update reranker tests to expect TypeError instead of NotImplementedError

* refactor: remove PendingQuery wrapper, use SearchQuery directly in MultiQuery path

* refactor: simplify MultiQuery path - remove seen_fields, merge field_names into main loop

* fix: address review comments - defensive checks and remove fields param from C API

- ScoreBasedReranker::rerank(): early return empty list when topn <= 0
- WeightedReranker::rescore(): null-check schema_ before use
- CallbackReranker::rerank(): check callback_ is not empty before invoke
- C API zvec_reranker_create_weighted(): remove unused fields parameter

* fix: remove duplicate field name test (check was intentionally removed)

* fix: address egolearner review comments

- Rename QueryResult to DocList for clarity (见名知义)
- Change docstring to #: comment for type alias
- Fix output_fields check: use 'is not None' instead of truthy check
  (None means unset, [] means explicit empty list - different semantics)
- Raise ValueError when search-by-id finds no document

* refactor: remove redundant output_fields assignment in _build_search_query

* refactor: address egolearner review comments (C++ refactoring)

- c_api.cc: simplify weighted reranker creation with inline vector ctor
- python_reranker.cc: refactor unwrap_rerank_result - take by value,
  early error return, move semantics
- Rename C API functions for consistent naming:
  zvec_reranker_create_rrf -> zvec_create_rrf_reranker
  zvec_reranker_create_weighted -> zvec_create_weighted_reranker
  zvec_reranker_destroy -> zvec_destroy_reranker
  zvec_reranker_get_rank_constant -> zvec_get_reranker_rank_constant
- reranker.h/cc: bind_schema returns Result<void>, caches
  vector<const FieldSchema*> to avoid repeated schema lookups in rescore
- python_param.cc: rename py::arg vector_query to search_query

* revert: rollback bind_schema refactoring due to thread-safety concern

The field_schemas_ caching approach introduces a data race when the same
WeightedReranker instance is shared across concurrent queries: bind_schema()
writes field_schemas_ while rerank() reads it concurrently.

Revert to storing schema_ + field_names_ and looking up fields in rescore().
Add @note thread-safety warning to WeightedReranker class documentation.

* fix: unify error message format in collection.cc

Change 'Vector field not found: X' to 'Invalid query: field X not found'
for consistent error formatting as suggested by zhourrr.

* fix: sort __all__ and remove duplicates in __init__.pyi

Fix RUF022 lint error: sort __all__ alphabetically and remove duplicate
entries (DenseEmbeddingFunction, ReRanker).

* style: format query_executor.py with ruff formatter

* fix: resolve Python test failures after FTS rebase integration

- test_query_executor.py: update method names to match refactored API
  (_do_build -> _build_queries, _do_merge_rerank_results -> _merge_and_rerank)
- test_reranker.py: fix expected exception type (TypeError from pybind11)
- test_collection_fts.py: update error message match patterns
- test_collection_fts_vector_hybrid.py: remove obsolete 'metrics' param,
  update weights from dict to positional list, adapt validation tests
  for multi-vector queries (now supported with reranker)
- test_collection_dql.py: remove 'metrics' param, update weights format
- collection.cc: distinguish FTS vs vector fields in MultiQuery path
  using get_fts_clause() to route field lookup correctly
- reranker.cc: use get_field() instead of get_vector_field() in rescore
  to support FTS+vector hybrid weighted reranking

* refactor: pass topn as rerank() parameter, move rerank_field to model rerankers

* fix: address review comments - rename test functions and restore duplicate field check

* refactor: simplify MultiQuery field lookup, let validate_and_sanitize handle type check
This commit is contained in:
Cuiys 2026-06-04 16:03:28 +08:00 committed by GitHub
parent f562bdd636
commit c46efe1241
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
29 changed files with 704 additions and 1098 deletions

View File

@ -824,7 +824,7 @@ class TestCollectionQuery:
for k, v in DEFAULT_VECTOR_FIELD_NAME.items():
multi_query_vectors.append(Query(field_name=v, vector=doc_vectors[v]))
rrf_reranker = RrfReRanker(topn=3)
rrf_reranker = RrfReRanker()
multi_query_result = full_collection.query(
multi_query_vectors,
reranker=rrf_reranker,
@ -876,8 +876,8 @@ class TestCollectionQuery:
batchdoc_and_check(full_collection, multiple_docs, doc_num, operator="insert")
doc_fields, doc_vectors = generate_vectordict_random(full_collection.schema)
metrics = {field: MetricType.IP for field in weights}
weighted_reranker = WeightedReRanker(topn=3, weights=weights, metrics=metrics)
weight_list = [weights[v] for v in DEFAULT_VECTOR_FIELD_NAME.values()]
weighted_reranker = WeightedReRanker(weights=weight_list)
single_query_results = {}
for k, v in DEFAULT_VECTOR_FIELD_NAME.items():
@ -1165,13 +1165,6 @@ class TestCollectionQuery:
), # Non-existent ID
"Expected exception for non-existent document ID",
),
(
"Both vector and id specified (invalid combination)",
lambda ref_dense_vector: Query(
field_name="vector_fp32_field", vector=ref_dense_vector, id="5"
),
"Expected exception for specifying both vector and id",
),
(
"Neither vector nor id specified",
lambda ref_dense_vector: Query(

View File

@ -976,18 +976,6 @@ class TestCollectionQuery:
result = collection_with_multiple_docs.query(filter="id in (1)", topk=100)
assert len(result) == 1
def test_collection_query_with_vector_and_id(
self, collection_with_single_doc: Collection, single_doc: Doc
):
with pytest.raises(ValueError):
collection_with_single_doc.query(
Query(
field_name="dense",
id=single_doc.id,
vector=single_doc.vector("dense"),
)
)
def test_collection_query_with_filter_not_in(
self, collection_with_multiple_docs: Collection, multiple_docs
):
@ -1013,30 +1001,6 @@ class TestCollectionQuery:
)
assert len(result) == 10
def test_collection_query_multi_vector_with_same_field(
self, collection_with_multiple_docs: Collection, multiple_docs
):
# Multi-vector query on same field without reranker should raise ValueError
with pytest.raises(ValueError, match="Reranker is required"):
collection_with_multiple_docs.query(
[
Query(field_name="dense", vector=multiple_docs[0].vector("dense")),
Query(field_name="dense", vector=multiple_docs[1].vector("dense")),
]
)
# Same field name with reranker should also raise ValueError
reranker = RrfReRanker(topn=10, rank_constant=60)
with pytest.raises(ValueError, match="appears more than once"):
collection_with_multiple_docs.query(
[
Query(field_name="dense", vector=multiple_docs[0].vector("dense")),
Query(field_name="dense", vector=multiple_docs[1].vector("dense")),
],
topk=10,
reranker=reranker,
)
def test_collection_query_by_dense_vector(
self, collection_with_multiple_docs: Collection, multiple_docs
):
@ -1087,7 +1051,7 @@ class TestCollectionQuery:
self, collection_with_multiple_docs: Collection, multiple_docs
):
"""Test multi-vector query with RRF reranker on multiple dense vectors."""
reranker = RrfReRanker(topn=10, rank_constant=60)
reranker = RrfReRanker(rank_constant=60)
result = collection_with_multiple_docs.query(
[
Query(field_name="dense", vector=multiple_docs[0].vector("dense")),
@ -1106,7 +1070,7 @@ class TestCollectionQuery:
self, collection_with_multiple_docs: Collection, multiple_docs
):
"""Test multi-vector query with RRF reranker on multiple sparse vectors."""
reranker = RrfReRanker(topn=10, rank_constant=60)
reranker = RrfReRanker(rank_constant=60)
result = collection_with_multiple_docs.query(
[
Query(field_name="sparse", vector=multiple_docs[0].vector("sparse")),
@ -1125,7 +1089,7 @@ class TestCollectionQuery:
self, collection_with_multiple_docs: Collection, multiple_docs
):
"""Test multi-vector query with RRF reranker combining dense + sparse."""
reranker = RrfReRanker(topn=10, rank_constant=60)
reranker = RrfReRanker(rank_constant=60)
result = collection_with_multiple_docs.query(
[
Query(field_name="dense", vector=multiple_docs[0].vector("dense")),
@ -1141,9 +1105,7 @@ class TestCollectionQuery:
self, collection_with_multiple_docs: Collection, multiple_docs
):
"""Test multi-vector query with Weighted reranker on multiple dense vectors."""
metrics = {"dense": MetricType.IP, "dense2": MetricType.IP}
weights = {"dense": 0.6, "dense2": 0.4}
reranker = WeightedReRanker(topn=10, metrics=metrics, weights=weights)
reranker = WeightedReRanker(weights=[0.6, 0.4])
result = collection_with_multiple_docs.query(
[
Query(field_name="dense", vector=multiple_docs[0].vector("dense")),
@ -1159,9 +1121,7 @@ class TestCollectionQuery:
self, collection_with_multiple_docs: Collection, multiple_docs
):
"""Test multi-vector query with Weighted reranker on multiple sparse vectors."""
metrics = {"sparse": MetricType.IP, "sparse2": MetricType.IP}
weights = {"sparse": 0.6, "sparse2": 0.4}
reranker = WeightedReRanker(topn=10, metrics=metrics, weights=weights)
reranker = WeightedReRanker(weights=[0.6, 0.4])
result = collection_with_multiple_docs.query(
[
Query(field_name="sparse", vector=multiple_docs[0].vector("sparse")),
@ -1180,9 +1140,7 @@ class TestCollectionQuery:
self, collection_with_multiple_docs: Collection, multiple_docs
):
"""Test multi-vector query with Weighted reranker combining dense + sparse."""
metrics = {"dense": MetricType.IP, "sparse": MetricType.IP}
weights = {"dense": 0.7, "sparse": 0.3}
reranker = WeightedReRanker(topn=10, metrics=metrics, weights=weights)
reranker = WeightedReRanker(weights=[0.7, 0.3])
result = collection_with_multiple_docs.query(
[
Query(field_name="dense", vector=multiple_docs[0].vector("dense")),
@ -1203,7 +1161,7 @@ class TestCollectionQuery:
def my_rerank_callback(query_results, topn):
callback_invoked.append(True)
all_docs = []
for docs in query_results.values():
for docs in query_results:
all_docs.extend(docs)
seen = set()
unique_docs = []
@ -1214,7 +1172,7 @@ class TestCollectionQuery:
unique_docs.sort(key=lambda d: d.score(), reverse=True)
return unique_docs[:topn]
reranker = CallbackReRanker(callback=my_rerank_callback, topn=10)
reranker = CallbackReRanker(callback=my_rerank_callback)
result = collection_with_multiple_docs.query(
[
Query(field_name="dense", vector=multiple_docs[0].vector("dense")),
@ -1234,7 +1192,7 @@ class TestCollectionQuery:
def my_rerank_callback(query_results, topn):
all_docs = []
for docs in query_results.values():
for docs in query_results:
all_docs.extend(docs)
seen = set()
unique_docs = []
@ -1245,7 +1203,7 @@ class TestCollectionQuery:
unique_docs.sort(key=lambda d: d.score(), reverse=True)
return unique_docs[:topn]
reranker = CallbackReRanker(callback=my_rerank_callback, topn=5)
reranker = CallbackReRanker(callback=my_rerank_callback)
result = collection_with_multiple_docs.query(
[
Query(field_name="dense", vector=multiple_docs[0].vector("dense")),

View File

@ -172,7 +172,7 @@ class TestFtsOnlyCollectionLifecycle:
class TestFtsOnlyCollectionQueryValidation:
def test_vector_query_rejected(self, fts_collection: Collection):
"""Vector query on a no-vector collection must raise."""
with pytest.raises(ValueError, match="vector or id"):
with pytest.raises(ValueError, match="No vector field found"):
fts_collection.query(
queries=Query(field_name="content", vector=[0.1, 0.2, 0.3]),
topk=5,
@ -181,7 +181,7 @@ class TestFtsOnlyCollectionQueryValidation:
def test_id_query_rejected(self, fts_collection: Collection):
"""ID-based query on a no-vector collection must raise."""
fts_collection.insert(_make_docs()[:1])
with pytest.raises(ValueError, match="vector or id"):
with pytest.raises(ValueError, match="No vector field found"):
fts_collection.query(
queries=Query(field_name="content", id="pk_0"),
topk=5,

View File

@ -29,7 +29,6 @@ from zvec import (
)
from zvec.extension.multi_vector_reranker import RrfReRanker, WeightedReRanker
from zvec.model.param.query import Fts, Query
from zvec.typing import MetricType
DIM = 16
@ -166,7 +165,7 @@ class TestFtsVectorHybridQuery:
def test_hybrid_fts_and_vector_basic(self, hybrid_collection_with_docs: Collection):
"""FTS + vector multi-query with RRF reranker returns results."""
reranker = RrfReRanker(topn=10, rank_constant=60)
reranker = RrfReRanker(rank_constant=60)
result = hybrid_collection_with_docs.query(
queries=[
Query(field_name="content", fts=Fts(match_string="retrieval")),
@ -185,7 +184,7 @@ class TestFtsVectorHybridQuery:
self, hybrid_collection_with_docs: Collection
):
"""Docs relevant in both FTS and vector should rank higher."""
reranker = RrfReRanker(topn=10, rank_constant=60)
reranker = RrfReRanker(rank_constant=60)
# FTS: "retrieval search" matches pk_3, pk_4
# Vector: ret_vec cluster matches pk_3, pk_4
# Both signals agree: pk_3 and pk_4 should rank top
@ -202,7 +201,7 @@ class TestFtsVectorHybridQuery:
def test_hybrid_scores_descending(self, hybrid_collection_with_docs: Collection):
"""Hybrid query results must be sorted by score descending."""
reranker = RrfReRanker(topn=10, rank_constant=60)
reranker = RrfReRanker(rank_constant=60)
result = hybrid_collection_with_docs.query(
queries=[
Query(field_name="content", fts=Fts(match_string="intelligence")),
@ -217,7 +216,7 @@ class TestFtsVectorHybridQuery:
def test_hybrid_with_filter(self, hybrid_collection_with_docs: Collection):
"""Hybrid query respects SQL filter."""
reranker = RrfReRanker(topn=10, rank_constant=60)
reranker = RrfReRanker(rank_constant=60)
result = hybrid_collection_with_docs.query(
queries=[
Query(field_name="content", fts=Fts(match_string="learning")),
@ -234,7 +233,7 @@ class TestFtsVectorHybridQuery:
self, hybrid_collection_with_docs: Collection
):
"""When FTS matches nothing, vector results still appear."""
reranker = RrfReRanker(topn=10, rank_constant=60)
reranker = RrfReRanker(rank_constant=60)
result = hybrid_collection_with_docs.query(
queries=[
Query(
@ -251,7 +250,7 @@ class TestFtsVectorHybridQuery:
def test_hybrid_query_string_syntax(self, hybrid_collection_with_docs: Collection):
"""Hybrid query works with FTS query_string (advanced syntax)."""
reranker = RrfReRanker(topn=10, rank_constant=60)
reranker = RrfReRanker(rank_constant=60)
result = hybrid_collection_with_docs.query(
queries=[
Query(
@ -283,35 +282,35 @@ class TestFtsVectorHybridValidation:
topk=5,
)
def test_duplicate_field_name_rejected(
def test_duplicate_field_name_allowed(
self, hybrid_collection_with_docs: Collection
):
"""Multi-query with duplicate field names should raise."""
reranker = RrfReRanker(topn=10, rank_constant=60)
with pytest.raises(ValueError, match="appears more than once"):
hybrid_collection_with_docs.query(
queries=[
Query(field_name="content", fts=Fts(match_string="hello")),
Query(field_name="content", fts=Fts(match_string="world")),
],
topk=5,
reranker=reranker,
)
"""Multi-query with duplicate field names is allowed and returns results."""
reranker = RrfReRanker(rank_constant=60)
result = hybrid_collection_with_docs.query(
queries=[
Query(field_name="content", fts=Fts(match_string="learning")),
Query(field_name="content", fts=Fts(match_string="intelligence")),
],
topk=5,
reranker=reranker,
)
assert len(result) > 0
assert len(result) <= 5
def test_multiple_vectors_without_fts_rejected(
self, hybrid_collection_with_docs: Collection
):
"""Two vector queries on a single-vector-field collection should raise."""
reranker = RrfReRanker(topn=10, rank_constant=60)
with pytest.raises(ValueError, match="cannot query with multiple vectors"):
hybrid_collection_with_docs.query(
queries=[
Query(field_name="embedding", vector=[1.0] * DIM),
Query(field_name="embedding", vector=[0.5] * DIM),
],
topk=5,
reranker=reranker,
)
def test_multiple_vectors_allowed(self, hybrid_collection_with_docs: Collection):
"""Two vector queries on the same field are allowed with a reranker."""
reranker = RrfReRanker(rank_constant=60)
result = hybrid_collection_with_docs.query(
queries=[
Query(field_name="embedding", vector=[1.0] * DIM),
Query(field_name="embedding", vector=[0.5] * DIM),
],
topk=5,
reranker=reranker,
)
assert len(result) > 0
assert len(result) <= 5
class TestFtsVectorHybridWeightedReranker:
@ -321,9 +320,8 @@ class TestFtsVectorHybridWeightedReranker:
self, hybrid_collection_with_docs: Collection
):
"""WeightedReranker correctly normalizes FTS scores alongside vector scores."""
metrics = {"embedding": MetricType.IP}
weights = {"content": 0.5, "embedding": 0.5}
reranker = WeightedReRanker(topn=10, metrics=metrics, weights=weights)
weights = [0.5, 0.5]
reranker = WeightedReRanker(weights=weights)
result = hybrid_collection_with_docs.query(
queries=[
Query(field_name="content", fts=Fts(match_string="retrieval search")),
@ -341,9 +339,8 @@ class TestFtsVectorHybridWeightedReranker:
self, hybrid_collection_with_docs: Collection
):
"""WeightedReranker hybrid results are sorted by score descending."""
metrics = {"embedding": MetricType.IP}
weights = {"content": 0.4, "embedding": 0.6}
reranker = WeightedReRanker(topn=10, metrics=metrics, weights=weights)
weights = [0.4, 0.6]
reranker = WeightedReRanker(weights=weights)
result = hybrid_collection_with_docs.query(
queries=[
Query(field_name="content", fts=Fts(match_string="intelligence")),
@ -361,11 +358,8 @@ class TestFtsVectorHybridWeightedReranker:
):
"""Higher FTS weight should boost FTS-relevant docs in ranking."""
# High FTS weight: FTS signal dominates
metrics = {"embedding": MetricType.IP}
weights_fts_heavy = {"content": 0.9, "embedding": 0.1}
reranker_fts = WeightedReRanker(
topn=10, metrics=metrics, weights=weights_fts_heavy
)
weights_fts_heavy = [0.9, 0.1]
reranker_fts = WeightedReRanker(weights=weights_fts_heavy)
result_fts = hybrid_collection_with_docs.query(
queries=[
Query(field_name="content", fts=Fts(match_string="retrieval")),
@ -376,10 +370,8 @@ class TestFtsVectorHybridWeightedReranker:
)
# High vector weight: vector signal dominates
weights_vec_heavy = {"content": 0.1, "embedding": 0.9}
reranker_vec = WeightedReRanker(
topn=10, metrics=metrics, weights=weights_vec_heavy
)
weights_vec_heavy = [0.1, 0.9]
reranker_vec = WeightedReRanker(weights=weights_vec_heavy)
result_vec = hybrid_collection_with_docs.query(
queries=[
Query(field_name="content", fts=Fts(match_string="retrieval")),

View File

@ -110,11 +110,11 @@ class TestFtsQueryBinding:
assert restored.query_string == "+vector search"
assert restored.match_string == ""
def test_vector_query_fts_field(self):
"""_VectorQuery should have fts field."""
from _zvec.param import _Fts, _VectorQuery
def test_search_query_fts_field(self):
"""_SearchQuery should have fts field."""
from _zvec.param import _Fts, _SearchQuery
vq = _VectorQuery()
vq = _SearchQuery()
# fts should be None by default (optional)
assert vq.fts is None
@ -125,11 +125,11 @@ class TestFtsQueryBinding:
assert vq.fts is not None
assert vq.fts.query_string == "hello"
def test_vector_query_pickle_with_fts(self):
"""_VectorQuery with fts should survive pickling."""
from _zvec.param import _Fts, _VectorQuery
def test_search_query_pickle_with_fts(self):
"""_SearchQuery with fts should survive pickling."""
from _zvec.param import _Fts, _SearchQuery
vq = _VectorQuery()
vq = _SearchQuery()
vq.topk = 10
vq.field_name = "embedding"
fts = _Fts()
@ -143,11 +143,11 @@ class TestFtsQueryBinding:
assert restored.fts is not None
assert restored.fts.match_string == "test query"
def test_vector_query_pickle_without_fts(self):
"""_VectorQuery without fts should survive pickling."""
from _zvec.param import _VectorQuery
def test_search_query_pickle_without_fts(self):
"""_SearchQuery without fts should survive pickling."""
from _zvec.param import _SearchQuery
vq = _VectorQuery()
vq = _SearchQuery()
vq.topk = 5
vq.field_name = "vec"

View File

@ -40,7 +40,7 @@ from zvec import (
VectorSchema,
)
from _zvec.param import _VectorQuery
from _zvec.param import _SearchQuery
# ----------------------------
# Invert Index Param Test Case

View File

@ -14,20 +14,16 @@
from __future__ import annotations
from typing import Dict, Union
from unittest.mock import MagicMock
from unittest.mock import MagicMock, patch
import numpy as np
import math
from _zvec.param import _VectorQuery
from _zvec.param import _SearchQuery
import pytest
from zvec.executor.query_executor import (
MultiVectorQueryExecutor,
NoVectorQueryExecutor,
QueryContext,
QueryExecutor,
QueryExecutorFactory,
SingleVectorQueryExecutor,
)
from zvec import (
RrfReRanker,
@ -43,21 +39,6 @@ from zvec import (
from zvec.extension.multi_vector_reranker import CallbackReRanker
# ----------------------------
# Mock Vector Schema
# ----------------------------
class MockVectorSchema(VectorSchema):
def __init__(self, name="test_vector"):
self._name = name
@property
def name(self):
return self._name
def _get_object(self):
return MagicMock()
# ----------------------------
# Mock Collection Schema
# ----------------------------
@ -110,7 +91,7 @@ class TestQuery:
assert query.has_vector()
def test_validate_dense_fp16_convert(self):
v = _VectorQuery()
v = _SearchQuery()
schema = VectorSchema(name="test", data_type=DataType.VECTOR_FP16)
vec = np.array([1.1, 2.1, 3.1], dtype=np.float16)
v.set_vector(schema._get_object(), vec)
@ -118,7 +99,7 @@ class TestQuery:
assert np.array_equal(vec, ret)
def test_validate_dense_fp32_convert(self):
v = _VectorQuery()
v = _SearchQuery()
schema = VectorSchema(name="test", data_type=DataType.VECTOR_FP32)
vec = np.array([1.1, 2.1, 3.1], dtype=np.float32)
v.set_vector(schema._get_object(), vec)
@ -126,7 +107,7 @@ class TestQuery:
assert np.array_equal(vec, ret)
def test_validate_dense_fp64_convert(self):
v = _VectorQuery()
v = _SearchQuery()
schema = VectorSchema(name="test", data_type=DataType.VECTOR_FP64)
vec = np.array([1.1, 2.1, 3.1], dtype=np.float64)
v.set_vector(schema._get_object(), vec)
@ -134,7 +115,7 @@ class TestQuery:
assert np.array_equal(vec, ret)
def test_validate_dense_int8_convert(self):
v = _VectorQuery()
v = _SearchQuery()
schema = VectorSchema(name="test", data_type=DataType.VECTOR_INT8)
vec = np.array([1, 2, 3], dtype=np.int8)
v.set_vector(schema._get_object(), vec)
@ -142,7 +123,7 @@ class TestQuery:
assert np.array_equal(vec, ret)
def test_validate_sparse_fp32_convert(self):
v = _VectorQuery()
v = _SearchQuery()
schema = VectorSchema(name="test", data_type=DataType.SPARSE_VECTOR_FP32)
vec = {1: 1.1, 2: 2.2, 3: 3.3}
v.set_vector(schema._get_object(), vec)
@ -151,7 +132,7 @@ class TestQuery:
assert math.isclose(vec[k], ret[k], abs_tol=1e-6)
def test_validate_sparse_fp16_convert(self):
v = _VectorQuery()
v = _SearchQuery()
schema = VectorSchema(name="test", data_type=DataType.SPARSE_VECTOR_FP16)
vec = {1: 1.1, 2: 2.2, 3: 3.3}
v.set_vector(schema._get_object(), vec)
@ -189,7 +170,6 @@ class TestQueryContext:
assert ctx.reranker is None
assert ctx.output_fields is None
assert ctx.include_vector is False
assert ctx.core_vectors == []
def test_properties(self):
queries = [Query(field_name="test")]
@ -215,9 +195,7 @@ class TestQueryContext:
def test_properties_with_weighted_reranker(self):
queries = [Query(field_name="test")]
reranker = WeightedReRanker(
topn=10,
metrics={"test": MetricType.L2},
weights={"test": 1.0},
weights=[1.0],
)
ctx = QueryContext(
@ -227,13 +205,12 @@ class TestQueryContext:
)
assert ctx.reranker == reranker
assert ctx.reranker.weights == {"test": 1.0}
assert ctx.reranker.metrics == {"test": MetricType.L2}
assert ctx.reranker.weights == [1.0]
def test_properties_with_callback_reranker(self):
queries = [Query(field_name="test")]
cb = lambda query_results, topn: []
reranker = CallbackReRanker(callback=cb, topn=10)
reranker = CallbackReRanker(callback=cb)
ctx = QueryContext(
topk=5,
@ -243,141 +220,83 @@ class TestQueryContext:
assert ctx.reranker == reranker
def test_core_vectors_setter(self):
ctx = QueryContext(topk=10)
core_vectors = [MagicMock()]
ctx.core_vectors = core_vectors
assert ctx.core_vectors == core_vectors
class TestNoVectorQueryExecutor:
class TestQueryExecutor:
def test_init(self):
schema = MockCollectionSchema()
executor = NoVectorQueryExecutor(schema)
executor = QueryExecutor(schema)
assert isinstance(executor, QueryExecutor)
def test_do_validate_with_queries(self):
def test_do_build_without_queries(self):
# When no queries are given, build a single vector-less query.
schema = MockCollectionSchema()
executor = NoVectorQueryExecutor(schema)
ctx = QueryContext(
topk=10, queries=[Query(field_name="test", vector=[0.1, 0.2, 0.3])]
)
with pytest.raises(
ValueError, match="Collection does not support query with vector or id"
):
executor._do_validate(ctx)
def test_do_validate_without_queries(self):
schema = MockCollectionSchema()
executor = NoVectorQueryExecutor(schema)
ctx = QueryContext(topk=10)
executor._do_validate(ctx)
def test_do_build(self):
schema = MockCollectionSchema()
executor = NoVectorQueryExecutor(schema)
executor = QueryExecutor(schema)
ctx = QueryContext(topk=5, filter="test_filter")
result = executor._do_build(ctx, MagicMock())
result = executor._build_queries(ctx, MagicMock())
assert len(result) == 1
assert result[0].topk == 5
assert result[0].filter == "test_filter"
class TestSingleVectorQueryExecutor:
def test_init(self):
def test_do_build_query_wo_vector(self):
# Vector-less core query should carry the context query params.
schema = MockCollectionSchema()
executor = SingleVectorQueryExecutor(schema)
assert isinstance(executor, NoVectorQueryExecutor)
executor = QueryExecutor(schema)
ctx = QueryContext(topk=7, filter="f", include_vector=True)
def test_do_validate_multiple_queries(self):
core_vector = executor._build_base_search_query(ctx)
assert core_vector.topk == 7
assert core_vector.filter == "f"
assert core_vector.include_vector is True
def test_do_merge_rerank_results_single_without_reranker(self):
# A single result list without a reranker is returned as-is.
schema = MockCollectionSchema()
executor = SingleVectorQueryExecutor(schema)
queries = [Query(field_name="test1"), Query(field_name="test2")]
ctx = QueryContext(topk=10, queries=queries)
executor = QueryExecutor(schema)
ctx = QueryContext(topk=5)
docs_list = [["doc1", "doc2"]]
with pytest.raises(
ValueError,
match="Collection has only one vector field, cannot query with multiple vectors",
):
executor._do_validate(ctx)
result = executor._merge_and_rerank(ctx, docs_list)
assert result == ["doc1", "doc2"]
def test_do_build_without_queries(self):
def test_do_merge_rerank_results_empty(self):
# Empty results should raise an error.
schema = MockCollectionSchema()
executor = SingleVectorQueryExecutor(schema)
executor = QueryExecutor(schema)
ctx = QueryContext(topk=5)
result = executor._do_build(ctx, MagicMock())
assert len(result) == 1
assert result[0].topk == 5
with pytest.raises(ValueError, match="Query results is empty"):
executor._merge_and_rerank(ctx, [])
class TestMultiVectorQueryExecutor:
def test_init(self):
def test_do_merge_rerank_results_with_reranker(self):
# Multiple result lists are merged through the reranker.
schema = MockCollectionSchema()
executor = MultiVectorQueryExecutor(schema)
assert isinstance(executor, SingleVectorQueryExecutor)
def test_do_validate_multiple_queries_without_reranker(self):
schema = MockCollectionSchema()
executor = MultiVectorQueryExecutor(schema)
queries = [Query(field_name="test1"), Query(field_name="test2")]
ctx = QueryContext(topk=10, queries=queries)
with pytest.raises(ValueError, match="Reranker is required for multi-query"):
executor._do_validate(ctx)
def test_do_validate_multiple_queries_with_reranker(self):
schema = MockCollectionSchema()
executor = MultiVectorQueryExecutor(schema)
queries = [Query(field_name="test1"), Query(field_name="test2")]
reranker = RrfReRanker()
ctx = QueryContext(topk=10, queries=queries, reranker=reranker)
executor._do_validate(ctx)
def test_do_validate_multiple_queries_with_weighted_reranker(self):
schema = MockCollectionSchema()
executor = MultiVectorQueryExecutor(schema)
queries = [Query(field_name="test1"), Query(field_name="test2")]
reranker = WeightedReRanker(
topn=10,
metrics={"test1": MetricType.L2, "test2": MetricType.L2},
weights={"test1": 0.7, "test2": 0.3},
executor = QueryExecutor(schema)
reranker = MagicMock()
reranker.rerank.return_value = ["merged"]
ctx = QueryContext(
topk=5,
queries=[Query(field_name="test1"), Query(field_name="test2")],
reranker=reranker,
)
ctx = QueryContext(topk=10, queries=queries, reranker=reranker)
docs_list = [["d1"], ["d2"]]
executor._do_validate(ctx)
result = executor._merge_and_rerank(ctx, docs_list)
assert result == ["merged"]
reranker.rerank.assert_called_once_with(docs_list, ctx.topk)
def test_do_validate_multiple_queries_with_callback_reranker(self):
def test_execute_python_pipeline(self):
# Each query is executed serially and converted into a result list.
schema = MockCollectionSchema()
executor = MultiVectorQueryExecutor(schema)
queries = [Query(field_name="test1"), Query(field_name="test2")]
reranker = CallbackReRanker(
callback=lambda query_results, topn: [],
topn=10,
)
ctx = QueryContext(topk=10, queries=queries, reranker=reranker)
executor = QueryExecutor(schema)
collection = MagicMock()
collection.Query.side_effect = [["raw1"], ["raw2"]]
vectors = [MagicMock(), MagicMock()]
executor._do_validate(ctx)
class TestQueryExecutorFactory:
def test_create_no_vectors(self):
schema = MockCollectionSchema()
executor = QueryExecutorFactory.create(schema)
assert isinstance(executor, NoVectorQueryExecutor)
def test_create_single_vector(self):
schema = MockCollectionSchema(vectors=MockVectorSchema())
executor = QueryExecutorFactory.create(schema)
assert isinstance(executor, SingleVectorQueryExecutor)
def test_create_multiple_vectors(self):
schema = MockCollectionSchema(
vectors={"test1": MockVectorSchema(), "test2": MockVectorSchema()}
)
executor = QueryExecutorFactory.create(schema)
assert isinstance(executor, MultiVectorQueryExecutor)
with patch(
"zvec.executor.query_executor.convert_to_py_doc",
side_effect=lambda doc, schema: doc,
):
results = executor._execute_python_pipeline(vectors, collection)
assert results == [["raw1"], ["raw2"]]
assert collection.Query.call_count == 2

View File

@ -15,10 +15,9 @@ from __future__ import annotations
from unittest.mock import patch, MagicMock
import pytest
import math
import os
from zvec import Doc, MetricType
from zvec import Doc
from zvec.extension.multi_vector_reranker import (
CallbackReRanker,
RrfReRanker,
@ -38,37 +37,27 @@ RUN_INTEGRATION_TESTS = os.environ.get("ZVEC_RUN_INTEGRATION_TESTS", "0") == "1"
# ----------------------------
class TestRrfReRanker:
def test_init(self):
reranker = RrfReRanker(topn=5, rerank_field="content", rank_constant=100)
assert reranker.topn == 5
assert reranker.rerank_field == "content"
reranker = RrfReRanker(rank_constant=100)
assert reranker.rank_constant == 100
def test_rrf_score(self):
reranker = RrfReRanker(rank_constant=60)
# 根据公式 1.0 / (k + rank + 1)其中k=60
assert reranker._rrf_score(0) == 1.0 / (60 + 0 + 1)
assert reranker._rrf_score(1) == 1.0 / (60 + 1 + 1)
assert reranker._rrf_score(10) == 1.0 / (60 + 10 + 1)
def test_rerank(self):
reranker = RrfReRanker(topn=3)
def test_rerank_delegates_to_cpp(self):
"""RrfReRanker.rerank() delegates to C++ (raises TypeError with Python Docs)."""
reranker = RrfReRanker()
doc1 = Doc(id="1", score=0.8)
doc2 = Doc(id="2", score=0.7)
doc3 = Doc(id="3", score=0.9)
doc4 = Doc(id="4", score=0.6)
query_results = {"vector1": [doc1, doc2, doc3], "vector2": [doc3, doc1, doc4]}
query_results = [[doc1, doc2, doc3], [doc3, doc1, doc4]]
results = reranker.rerank(query_results)
with pytest.raises((TypeError, RuntimeError)):
reranker.rerank(query_results, topn=3)
assert len(results) <= reranker.topn
for doc in results:
assert hasattr(doc, "score")
scores = [doc.score for doc in results]
assert scores == sorted(scores, reverse=True)
def test_get_object_returns_cpp_reranker(self):
"""_get_object() returns a valid C++ reranker instance."""
reranker = RrfReRanker()
assert reranker._get_object() is not None
# ----------------------------
@ -76,64 +65,30 @@ class TestRrfReRanker:
# ----------------------------
class TestWeightedReRanker:
def test_init(self):
metrics = {"vector1": MetricType.L2, "vector2": MetricType.COSINE}
weights = {"vector1": 0.7, "vector2": 0.3}
weights = [0.7, 0.3]
reranker = WeightedReRanker(
topn=5,
rerank_field="content",
metrics=metrics,
weights=weights,
)
assert reranker.topn == 5
assert reranker.rerank_field == "content"
assert reranker.metrics == metrics
assert reranker.weights == weights
assert list(reranker.weights) == weights
def test_normalize_score(self):
reranker = WeightedReRanker()
score = reranker._normalize_score(1.0, MetricType.L2)
expected = 1.0 - 2 * math.atan(1.0) / math.pi
assert score == expected
score = reranker._normalize_score(1.0, MetricType.IP)
expected = 0.5 + math.atan(1.0) / math.pi
assert score == expected
score = reranker._normalize_score(1.0, MetricType.COSINE)
expected = 1.0 - 1.0 / 2.0
assert score == expected
with pytest.raises(ValueError, match="Unsupported metric type"):
reranker._normalize_score(1.0, "unsupported_metric")
def test_rerank(self):
metrics = {"vector1": MetricType.L2, "vector2": MetricType.L2}
weights = {"vector1": 0.7, "vector2": 0.3}
reranker = WeightedReRanker(topn=3, weights=weights, metrics=metrics)
def test_rerank_delegates_to_cpp(self):
"""WeightedReRanker.rerank() delegates to C++ (raises TypeError with Python Docs)."""
weights = [0.7, 0.3]
reranker = WeightedReRanker(weights=weights)
doc1 = Doc(id="1", score=0.8)
doc2 = Doc(id="2", score=0.7)
doc3 = Doc(id="3", score=0.9)
query_results = {"vector1": [doc1, doc2], "vector2": [doc2, doc3]}
query_results = [[doc1, doc2], [doc2, doc3]]
results = reranker.rerank(query_results)
with pytest.raises((TypeError, RuntimeError)):
reranker.rerank(query_results, topn=3)
assert len(results) <= reranker.topn
for doc in results:
assert hasattr(doc, "score")
def test_rerank_missing_metric_raises(self):
metrics = {"vector1": MetricType.L2}
reranker = WeightedReRanker(topn=3, metrics=metrics)
doc1 = Doc(id="1", score=0.8)
query_results = {"vector1": [doc1], "vector2": [doc1]}
with pytest.raises(ValueError, match="no metric type specified"):
reranker.rerank(query_results)
def test_get_object_returns_cpp_reranker(self):
"""_get_object() returns a valid C++ reranker instance."""
reranker = WeightedReRanker(weights=[0.5, 0.5])
assert reranker._get_object() is not None
# ----------------------------
@ -144,27 +99,27 @@ class TestCallbackReRanker:
def my_callback(query_results, topn):
return []
reranker = CallbackReRanker(callback=my_callback, topn=5)
assert reranker.topn == 5
reranker = CallbackReRanker(callback=my_callback)
assert reranker._get_object() is not None
def test_rerank(self):
def my_callback(query_results, topn):
all_docs = []
for docs in query_results.values():
for docs in query_results:
all_docs.extend(docs)
all_docs.sort(key=lambda d: d.score, reverse=True)
return all_docs[:topn]
reranker = CallbackReRanker(callback=my_callback, topn=3)
reranker = CallbackReRanker(callback=my_callback)
doc1 = Doc(id="1", score=0.8)
doc2 = Doc(id="2", score=0.9)
doc3 = Doc(id="3", score=0.7)
doc4 = Doc(id="4", score=0.6)
query_results = {"vector1": [doc1, doc2], "vector2": [doc3, doc4]}
query_results = [[doc1, doc2], [doc3, doc4]]
results = reranker.rerank(query_results)
results = reranker.rerank(query_results, topn=3)
assert len(results) == 3
scores = [doc.score for doc in results]
@ -177,8 +132,8 @@ class TestCallbackReRanker:
received_topn.append(topn)
return []
reranker = CallbackReRanker(callback=my_callback, topn=7)
reranker.rerank({"v1": [Doc(id="1", score=0.5)]})
reranker = CallbackReRanker(callback=my_callback)
reranker.rerank([[Doc(id="1", score=0.5)]], topn=7)
assert received_topn == [7]
@ -237,12 +192,6 @@ class TestQwenReRanker:
)
assert reranker.query == "test query"
def test_topn_property(self):
reranker = QwenReRanker(
query="test", topn=5, api_key="test_key", rerank_field="content"
)
assert reranker.topn == 5
def test_rerank_field_property(self):
reranker = QwenReRanker(query="test", api_key="test_key", rerank_field="title")
assert reranker.rerank_field == "title"
@ -251,7 +200,7 @@ class TestQwenReRanker:
reranker = QwenReRanker(
query="test", api_key="test_key", rerank_field="content"
)
results = reranker.rerank({})
results = reranker.rerank([], topn=10)
assert results == []
def test_rerank_no_valid_documents(self):
@ -259,22 +208,22 @@ class TestQwenReRanker:
query="test", api_key="test_key", rerank_field="content"
)
# Document without the rerank_field
query_results = {"vector1": [Doc(id="1")]}
query_results = [[Doc(id="1")]]
with pytest.raises(ValueError, match="No documents to rerank"):
reranker.rerank(query_results)
reranker.rerank(query_results, topn=10)
def test_rerank_skip_empty_content(self):
reranker = QwenReRanker(
query="test", api_key="test_key", rerank_field="content"
)
query_results = {
"vector1": [
query_results = [
[
Doc(id="1", fields={"content": ""}),
Doc(id="2", fields={"content": " "}),
]
}
]
with pytest.raises(ValueError, match="No documents to rerank"):
reranker.rerank(query_results)
reranker.rerank(query_results, topn=10)
@patch("zvec.extension.qwen_function.require_module")
def test_rerank_success(self, mock_require_module):
@ -294,17 +243,17 @@ class TestQwenReRanker:
mock_dashscope.TextReRank.call.return_value = mock_response
reranker = QwenReRanker(
query="test query", topn=2, api_key="test_key", rerank_field="content"
query="test query", api_key="test_key", rerank_field="content"
)
query_results = {
"vector1": [
query_results = [
[
Doc(id="1", fields={"content": "Document 1"}),
Doc(id="2", fields={"content": "Document 2"}),
]
}
]
results = reranker.rerank(query_results)
results = reranker.rerank(query_results, topn=2)
assert len(results) == 2
assert results[0].id == "1"
@ -338,14 +287,14 @@ class TestQwenReRanker:
mock_dashscope.TextReRank.call.return_value = mock_response
reranker = QwenReRanker(
query="test", topn=5, api_key="test_key", rerank_field="content"
query="test", api_key="test_key", rerank_field="content"
)
# Same document in multiple vector results
doc1 = Doc(id="1", fields={"content": "Document 1"})
query_results = {"vector1": [doc1], "vector2": [doc1]}
query_results = [[doc1], [doc1]]
results = reranker.rerank(query_results)
results = reranker.rerank(query_results, topn=5)
# Should only call API with document once
call_args = mock_dashscope.TextReRank.call.call_args
@ -368,10 +317,10 @@ class TestQwenReRanker:
query="test", api_key="test_key", rerank_field="content"
)
query_results = {"vector1": [Doc(id="1", fields={"content": "Document 1"})]}
query_results = [[Doc(id="1", fields={"content": "Document 1"})]]
with pytest.raises(ValueError, match="DashScope API error"):
reranker.rerank(query_results)
reranker.rerank(query_results, topn=10)
@patch("zvec.extension.qwen_function.require_module")
def test_rerank_runtime_error(self, mock_require_module):
@ -384,10 +333,10 @@ class TestQwenReRanker:
query="test", api_key="test_key", rerank_field="content"
)
query_results = {"vector1": [Doc(id="1", fields={"content": "Document 1"})]}
query_results = [[Doc(id="1", fields={"content": "Document 1"})]]
with pytest.raises(RuntimeError, match="Failed to call DashScope API"):
reranker.rerank(query_results)
reranker.rerank(query_results, topn=10)
@pytest.mark.skipif(
not RUN_INTEGRATION_TESTS,
@ -403,14 +352,13 @@ class TestQwenReRanker:
# Create reranker with real API
reranker = QwenReRanker(
query="What is machine learning?",
topn=3,
rerank_field="content",
model="gte-rerank-v2",
)
# Prepare test documents
query_results = {
"vector1": [
query_results = [
[
Doc(
id="1",
score=0.8,
@ -433,7 +381,7 @@ class TestQwenReRanker:
},
),
],
"vector2": [
[
Doc(
id="4",
score=0.6,
@ -449,10 +397,10 @@ class TestQwenReRanker:
},
),
],
}
]
# Call real API
results = reranker.rerank(query_results)
results = reranker.rerank(query_results, topn=3)
# Verify results
assert len(results) <= 3, "Should return at most topn documents"
@ -523,13 +471,11 @@ class TestDefaultLocalReRanker:
reranker = DefaultLocalReRanker(
query="test query",
topn=5,
rerank_field="content",
model_name="cross-encoder/ms-marco-MiniLM-L6-v2",
)
assert reranker.query == "test query"
assert reranker.topn == 5
assert reranker.rerank_field == "content"
assert reranker.model_name == "cross-encoder/ms-marco-MiniLM-L6-v2"
assert reranker.model_source == "huggingface"
@ -551,7 +497,6 @@ class TestDefaultLocalReRanker:
reranker = DefaultLocalReRanker(
query="custom query",
topn=10,
rerank_field="title",
model_name="cross-encoder/ms-marco-MiniLM-L12-v2",
model_source="modelscope",
@ -560,7 +505,6 @@ class TestDefaultLocalReRanker:
)
assert reranker.query == "custom query"
assert reranker.topn == 10
assert reranker.rerank_field == "title"
assert reranker.model_name == "cross-encoder/ms-marco-MiniLM-L12-v2"
assert reranker.model_source == "modelscope"
@ -593,23 +537,6 @@ class TestDefaultLocalReRanker:
reranker = DefaultLocalReRanker(query="test query", rerank_field="content")
assert reranker.query == "test query"
def test_topn_property(self):
"""Test topn property."""
mock_model = MagicMock()
mock_model.predict = MagicMock()
mock_st = MagicMock()
mock_st.CrossEncoder.return_value = mock_model
with patch(
"zvec.extension.sentence_transformer_rerank_function.require_module",
return_value=mock_st,
):
reranker = DefaultLocalReRanker(
query="test", topn=15, rerank_field="content"
)
assert reranker.topn == 15
def test_rerank_field_property(self):
"""Test rerank_field property."""
mock_model = MagicMock()
@ -655,7 +582,7 @@ class TestDefaultLocalReRanker:
return_value=mock_st,
):
reranker = DefaultLocalReRanker(query="test", rerank_field="content")
results = reranker.rerank({})
results = reranker.rerank([], topn=10)
assert results == []
def test_rerank_no_valid_documents(self):
@ -673,9 +600,9 @@ class TestDefaultLocalReRanker:
reranker = DefaultLocalReRanker(query="test", rerank_field="content")
# Document without the rerank_field
query_results = {"vector1": [Doc(id="1")]}
query_results = [[Doc(id="1")]]
with pytest.raises(ValueError, match="No documents to rerank"):
reranker.rerank(query_results)
reranker.rerank(query_results, topn=10)
def test_rerank_skip_empty_content(self):
"""Test rerank skips documents with empty content."""
@ -691,14 +618,14 @@ class TestDefaultLocalReRanker:
):
reranker = DefaultLocalReRanker(query="test", rerank_field="content")
query_results = {
"vector1": [
query_results = [
[
Doc(id="1", fields={"content": ""}),
Doc(id="2", fields={"content": " "}),
]
}
]
with pytest.raises(ValueError, match="No documents to rerank"):
reranker.rerank(query_results)
reranker.rerank(query_results, topn=10)
def test_rerank_success(self):
"""Test successful rerank with mocked model."""
@ -720,19 +647,17 @@ class TestDefaultLocalReRanker:
"zvec.extension.sentence_transformer_rerank_function.require_module",
return_value=mock_st,
):
reranker = DefaultLocalReRanker(
query="test query", topn=3, rerank_field="content"
)
reranker = DefaultLocalReRanker(query="test query", rerank_field="content")
query_results = {
"vector1": [
query_results = [
[
Doc(id="1", score=0.8, fields={"content": "Document 1"}),
Doc(id="2", score=0.7, fields={"content": "Document 2"}),
Doc(id="3", score=0.6, fields={"content": "Document 3"}),
]
}
]
results = reranker.rerank(query_results)
results = reranker.rerank(query_results, topn=3)
# Verify results
assert len(results) == 3
@ -771,21 +696,19 @@ class TestDefaultLocalReRanker:
"zvec.extension.sentence_transformer_rerank_function.require_module",
return_value=mock_st,
):
reranker = DefaultLocalReRanker(
query="test", topn=2, rerank_field="content"
)
reranker = DefaultLocalReRanker(query="test", rerank_field="content")
query_results = {
"vector1": [
query_results = [
[
Doc(id="1", fields={"content": "Doc 1"}),
Doc(id="2", fields={"content": "Doc 2"}),
Doc(id="3", fields={"content": "Doc 3"}),
Doc(id="4", fields={"content": "Doc 4"}),
Doc(id="5", fields={"content": "Doc 5"}),
]
}
]
results = reranker.rerank(query_results)
results = reranker.rerank(query_results, topn=2)
# Should only return top 2
assert len(results) == 2
@ -811,20 +734,18 @@ class TestDefaultLocalReRanker:
"zvec.extension.sentence_transformer_rerank_function.require_module",
return_value=mock_st,
):
reranker = DefaultLocalReRanker(
query="test", topn=5, rerank_field="content"
)
reranker = DefaultLocalReRanker(query="test", rerank_field="content")
# Same document in multiple vector results
doc1 = Doc(id="1", fields={"content": "Document 1"})
doc2 = Doc(id="2", fields={"content": "Document 2"})
query_results = {
"vector1": [doc1, doc2],
"vector2": [doc1], # doc1 appears in both
}
query_results = [
[doc1, doc2],
[doc1], # doc1 appears in both
]
results = reranker.rerank(query_results)
results = reranker.rerank(query_results, topn=5)
# Should only process each document once
assert len(results) == 2
@ -852,19 +773,17 @@ class TestDefaultLocalReRanker:
"zvec.extension.sentence_transformer_rerank_function.require_module",
return_value=mock_st,
):
reranker = DefaultLocalReRanker(
query="test", topn=3, rerank_field="content"
)
reranker = DefaultLocalReRanker(query="test", rerank_field="content")
query_results = {
"vector1": [
query_results = [
[
Doc(id="1", fields={"content": "Doc 1"}),
Doc(id="2", fields={"content": "Doc 2"}),
Doc(id="3", fields={"content": "Doc 3"}),
]
}
]
results = reranker.rerank(query_results)
results = reranker.rerank(query_results, topn=3)
# Should be sorted by score (descending)
assert len(results) == 3
@ -892,10 +811,10 @@ class TestDefaultLocalReRanker:
):
reranker = DefaultLocalReRanker(query="test", rerank_field="content")
query_results = {"vector1": [Doc(id="1", fields={"content": "Document 1"})]}
query_results = [[Doc(id="1", fields={"content": "Document 1"})]]
with pytest.raises(RuntimeError, match="Failed to compute rerank scores"):
reranker.rerank(query_results)
reranker.rerank(query_results, topn=10)
def test_rerank_with_custom_batch_size(self):
"""Test rerank uses custom batch_size."""
@ -918,14 +837,14 @@ class TestDefaultLocalReRanker:
query="test", rerank_field="content", batch_size=64
)
query_results = {
"vector1": [
query_results = [
[
Doc(id="1", fields={"content": "Doc 1"}),
Doc(id="2", fields={"content": "Doc 2"}),
]
}
]
reranker.rerank(query_results)
reranker.rerank(query_results, topn=10)
# Verify batch_size is passed to predict
call_args = mock_model.predict.call_args
@ -947,13 +866,12 @@ class TestDefaultLocalReRanker:
# Create reranker with real model (using default lightweight model)
reranker = DefaultLocalReRanker(
query="What is machine learning?",
topn=3,
rerank_field="content",
)
# Prepare test documents
query_results = {
"vector1": [
query_results = [
[
Doc(
id="1",
score=0.8,
@ -976,7 +894,7 @@ class TestDefaultLocalReRanker:
},
),
],
"vector2": [
[
Doc(
id="4",
score=0.6,
@ -992,10 +910,10 @@ class TestDefaultLocalReRanker:
},
),
],
}
]
# Call real model
results = reranker.rerank(query_results)
results = reranker.rerank(query_results, topn=3)
# Verify results
assert len(results) <= 3, "Should return at most topn documents"
@ -1030,3 +948,49 @@ class TestDefaultLocalReRanker:
content = doc.field("content")
if content:
print(f" Content: {content[:80]}...")
# ----------------------------
# DocList Type and Delegation Tests
# ----------------------------
class TestDocList:
def test_type_alias(self):
"""DocList is list[Doc]."""
from zvec.model.doc import DocList
from zvec import Doc, DocList as QR
assert DocList == list[Doc]
assert QR == list[Doc]
def test_rrf_reranker_delegates_to_cpp(self):
"""RrfReRanker.rerank() delegates to C++ (raises TypeError with Python Docs)."""
reranker = RrfReRanker()
with pytest.raises(TypeError):
reranker.rerank([[Doc(id="1", score=0.5)]], topn=5)
def test_weighted_reranker_delegates_to_cpp(self):
"""WeightedReRanker.rerank() delegates to C++ (raises TypeError with Python Docs)."""
reranker = WeightedReRanker(weights=[0.7, 0.3])
with pytest.raises(TypeError):
reranker.rerank(
[[Doc(id="1", score=0.5)], [Doc(id="2", score=0.3)]], topn=5
)
def test_single_route_query_results(self):
"""CallbackReRanker works with single-route (one element list)."""
def cb(query_results, topn):
return query_results[0][:topn]
reranker = CallbackReRanker(callback=cb)
results = reranker.rerank(
[
[
Doc(id="1", score=0.9),
Doc(id="2", score=0.8),
Doc(id="3", score=0.7),
]
],
topn=2,
)
assert len(results) == 2

View File

@ -71,7 +71,7 @@ from .model import schema as schema
# —— Core data structures ——
from .model.collection import Collection
from .model.doc import Doc
from .model.doc import Doc, DocList
# —— Query & index parameters ——
# —— FTS params (C++ binding) ——
@ -127,6 +127,7 @@ __all__ = [
# Core classes
"Collection",
"Doc",
"DocList",
# Schema
"CollectionSchema",
"FieldSchema",

View File

@ -11,7 +11,7 @@ from .extension import ReRanker, RrfReRanker, WeightedReRanker
from .extension.embedding import DenseEmbeddingFunction
from .model import param, schema
from .model.collection import Collection
from .model.doc import Doc
from .model.doc import Doc, DocList
from .model.param import (
AddColumnOption,
AlterColumnOption,
@ -52,8 +52,8 @@ __all__: list = [
"CollectionStats",
"DataType",
"DenseEmbeddingFunction",
"DenseEmbeddingFunction",
"Doc",
"DocList",
"FieldSchema",
"FlatIndexParam",
"HnswIndexParam",
@ -72,7 +72,6 @@ __all__: list = [
"QuantizeType",
"Query",
"ReRanker",
"ReRanker",
"RrfReRanker",
"Status",
"StatusCode",
@ -127,7 +126,7 @@ class _Collection:
def Optimize(self, arg0: param.OptimizeOption) -> None: ...
def Options(self) -> param.CollectionOption: ...
def Path(self) -> str: ...
def Query(self, arg0: param._VectorQuery) -> list[_Doc]: ...
def Query(self, arg0: param._SearchQuery) -> list[_Doc]: ...
def Schema(self) -> schema._CollectionSchema: ...
def Stats(self) -> schema.CollectionStats: ...
def Update(self, arg0: collections.abc.Sequence[_Doc]) -> list[typing.Status]: ...

View File

@ -16,11 +16,9 @@ from __future__ import annotations
from .query_executor import (
QueryContext,
QueryExecutor,
QueryExecutorFactory,
)
__all__ = [
"QueryContext",
"QueryExecutor",
"QueryExecutorFactory",
]

View File

@ -13,18 +13,15 @@
# limitations under the License.
from __future__ import annotations
import os
from abc import ABC, abstractmethod
from concurrent.futures import ThreadPoolExecutor, as_completed
from typing import Optional, Union, final
from typing import Optional, Union
import numpy as np
from _zvec import _Collection, _MultiQuery
from _zvec.param import _Fts, _SubQuery, _VectorQuery
from _zvec.param import _Fts, _SearchQuery, _SubQuery
from ..extension import ReRanker, RrfReRanker, WeightedReRanker
from ..extension import ReRanker
from ..model.convert import convert_to_py_doc
from ..model.doc import Doc
from ..model.doc import DocList
from ..model.param.query import Query
from ..model.schema import CollectionSchema
from ..typing import DataType
@ -32,7 +29,6 @@ from ..typing import DataType
__all__ = [
"QueryContext",
"QueryExecutor",
"QueryExecutorFactory",
]
DTYPE_MAP = {
@ -80,9 +76,6 @@ class QueryContext:
# reranker
self._reranker = reranker
# core vectors
self._core_vectors = []
@property
def topk(self):
return self._topk
@ -107,61 +100,120 @@ class QueryContext:
def include_vector(self):
return self._include_vector
@property
def core_vectors(self):
return self._core_vectors
@core_vectors.setter
def core_vectors(self, core_vectors: list[_VectorQuery]):
self._core_vectors = core_vectors
class QueryExecutor:
"""Unified query executor that routes based on query count and reranker type."""
class QueryExecutor(ABC):
def __init__(self, schema: CollectionSchema):
self._schema = schema
self._concurrency = max(1, int(os.getenv("ZVEC_QUERY_CONCURRENCY", "1")))
@abstractmethod
def _do_validate(self, ctx: QueryContext) -> None:
pass
@abstractmethod
def _do_build(
def _build_queries(
self, ctx: QueryContext, collection: _Collection
) -> list[_VectorQuery]:
pass
) -> list[_SearchQuery]:
"""Build query vector list (no validation, conversion only)."""
if not ctx.queries:
return [self._build_base_search_query(ctx)]
return [
self._build_search_query(ctx, query, collection) for query in ctx.queries
]
def _do_build_query_wo_vector(self, ctx: QueryContext) -> _VectorQuery:
core_vector = _VectorQuery()
core_vector.topk = ctx.topk
core_vector.include_vector = ctx.include_vector
def execute(self, ctx: QueryContext, collection: _Collection) -> DocList:
"""Execute a query, routing by query count.
A single (or vector-less) query is sent to C++ as a ``_SearchQuery``;
multiple queries are assembled into a ``_MultiQuery``.
"""
queries = self._build_queries(ctx, collection)
if not queries:
raise ValueError("No query to execute")
if len(queries) == 1:
return self._execute_single_query(queries[0], collection)
return self._execute_multi_query(ctx, queries, collection)
def _execute_single_query(
self, query: _SearchQuery, collection: _Collection
) -> DocList:
"""Single/vector-less query: send a ``_SearchQuery`` to C++."""
docs = collection.Query(query)
return [convert_to_py_doc(doc, self._schema) for doc in docs]
def _execute_multi_query(
self, ctx: QueryContext, queries: list[_SearchQuery], collection: _Collection
) -> DocList:
"""Multiple queries: send a ``_MultiQuery`` to C++.
A Python-only reranker (``_get_object()`` returns None) cannot run
inside the C++ MultiQuery, so each route is executed individually and
merged by the reranker in Python.
"""
reranker = ctx.reranker
if reranker is not None and reranker._get_object() is None:
docs_list = self._execute_python_pipeline(queries, collection)
return self._merge_and_rerank(ctx, docs_list)
multi_query = self._build_multi_query(ctx, queries)
docs = collection.Query(multi_query)
return [convert_to_py_doc(doc, self._schema) for doc in docs]
def _build_multi_query(
self, ctx: QueryContext, queries: list[_SearchQuery]
) -> _MultiQuery:
"""Assemble a C++ ``_MultiQuery`` from per-route ``_SearchQuery`` objects."""
multi_query = _MultiQuery()
multi_query.queries = [_SubQuery.from_search_query(query) for query in queries]
multi_query.topk = ctx.topk
if ctx.filter:
core_vector.filter = ctx.filter
if ctx.output_fields:
core_vector.output_fields = ctx.output_fields
return core_vector
multi_query.filter = ctx.filter
multi_query.include_vector = ctx.include_vector
if ctx.output_fields is not None:
multi_query.output_fields = ctx.output_fields
if ctx.reranker is not None:
multi_query.reranker = ctx.reranker._get_object()
return multi_query
def _do_build_fts_query(self, query: Query, core_vector: _VectorQuery) -> None:
"""Set FTS query on core_vector if the query has FTS parameters."""
def _execute_python_pipeline(
self, vectors: list[_SearchQuery], collection: _Collection
) -> list[DocList]:
"""Execute queries serially for the Python-only reranker path."""
return [self._execute_single_query(query, collection) for query in vectors]
def _merge_and_rerank(self, ctx: QueryContext, docs_list: list[DocList]) -> DocList:
"""Merge and rerank results from the Python pipeline path."""
if not docs_list:
raise ValueError("Query results is empty")
if len(docs_list) == 1 and not ctx.reranker:
return docs_list[0]
return ctx.reranker.rerank(docs_list, ctx.topk)
def _build_base_search_query(self, ctx: QueryContext) -> _SearchQuery:
search_query = _SearchQuery()
search_query.topk = ctx.topk
search_query.include_vector = ctx.include_vector
if ctx.filter:
search_query.filter = ctx.filter
if ctx.output_fields is not None:
search_query.output_fields = ctx.output_fields
return search_query
def _apply_fts(self, query: Query, search_query: _SearchQuery) -> None:
"""Set FTS query on search_query if the query has FTS parameters."""
if query.has_fts():
fts = _Fts()
fts.query_string = query.fts.query_string or ""
fts.match_string = query.fts.match_string or ""
core_vector.fts = fts
search_query.fts = fts
def _do_build_query_with_vector(
def _build_search_query(
self, ctx: QueryContext, query: Query, collection: _Collection
) -> _VectorQuery:
core_vector = self._do_build_query_wo_vector(ctx)
core_vector.field_name = query.field_name
) -> _SearchQuery:
search_query = self._build_base_search_query(ctx)
search_query.field_name = query.field_name
if query.param:
core_vector.query_params = query.param
search_query.query_params = query.param
# set FTS query if provided
self._do_build_fts_query(query, core_vector)
# set output_fields
core_vector.output_fields = ctx.output_fields
self._apply_fts(query, search_query)
vector_schema = None
if query.has_vector() or query.has_id():
@ -181,189 +233,14 @@ class QueryExecutor(ABC):
fetched = collection.Fetch([query.id])
doc = next(iter(fetched.values()))
if not doc:
return core_vector
raise ValueError(f"Document with id '{query.id}' not found")
vec_data = doc.get_any(vector_schema.name, vector_schema.data_type)
else:
return core_vector
return search_query
target_dtype = DTYPE_MAP.get(vector_schema.data_type.value)
core_vector.set_vector(
search_query.set_vector(
vector_schema._get_object(),
convert_to_numpy(vec_data, target_dtype) if target_dtype else vec_data,
)
return core_vector
def _do_execute(
self, vectors: list[_VectorQuery], collection: _Collection
) -> dict[str, list[Doc]]:
query_cnt = len(vectors)
if query_cnt == 0:
raise ValueError("No query to execute")
if len(vectors) == 1 or self._concurrency == 1:
results = {}
for query in vectors:
docs = collection.Query(query)
results[query.field_name] = [
convert_to_py_doc(doc, self._schema) for doc in docs
]
return results
results = {}
with ThreadPoolExecutor(max_workers=self._concurrency) as executor:
future_to_query = {
executor.submit(collection.Query, query): query.field_name
for query in vectors
}
for future in as_completed(future_to_query):
field_name = future_to_query[future]
try:
docs = future.result()
results[field_name] = [
convert_to_py_doc(doc, self._schema) for doc in docs
]
except Exception as e:
raise e
return results
def _do_merge_rerank_results(
self, ctx: QueryContext, docs_map: dict[str, list[Doc]]
) -> list[Doc]:
query_result_cnt = len(docs_map) if docs_map else 0
if query_result_cnt == 0:
raise ValueError("Query results is none and dost not to rerank")
if query_result_cnt == 1:
if not ctx.reranker or isinstance(
ctx.reranker, (RrfReRanker, WeightedReRanker)
):
return next(iter(docs_map.values()))
return ctx.reranker.rerank(docs_map)
return ctx.reranker.rerank(docs_map)
@final
def execute(self, ctx: QueryContext, collection: _Collection) -> list[Doc]:
# 1. validate query
self._do_validate(ctx)
# 2. build query vector
query_vectors = self._do_build(ctx, collection)
if not query_vectors:
raise ValueError("No query to execute")
# 3. execute query
docs = self._do_execute(query_vectors, collection)
# 4. merge and rerank result
return self._do_merge_rerank_results(ctx, docs)
class NoVectorQueryExecutor(QueryExecutor):
def __init__(self, schema: CollectionSchema):
super().__init__(schema)
def _do_validate(self, ctx: QueryContext) -> None:
for query in ctx.queries:
if query.has_vector() or query.has_id():
raise ValueError("Collection does not support query with vector or id")
query._validate()
def _do_build(
self, ctx: QueryContext, collection: _Collection
) -> list[_VectorQuery]:
if len(ctx.queries) == 0:
return [self._do_build_query_wo_vector(ctx)]
# FTS-only branch in _do_build_query_with_vector skips vector resolution.
return [
self._do_build_query_with_vector(ctx, query, collection)
for query in ctx.queries
]
class SingleVectorQueryExecutor(NoVectorQueryExecutor):
def __init__(self, schema: CollectionSchema) -> None:
super().__init__(schema)
def _validate_multi_query(self, ctx: QueryContext) -> None:
"""Shared validation for multi-query: reranker required + no duplicate fields."""
if ctx.reranker is None:
raise ValueError("Reranker is required for multi-query")
seen_fields = set()
for query in ctx.queries:
query._validate()
if query.field_name in seen_fields:
raise ValueError(
f"Query field name '{query.field_name}' appears more than once"
)
seen_fields.add(query.field_name)
def _do_validate(self, ctx: QueryContext) -> None:
if len(ctx.queries) > 1:
# Allow FTS + vector hybrid multi-query (requires reranker)
if not any(q.has_fts() for q in ctx.queries):
raise ValueError(
"Collection has only one vector field, cannot query with multiple vectors"
)
self._validate_multi_query(ctx)
return
for query in ctx.queries:
query._validate()
def _do_build(
self, ctx: QueryContext, collection: _Collection
) -> list[_VectorQuery]:
if len(ctx.queries) == 0:
return [self._do_build_query_wo_vector(ctx)]
vectors = []
for query in ctx.queries:
vectors.append(self._do_build_query_with_vector(ctx, query, collection))
return vectors
def execute(self, ctx: QueryContext, collection: _Collection) -> list[Doc]:
# 1. validate query
self._do_validate(ctx)
# 2. build query vectors
query_vectors = self._do_build(ctx, collection)
if not query_vectors:
raise ValueError("No query to execute")
# Multi-query fast path: route FTS + vector hybrid to C++ MultiQuery
if len(query_vectors) > 1 and ctx.reranker is not None:
cpp_reranker = ctx.reranker._get_object()
if cpp_reranker is not None:
mvq = _MultiQuery()
mvq.queries = [_SubQuery.from_vector_query(vq) for vq in query_vectors]
mvq.topk = ctx.topk
if ctx.filter:
mvq.filter = ctx.filter
mvq.include_vector = ctx.include_vector
if ctx.output_fields:
mvq.output_fields = ctx.output_fields
mvq.reranker = cpp_reranker
docs = collection.Query(mvq)
return [convert_to_py_doc(doc, self._schema) for doc in docs]
# 3. execute query
docs = self._do_execute(query_vectors, collection)
# 4. merge and rerank result
return self._do_merge_rerank_results(ctx, docs)
class MultiVectorQueryExecutor(SingleVectorQueryExecutor):
def __init__(self, schema: CollectionSchema) -> None:
super().__init__(schema)
def _do_validate(self, ctx: QueryContext) -> None:
if len(ctx.queries) > 1:
self._validate_multi_query(ctx)
return
for query in ctx.queries:
query._validate()
class QueryExecutorFactory:
@staticmethod
def create(schema: CollectionSchema) -> QueryExecutor:
vectors = schema.vectors
if len(vectors) == 0:
return NoVectorQueryExecutor(schema)
if len(vectors) == 1:
return SingleVectorQueryExecutor(schema)
return MultiVectorQueryExecutor(schema)
return search_query

View File

@ -13,16 +13,12 @@
# limitations under the License.
from __future__ import annotations
import heapq
import math
from collections import defaultdict
from collections.abc import Callable
from typing import Optional
from _zvec import _CallbackReranker, _RrfReranker, _WeightedReranker
from ..model.doc import Doc
from ..typing import MetricType
from ..model.doc import DocList
from .rerank_function import RerankFunction
@ -35,24 +31,15 @@ class RrfReRanker(RerankFunction):
The RRF score for a document at rank ``r`` is: ``1 / (k + r + 1)``,
where ``k`` is the rank constant.
Note:
This re-ranker is specifically designed for multi-vector scenarios where
query results from multiple vector fields need to be combined.
Args:
topn (int, optional): Number of top documents to return. Defaults to 10.
rerank_field (Optional[str], optional): Ignored by RRF. Defaults to None.
rank_constant (int, optional): Smoothing constant ``k`` in RRF formula.
Larger values reduce the impact of early ranks. Defaults to 60.
"""
def __init__(
self,
topn: int = 10,
rerank_field: Optional[str] = None,
rank_constant: int = 60,
):
super().__init__(topn=topn, rerank_field=rerank_field)
self._rank_constant = rank_constant
# Use C++ implementation for performance
self._cpp_reranker = _RrfReranker(rank_constant)
@ -65,36 +52,18 @@ class RrfReRanker(RerankFunction):
"""Return the underlying C++ RrfReranker instance."""
return self._cpp_reranker
def _rrf_score(self, rank: int) -> float:
return 1.0 / (self._rank_constant + rank + 1)
def rerank(self, query_results: dict[str, list[Doc]]) -> list[Doc]:
"""Apply Reciprocal Rank Fusion to combine multiple query results.
def rerank(self, query_results: list[DocList], topn: int) -> DocList:
"""Re-rank using C++ RRF implementation.
Args:
query_results (dict[str, list[Doc]]): Results from one or more vector queries.
query_results (list[DocList]): Multi-route recall results,
positionally aligned with queries.
topn (int): Number of top documents to return.
Returns:
list[Doc]: Re-ranked documents with RRF scores in the ``score`` field.
DocList: Re-ranked documents.
"""
rrf_scores: dict[str, float] = defaultdict(float)
id_to_doc: dict[str, Doc] = {}
for _, query_result in query_results.items():
for rank, doc in enumerate(query_result):
doc_id = doc.id
rrf_score = self._rrf_score(rank)
rrf_scores[doc_id] += rrf_score
if doc_id not in id_to_doc:
id_to_doc[doc_id] = doc
top_docs = heapq.nlargest(self.topn, rrf_scores.items(), key=lambda x: x[1])
results: list[Doc] = []
for doc_id, rrf_score in top_docs:
doc = id_to_doc[doc_id]
new_doc = doc._replace(score=rrf_score)
results.append(new_doc)
return results
return self._cpp_reranker.rerank(query_results, topn)
class WeightedReRanker(RerankFunction):
@ -102,98 +71,41 @@ class WeightedReRanker(RerankFunction):
Each vector field's relevance score is normalized based on its own metric
type, then scaled by a user-provided weight. Final scores are summed across
fields.
Note:
This re-ranker is specifically designed for multi-vector scenarios where
query results from multiple vector fields need to be combined with
configurable weights.
fields. The actual re-ranking logic lives in the C++ implementation.
Args:
topn (int, optional): Number of top documents to return. Defaults to 10.
rerank_field (Optional[str], optional): Ignored. Defaults to None.
metrics (Optional[dict[str, MetricType]], optional): Per-field distance
metric used for score normalization. Every queried field must have
a metric specified; missing fields will raise an error at rerank time.
Defaults to None.
weights (Optional[dict[str, float]], optional): Weight per vector field.
Fields not listed use weight 1.0. Defaults to None.
Note:
Supported metrics: L2, IP, COSINE. Scores are normalized to [0, 1].
weights (Optional[list[float]], optional): Weight per vector field,
aligned by position with the queries supplied to ``collection.query()``.
Defaults to None (treated as an empty list).
"""
def __init__(
self,
topn: int = 10,
rerank_field: Optional[str] = None,
metrics: Optional[dict[str, MetricType]] = None,
weights: Optional[dict[str, float]] = None,
weights: Optional[list[float]] = None,
):
super().__init__(topn=topn, rerank_field=rerank_field)
self._weights = weights or {}
self._metrics = metrics or {}
self._cpp_reranker = _WeightedReranker(weights or [])
@property
def weights(self) -> dict[str, float]:
"""dict[str, float]: Weight mapping for vector fields."""
return self._weights
@property
def metrics(self) -> dict[str, MetricType]:
"""dict[str, MetricType]: Per-field metric type mapping."""
return self._metrics
def weights(self) -> list[float]:
"""list[float]: Weight list for vector fields, aligned with queries."""
return self._cpp_reranker.weights
def _get_object(self):
"""Return a C++ WeightedReranker instance."""
return _WeightedReranker(self._weights)
"""Return the underlying C++ WeightedReranker instance."""
return self._cpp_reranker
def rerank(self, query_results: dict[str, list[Doc]]) -> list[Doc]:
"""Combine scores from multiple vector fields using weighted sum.
def rerank(self, query_results: list[DocList], topn: int) -> DocList:
"""Re-rank using C++ Weighted implementation.
Args:
query_results (dict[str, list[Doc]]): Results per vector field.
query_results (list[DocList]): Multi-route recall results,
positionally aligned with queries.
topn (int): Number of top documents to return.
Returns:
list[Doc]: Re-ranked documents with combined scores in ``score`` field.
DocList: Re-ranked documents.
"""
weighted_scores: dict[str, float] = defaultdict(float)
id_to_doc: dict[str, Doc] = {}
for vector_name, query_result in query_results.items():
if vector_name not in self._metrics:
raise ValueError(
f"WeightedReRanker: no metric type specified for field "
f"'{vector_name}'"
)
metric = self._metrics[vector_name]
for _, doc in enumerate(query_result):
doc_id = doc.id
weighted_score = self._normalize_score(
doc.score, metric
) * self.weights.get(vector_name, 1.0)
weighted_scores[doc_id] += weighted_score
if doc_id not in id_to_doc:
id_to_doc[doc_id] = doc
top_docs = heapq.nlargest(
self.topn, weighted_scores.items(), key=lambda x: x[1]
)
results: list[Doc] = []
for doc_id, weighted_score in top_docs:
doc = id_to_doc[doc_id]
new_doc = doc._replace(score=weighted_score)
results.append(new_doc)
return results
def _normalize_score(self, score: float, metric: MetricType) -> float:
if metric == MetricType.L2:
return 1.0 - 2 * math.atan(score) / math.pi
if metric == MetricType.IP:
return 0.5 + math.atan(score) / math.pi
if metric == MetricType.COSINE:
return 1.0 - score / 2.0
raise ValueError("Unsupported metric type")
return self._cpp_reranker.rerank(query_results, topn)
class CallbackReRanker(RerankFunction):
@ -202,21 +114,18 @@ class CallbackReRanker(RerankFunction):
This bridges a Python callable into the C++ reranker interface, enabling
custom re-ranking logic to be executed within the C++ MultiQuery path.
The callback receives the raw C++ Doc objects (as ``_Doc`` instances) grouped
by vector field name, and must return a list of ``_Doc`` instances.
The callback receives raw C++ ``_Doc`` objects grouped per query (as a
``list[list[_Doc]]``) and must return a ``list[_Doc]``.
Args:
callback: A callable with signature
``(query_results: dict[str, list[_Doc]], topn: int) -> list[_Doc]``.
topn (int, optional): Number of top documents to return. Defaults to 10.
``(query_results: list[list[_Doc]], topn: int) -> list[_Doc]``.
"""
def __init__(
self,
callback: Callable,
topn: int = 10,
):
super().__init__(topn=topn)
self._callback = callback
self._cpp_reranker = _CallbackReranker(callback)
@ -224,13 +133,15 @@ class CallbackReRanker(RerankFunction):
"""Return the underlying C++ CallbackReranker instance."""
return self._cpp_reranker
def rerank(self, query_results: dict[str, list[Doc]]) -> list[Doc]:
def rerank(self, query_results: list[DocList], topn: int) -> DocList:
"""Invoke the callback to re-rank documents.
Args:
query_results (dict[str, list[Doc]]): Results per vector field.
query_results (list[DocList]): Multi-route recall results,
positionally aligned with queries.
topn (int): Number of top documents to return.
Returns:
list[Doc]: Re-ranked documents.
DocList: Re-ranked documents.
"""
return self._callback(query_results, self.topn)
return self._callback(query_results, topn)

View File

@ -15,7 +15,7 @@ from __future__ import annotations
from typing import Optional
from ..model.doc import Doc
from ..model.doc import Doc, DocList
from .qwen_function import QwenFunctionBase
from .rerank_function import RerankFunction
@ -32,8 +32,6 @@ class QwenReRanker(QwenFunctionBase, RerankFunction):
Args:
query (str): Query text for semantic re-ranking. **Required**.
topn (int, optional): Maximum number of documents to return after re-ranking.
Defaults to 10.
rerank_field (str): Document field name to use as re-ranking input text.
**Required** (e.g., "content", "title", "body").
model (str, optional): DashScope re-ranking model identifier.
@ -53,7 +51,6 @@ class QwenReRanker(QwenFunctionBase, RerankFunction):
Example:
>>> reranker = QwenReRanker(
... query="machine learning algorithms",
... topn=5,
... rerank_field="content",
... model="gte-rerank-v2",
... api_key="your-api-key"
@ -64,7 +61,6 @@ class QwenReRanker(QwenFunctionBase, RerankFunction):
def __init__(
self,
query: Optional[str] = None,
topn: int = 10,
rerank_field: Optional[str] = None,
model: str = "gte-rerank-v2",
api_key: Optional[str] = None,
@ -73,7 +69,6 @@ class QwenReRanker(QwenFunctionBase, RerankFunction):
Args:
query (Optional[str]): Query text for semantic matching. Required.
topn (int): Number of top results to return.
rerank_field (Optional[str]): Document field for re-ranking input.
model (str): DashScope model name.
api_key (Optional[str]): API key or None to use environment variable.
@ -82,37 +77,45 @@ class QwenReRanker(QwenFunctionBase, RerankFunction):
ValueError: If query is empty or API key is unavailable.
"""
QwenFunctionBase.__init__(self, model=model, api_key=api_key)
RerankFunction.__init__(self, topn=topn, rerank_field=rerank_field)
RerankFunction.__init__(self)
if not query:
raise ValueError("Query is required for QwenReRanker")
self._query = query
self._rerank_field = rerank_field
@property
def query(self) -> str:
"""str: Query text used for semantic re-ranking."""
return self._query
def rerank(self, query_results: dict[str, list[Doc]]) -> list[Doc]:
@property
def rerank_field(self) -> Optional[str]:
"""Optional[str]: Field name used as re-ranking input."""
return self._rerank_field
def rerank(self, query_results: list[DocList], topn: int) -> DocList:
"""Re-rank documents using Qwen's TextReRank API.
Sends document texts to DashScope TextReRank service along with the query.
Returns documents sorted by relevance scores from the cross-encoder model.
Args:
query_results (dict[str, list[Doc]]): Mapping from vector field names
to lists of retrieved documents. Documents from all fields are
query_results (list[DocList]): Multi-route recall results,
positionally aligned with the queries supplied to
``collection.query()``. Documents from all routes are
deduplicated and re-ranked together.
topn (int): Maximum number of documents to return after re-ranking.
Returns:
list[Doc]: Re-ranked documents (up to ``topn``) with updated ``score``
fields containing relevance scores from the API.
DocList: Re-ranked documents (up to ``topn``) with updated
``score`` fields containing relevance scores from the API.
Raises:
ValueError: If no valid documents are found or API call fails.
Note:
- Duplicate documents (same ID) across fields are processed once
- Duplicate documents (same ID) across routes are processed once
- Documents with empty/missing ``rerank_field`` content are skipped
- Returned scores are relevance scores from the cross-encoder model
"""
@ -124,7 +127,7 @@ class QwenReRanker(QwenFunctionBase, RerankFunction):
doc_ids: list[str] = []
contents: list[str] = []
for _, query_result in query_results.items():
for query_result in query_results:
for doc in query_result:
doc_id = doc.id
if doc_id in id_to_doc:
@ -147,11 +150,11 @@ class QwenReRanker(QwenFunctionBase, RerankFunction):
output = self._call_rerank_api(
query=self.query,
documents=contents,
top_n=self.topn,
top_n=topn,
)
# Build result list with updated scores
results: list[Doc] = []
results: DocList = []
for item in output["results"]:
idx = item["index"]
doc_id = doc_ids[idx]

View File

@ -14,9 +14,8 @@
from __future__ import annotations
from abc import ABC, abstractmethod
from typing import Optional
from ..model.doc import Doc
from ..model.doc import DocList
class RerankFunction(ABC):
@ -26,44 +25,22 @@ class RerankFunction(ABC):
a secondary scoring strategy. They are used in the ``query()`` method of
``Collection`` via the ``reranker`` parameter.
Args:
topn (int, optional): Number of top documents to return after re-ranking.
Defaults to 10.
rerank_field (Optional[str], optional): Field name used as input for
re-ranking (e.g., document title or body). Defaults to None.
Note:
Subclasses must implement the ``rerank()`` method.
"""
def __init__(
self,
topn: int = 10,
rerank_field: Optional[str] = None,
):
self._topn = topn
self._rerank_field = rerank_field
@property
def topn(self) -> int:
"""int: Number of top documents to return after re-ranking."""
return self._topn
@property
def rerank_field(self) -> Optional[str]:
"""Optional[str]: Field name used as re-ranking input."""
return self._rerank_field
@abstractmethod
def rerank(self, query_results: dict[str, list[Doc]]) -> list[Doc]:
"""Re-rank documents from one or more vector queries.
def rerank(self, query_results: list[DocList], topn: int) -> DocList:
"""Re-rank documents from multi-route recall results.
Args:
query_results (dict[str, list[Doc]]): Mapping from vector field name
to list of retrieved documents (sorted by relevance).
query_results (list[DocList]): List of query results from
multi-route recall. Each element corresponds to a Query in the
collection.query(queries=List[Query]) call, aligned by position.
topn (int): Number of top documents to return after re-ranking.
Returns:
list[Doc]: Re-ranked list of documents (length ``topn``),
DocList: Re-ranked list of documents (length ``topn``),
with updated ``score`` fields.
"""
...

View File

@ -15,7 +15,7 @@ from __future__ import annotations
from typing import Literal, Optional
from ..model.doc import Doc
from ..model.doc import Doc, DocList
from ..tool import require_module
from .rerank_function import RerankFunction
from .sentence_transformer_function import SentenceTransformerFunctionBase
@ -33,8 +33,6 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction):
Args:
query (str): Query text for semantic re-ranking. **Required**.
topn (int, optional): Maximum number of documents to return after re-ranking.
Defaults to 10.
rerank_field (Optional[str], optional): Document field name to use as
re-ranking input text. **Required** (e.g., "content", "title", "body").
model_name (str, optional): Cross-encoder model identifier or local path.
@ -56,7 +54,6 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction):
Attributes:
query (str): The query text used for re-ranking.
topn (int): Maximum number of documents to return.
rerank_field (Optional[str]): Field name used for re-ranking input.
model_name (str): The cross-encoder model being used.
model_source (str): The model source ("huggingface" or "modelscope").
@ -113,7 +110,6 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction):
>>>
>>> reranker = SentenceTransformerReRanker(
... query="machine learning algorithms",
... topn=5,
... rerank_field="content"
... )
>>>
@ -127,7 +123,6 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction):
>>> # Using ModelScope for users in China
>>> reranker = SentenceTransformerReRanker(
... query="深度学习",
... topn=10,
... rerank_field="content",
... model_source="modelscope"
... )
@ -135,7 +130,6 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction):
>>> # Using larger model for better quality
>>> reranker = SentenceTransformerReRanker(
... query="neural networks",
... topn=5,
... rerank_field="content",
... model_name="BAAI/bge-reranker-large",
... device="cuda",
@ -143,13 +137,13 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction):
... )
>>> # Direct rerank call (for testing)
>>> query_results = {
... "vector1": [
>>> query_results = [
... [
... Doc(id="1", score=0.9, fields={"content": "Machine learning is..."}),
... Doc(id="2", score=0.8, fields={"content": "Deep learning is..."}),
... ]
... }
>>> reranked = reranker.rerank(query_results)
... ]
>>> reranked = reranker.rerank(query_results, topn=5)
>>> for doc in reranked:
... print(f"ID: {doc.id}, Score: {doc.score:.4f}")
ID: 2, Score: 0.9234
@ -170,7 +164,6 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction):
def __init__(
self,
query: Optional[str] = None,
topn: int = 10,
rerank_field: Optional[str] = None,
model_name: str = "cross-encoder/ms-marco-MiniLM-L6-v2",
model_source: Literal["huggingface", "modelscope"] = "huggingface",
@ -181,7 +174,6 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction):
Args:
query (Optional[str]): Query text for semantic matching. Required.
topn (int): Number of top results to return.
rerank_field (Optional[str]): Document field for re-ranking input.
model_name (str): Cross-encoder model identifier.
model_source (Literal["huggingface", "modelscope"]): Model source.
@ -197,12 +189,13 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction):
)
# Initialize rerank function
RerankFunction.__init__(self, topn=topn, rerank_field=rerank_field)
RerankFunction.__init__(self)
# Validate query
if not query:
raise ValueError("Query is required for DefaultLocalReRanker")
self._query = query
self._rerank_field = rerank_field
self._batch_size = batch_size
# Load and validate cross-encoder model
@ -273,12 +266,17 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction):
"""str: Query text used for semantic re-ranking."""
return self._query
@property
def rerank_field(self) -> Optional[str]:
"""Optional[str]: Field name used as re-ranking input."""
return self._rerank_field
@property
def batch_size(self) -> int:
"""int: Batch size for processing query-document pairs."""
return self._batch_size
def rerank(self, query_results: dict[str, list[Doc]]) -> list[Doc]:
def rerank(self, query_results: list[DocList], topn: int) -> DocList:
"""Re-rank documents using Sentence Transformer cross-encoder model.
Evaluates each query-document pair using the cross-encoder model to compute
@ -286,19 +284,22 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction):
results are returned.
Args:
query_results (dict[str, list[Doc]]): Mapping from vector field names
to lists of retrieved documents. Documents from all fields are
query_results (list[DocList]): Multi-route recall results,
positionally aligned with the queries supplied to
``collection.query()``. Documents from all routes are
deduplicated and re-ranked together.
topn (int): Maximum number of documents to return after re-ranking.
Returns:
list[Doc]: Re-ranked documents (up to ``topn``) with updated ``score``
fields containing relevance scores from the cross-encoder model.
DocList: Re-ranked documents (up to ``topn``) with updated
``score`` fields containing relevance scores from the
cross-encoder model.
Raises:
ValueError: If no valid documents are found or model inference fails.
Note:
- Duplicate documents (same ID) across fields are processed once
- Duplicate documents (same ID) across routes are processed once
- Documents with empty/missing ``rerank_field`` content are skipped
- Returned scores are logits from the cross-encoder model
- Higher scores indicate higher relevance
@ -310,13 +311,13 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction):
... topn=3,
... rerank_field="content"
... )
>>> query_results = {
... "vector1": [
>>> query_results = [
... [
... Doc(id="1", score=0.9, fields={"content": "ML basics"}),
... Doc(id="2", score=0.8, fields={"content": "DL tutorial"}),
... ]
... }
>>> reranked = reranker.rerank(query_results)
... ]
>>> reranked = reranker.rerank(query_results, topn=3)
>>> len(reranked) <= 3
True
"""
@ -328,7 +329,7 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction):
doc_ids: list[str] = []
contents: list[str] = []
for _, query_result in query_results.items():
for query_result in query_results:
for doc in query_result:
doc_id = doc.id
if doc_id in id_to_doc:
@ -373,10 +374,10 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction):
# Sort by score (descending) and take top-k
scored_docs.sort(key=lambda x: x[2], reverse=True)
top_scored_docs = scored_docs[: self.topn]
top_scored_docs = scored_docs[:topn]
# Build result list with updated scores
results: list[Doc] = []
results: DocList = []
for _, doc, score in top_scored_docs:
new_doc = doc._replace(score=score)
results.append(new_doc)

View File

@ -18,11 +18,11 @@ from typing import Optional, Union, overload
from _zvec import _Collection
from ..executor import QueryContext, QueryExecutorFactory
from ..executor import QueryContext, QueryExecutor
from ..extension import ReRanker
from ..typing import Status
from .convert import convert_to_cpp_doc, convert_to_py_doc
from .doc import Doc
from .doc import Doc, DocList
from .param import (
AddColumnOption,
AlterColumnOption,
@ -63,7 +63,7 @@ class Collection:
inst._obj = core_collection
schema = CollectionSchema._from_core(core_collection.Schema())
inst._schema = schema
inst._querier = QueryExecutorFactory.create(schema)
inst._querier = QueryExecutor(schema)
return inst
@property
@ -381,7 +381,7 @@ class Collection:
include_vector: bool = False,
output_fields: Optional[list[str]] = None,
reranker: Optional[ReRanker] = None,
) -> list[Doc]:
) -> DocList:
"""Perform vector similarity search with optional filtering and re-ranking.
At least one `Query` must be provided via `queries`.
@ -403,7 +403,7 @@ class Collection:
Defaults to None.
Returns:
list[Doc]: Top-k matching documents, sorted by relevance score.
DocList: Top-k matching documents, sorted by relevance score.
Examples:
>>> from zvec import Query

View File

@ -20,6 +20,7 @@ from ..common import VectorType
__all__ = [
"Doc",
"DocList",
]
@ -171,3 +172,7 @@ class Doc:
else:
obj.vectors = {}
return obj
#: Type alias for query results: a list of documents returned by a single query route.
DocList = list[Doc]

View File

@ -817,7 +817,7 @@ class VectorIndexParam(IndexParam):
QuantizeType: Vector quantization type (e.g., FP16, INT8).
"""
class _VectorQuery:
class _SearchQuery:
field_name: str
filter: str
include_vector: bool

View File

@ -5622,7 +5622,7 @@ zvec_error_code_t zvec_group_by_vector_query_set_flat_params(
// Reranker Implementation
// =============================================================================
zvec_reranker_t *zvec_reranker_create_rrf(int rank_constant) {
zvec_reranker_t *zvec_create_rrf_reranker(int rank_constant) {
ZVEC_TRY_RETURN_NULL("Failed to create RRF Reranker",
auto *reranker =
new zvec::Reranker::Ptr(
@ -5632,39 +5632,29 @@ zvec_reranker_t *zvec_reranker_create_rrf(int rank_constant) {
return nullptr;
}
zvec_reranker_t *zvec_reranker_create_weighted(const char **fields,
const double *weights,
size_t field_count) {
if ((!fields || !weights) && field_count > 0) {
set_last_error(
"Fields and weights pointers cannot be null when field_count > 0");
zvec_reranker_t *zvec_create_weighted_reranker(const double *weights,
size_t weight_count) {
if (!weights && weight_count > 0) {
set_last_error("Weights pointer cannot be null when weight_count > 0");
return nullptr;
}
ZVEC_TRY_RETURN_NULL(
"Failed to create Weighted Reranker",
std::map<std::string, double> weight_map;
for (size_t i = 0; i < field_count; ++i) {
if (!fields[i]) {
set_last_error("Null field name at index " + std::to_string(i));
return nullptr;
}
weight_map[fields[i]] = weights[i];
}
auto *reranker = new zvec::Reranker::Ptr(
std::make_shared<zvec::WeightedReranker>(weight_map));
std::make_shared<zvec::WeightedReranker>(
std::vector<double>(weights, weights + weight_count)));
return reinterpret_cast<zvec_reranker_t *>(reranker);)
return nullptr;
}
void zvec_reranker_destroy(zvec_reranker_t *reranker) {
void zvec_destroy_reranker(zvec_reranker_t *reranker) {
if (reranker) {
delete reinterpret_cast<zvec::Reranker::Ptr *>(reranker);
}
}
int zvec_reranker_get_rank_constant(const zvec_reranker_t *reranker) {
int zvec_get_reranker_rank_constant(const zvec_reranker_t *reranker) {
if (!reranker) return -1;
auto *ptr = reinterpret_cast<const zvec::Reranker::Ptr *>(reranker);
auto *rrf = dynamic_cast<const zvec::RrfReranker *>(ptr->get());

View File

@ -1538,19 +1538,19 @@ void ZVecPyParams::bind_vector_query(py::module_ &m) {
.def(py::init<>())
.def_readwrite("num_candidates", &SubQuery::num_candidates_)
.def_static(
"from_vector_query",
"from_search_query",
[](const SearchQuery &sq) {
SubQuery sub;
sub.num_candidates_ = sq.topk_;
sub.target_ = sq.target_;
return sub;
},
py::arg("vector_query"),
py::arg("search_query"),
"Create a SubQuery from a single-target search query.");
// _VectorQuery is the historical Python class name; it now wraps the
// _SearchQuery is the Python class name; it wraps the
// single-target SearchQuery so external Python code keeps working unchanged.
py::class_<SearchQuery>(m, "_VectorQuery")
py::class_<SearchQuery>(m, "_SearchQuery")
.def(py::init<>())
// properties
.def_readwrite("topk", &SearchQuery::topk_)
@ -1805,7 +1805,7 @@ void ZVecPyParams::bind_vector_query(py::module_ &m) {
},
[](py::tuple t) {
if (t.size() != 10)
throw std::runtime_error("Invalid pickle data for _VectorQuery");
throw std::runtime_error("Invalid pickle data for _SearchQuery");
SearchQuery obj{};
obj.topk_ = t[0].cast<int>();

View File

@ -13,6 +13,7 @@
// limitations under the License.
#include "python_reranker.h"
#include <stdexcept>
#include <pybind11/functional.h>
#include <pybind11/stl.h>
#include <zvec/db/collection.h>
@ -20,9 +21,40 @@
namespace zvec {
namespace {
inline void reranker_throw_if_error(const Status &status) {
switch (status.code()) {
case StatusCode::OK:
return;
case StatusCode::NOT_FOUND:
throw py::key_error(status.message());
case StatusCode::INVALID_ARGUMENT:
throw py::value_error(status.message());
default:
throw std::runtime_error(status.message());
}
}
inline DocPtrList unwrap_rerank_result(Result<DocPtrList> result) {
if (!result.has_value()) {
reranker_throw_if_error(result.error());
}
return std::move(result).value();
}
} // namespace
void ZVecPyReranker::Initialize(py::module_ &m) {
// Bind Reranker base class (abstract, cannot be instantiated directly)
py::class_<Reranker, Reranker::Ptr>(m, "_Reranker");
py::class_<Reranker, Reranker::Ptr>(m, "_Reranker")
.def(
"rerank",
[](const Reranker &self, const std::vector<DocPtrList> &query_results,
int topn) {
return unwrap_rerank_result(self.rerank(query_results, topn));
},
py::arg("query_results"), py::arg("topn") = 10);
// Bind ScoreBasedReranker intermediate class
py::class_<ScoreBasedReranker, Reranker, std::shared_ptr<ScoreBasedReranker>>(
@ -37,7 +69,7 @@ void ZVecPyReranker::Initialize(py::module_ &m) {
// Bind WeightedReranker
py::class_<WeightedReranker, ScoreBasedReranker,
std::shared_ptr<WeightedReranker>>(m, "_WeightedReranker")
.def(py::init<std::map<std::string, double>>(), py::arg("weights"))
.def(py::init<std::vector<double>>(), py::arg("weights"))
.def_property_readonly("weights", &WeightedReranker::weights);
// Bind CallbackReranker

View File

@ -16,7 +16,6 @@
#include <cstdint>
#include <memory>
#include <mutex>
#include <set>
#include <shared_mutex>
#include <string>
#include <variant>
@ -1706,29 +1705,16 @@ Result<DocPtrList> CollectionImpl::Query(const MultiQuery &query) const {
return DocPtrList();
}
struct PendingQuery {
std::string field_name;
SearchQuery query;
};
// Convert each SubQuery to a SearchQuery and validate.
std::set<std::string> seen_fields;
std::vector<PendingQuery> pending_queries;
pending_queries.reserve(query.queries.size());
std::vector<SearchQuery> search_queries;
std::vector<std::string> field_names;
search_queries.reserve(query.queries.size());
field_names.reserve(query.queries.size());
for (const auto &sub : query.queries) {
const auto &target = sub.target_;
auto [_, inserted] = seen_fields.insert(target.field_name_);
if (!inserted) {
return tl::make_unexpected(Status::InvalidArgument(
"Duplicate field name in multi-query: ", target.field_name_));
}
// Use get_field uniformly; validate_and_sanitize checks type compatibility.
auto *field_schema = schema_->get_field(target.field_name_);
if (!field_schema) {
return tl::make_unexpected(
Status::InvalidArgument("Field not found: ", target.field_name_));
}
SearchQuery sq;
sq.target_ = target;
@ -1740,43 +1726,44 @@ Result<DocPtrList> CollectionImpl::Query(const MultiQuery &query) const {
auto s = sq.validate_and_sanitize(field_schema);
CHECK_RETURN_STATUS_EXPECTED(s);
pending_queries.push_back({target.field_name_, std::move(sq)});
field_names.push_back(target.field_name_);
search_queries.push_back(std::move(sq));
}
std::map<std::string, DocPtrList> query_results;
auto execute_query = [&](PendingQuery &pending) -> Result<DocPtrList> {
// Execute sub-queries.
auto execute_query = [&](SearchQuery &sq) -> Result<DocPtrList> {
auto engine = sqlengine::SQLEngine::create(std::make_shared<Profiler>());
return engine->execute(schema_, std::move(pending.query), segments);
return engine->execute(schema_, std::move(sq), segments);
};
std::vector<Result<DocPtrList>> results(pending_queries.size());
std::vector<Result<DocPtrList>> results(search_queries.size());
// Single-segment queries have no segment-level fanout; multi-segment queries
// already use the query pool per sub-query.
if (segments.size() == 1) {
auto group = GlobalResource::Instance().query_thread_pool()->make_group();
for (size_t i = 0; i < pending_queries.size(); ++i) {
for (size_t i = 0; i < search_queries.size(); ++i) {
group->execute(
[&, i]() { results[i] = execute_query(pending_queries[i]); });
[&, i]() { results[i] = execute_query(search_queries[i]); });
}
group->wait_finish();
} else {
for (size_t i = 0; i < pending_queries.size(); ++i) {
results[i] = execute_query(pending_queries[i]);
for (size_t i = 0; i < search_queries.size(); ++i) {
results[i] = execute_query(search_queries[i]);
}
}
for (size_t i = 0; i < pending_queries.size(); ++i) {
if (!results[i]) {
return tl::make_unexpected(results[i].error());
// Collect results and rerank.
std::vector<DocPtrList> query_results;
query_results.reserve(results.size());
for (auto &result : results) {
if (!result) {
return tl::make_unexpected(result.error());
}
query_results[pending_queries[i].field_name] =
std::move(results[i].value());
query_results.push_back(std::move(result.value()));
}
// Merge and rerank results
query.reranker->bind_schema(schema_);
query.reranker->bind_schema(schema_, field_names);
return query.reranker->rerank(query_results, query.topk);
}

View File

@ -28,16 +28,22 @@ namespace zvec {
// ==================== ScoreBasedReranker ====================
Result<DocPtrList> ScoreBasedReranker::rerank(
const std::map<std::string, DocPtrList> &query_results, int topn) const {
const std::vector<DocPtrList> &query_results, int topn) const {
if (topn <= 0) {
return DocPtrList();
}
std::unordered_map<std::string, double> scores;
std::unordered_map<std::string, Doc::Ptr> id_to_doc;
for (const auto &[field_name, docs] : query_results) {
for (size_t query_index = 0; query_index < query_results.size();
++query_index) {
const auto &docs = query_results[query_index];
for (size_t rank = 0; rank < docs.size(); ++rank) {
const auto &doc = docs[rank];
const std::string &doc_id = doc->pk();
auto rs = rescore(static_cast<double>(doc->score()),
static_cast<int>(rank), field_name);
static_cast<int>(rank), static_cast<int>(query_index));
if (!rs.has_value()) {
return tl::make_unexpected(rs.error());
}
@ -79,18 +85,20 @@ Result<DocPtrList> ScoreBasedReranker::rerank(
// ==================== RrfReranker ====================
Result<double> RrfReranker::rescore(double /*score*/, int rank,
const std::string & /*field_name*/) const {
int /*query_index*/) const {
return 1.0 / (static_cast<double>(rank_constant_) +
static_cast<double>(rank) + 1.0);
}
// ==================== WeightedReranker ====================
WeightedReranker::WeightedReranker(const std::map<std::string, double> &weights)
WeightedReranker::WeightedReranker(const std::vector<double> &weights)
: weights_(weights) {}
void WeightedReranker::bind_schema(CollectionSchema::Ptr schema) {
void WeightedReranker::bind_schema(
CollectionSchema::Ptr schema, const std::vector<std::string> &field_names) {
schema_ = std::move(schema);
field_names_ = field_names;
}
Result<double> WeightedReranker::normalize_score(double score,
@ -122,7 +130,18 @@ Result<double> WeightedReranker::normalize_score(double score,
}
Result<double> WeightedReranker::rescore(double score, int /*rank*/,
const std::string &field_name) const {
int query_index) const {
if (!schema_) {
return tl::make_unexpected(
Status::InvalidArgument("WeightedReranker: schema is null"));
}
if (query_index < 0 ||
static_cast<size_t>(query_index) >= field_names_.size()) {
return tl::make_unexpected(
Status::InvalidArgument("WeightedReranker: query_index out of range: ",
std::to_string(query_index)));
}
const auto &field_name = field_names_[query_index];
const auto *field = schema_->get_field(field_name);
if (!field) {
return tl::make_unexpected(Status::InvalidArgument(
@ -133,9 +152,8 @@ Result<double> WeightedReranker::rescore(double score, int /*rank*/,
return tl::make_unexpected(normalized.error());
}
double weight = 1.0;
auto weight_it = weights_.find(field_name);
if (weight_it != weights_.end()) {
weight = weight_it->second;
if (static_cast<size_t>(query_index) < weights_.size()) {
weight = weights_[query_index];
}
return normalized.value() * weight;
}

View File

@ -1140,8 +1140,8 @@ typedef struct zvec_doc_t zvec_doc_t;
/**
* @brief Reranker structure (opaque pointer)
* Aligned with zvec::Reranker
* Use zvec_reranker_create_rrf() or zvec_reranker_create_weighted() to create
* and zvec_reranker_destroy() to destroy
* Use zvec_create_rrf_reranker() or zvec_create_weighted_reranker() to create
* and zvec_destroy_reranker() to destroy
*/
typedef struct zvec_reranker_t zvec_reranker_t;
typedef struct zvec_collection_schema_t zvec_collection_schema_t;
@ -1961,23 +1961,22 @@ zvec_group_by_vector_query_set_flat_params(
* @return zvec_reranker_t* Pointer to the newly created reranker
*/
ZVEC_EXPORT zvec_reranker_t *ZVEC_CALL
zvec_reranker_create_rrf(int rank_constant);
zvec_create_rrf_reranker(int rank_constant);
/**
* @brief Create a Weighted reranker
* @param fields Array of field names
* @param weights Array of weights corresponding to fields
* @param field_count Number of field/weight entries
* @param weights Array of weights for each query
* @param weight_count Number of weight entries
* @return zvec_reranker_t* Pointer to the newly created reranker
*/
ZVEC_EXPORT zvec_reranker_t *ZVEC_CALL zvec_reranker_create_weighted(
const char **fields, const double *weights, size_t field_count);
ZVEC_EXPORT zvec_reranker_t *ZVEC_CALL
zvec_create_weighted_reranker(const double *weights, size_t weight_count);
/**
* @brief Destroy reranker
* @param reranker Reranker pointer
*/
ZVEC_EXPORT void ZVEC_CALL zvec_reranker_destroy(zvec_reranker_t *reranker);
ZVEC_EXPORT void ZVEC_CALL zvec_destroy_reranker(zvec_reranker_t *reranker);
/**
* @brief Get RRF rank constant (only valid for RRF reranker)
@ -1985,7 +1984,7 @@ ZVEC_EXPORT void ZVEC_CALL zvec_reranker_destroy(zvec_reranker_t *reranker);
* @return int Rank constant, or -1 if not an RRF reranker
*/
ZVEC_EXPORT int ZVEC_CALL
zvec_reranker_get_rank_constant(const zvec_reranker_t *reranker);
zvec_get_reranker_rank_constant(const zvec_reranker_t *reranker);
// -----------------------------------------------------------------------------
// zvec_multi_query_t (Multi Query)
@ -2100,7 +2099,7 @@ ZVEC_EXPORT zvec_error_code_t ZVEC_CALL zvec_multi_query_get_output_fields(
* reranker)
* @param query Multi-vector query pointer
* @param reranker Reranker pointer (remains valid, caller must call
* zvec_reranker_destroy after use)
* zvec_destroy_reranker after use)
* @return zvec_error_code_t Error code
*/
ZVEC_EXPORT zvec_error_code_t ZVEC_CALL zvec_multi_query_set_reranker(

View File

@ -14,9 +14,9 @@
#pragma once
#include <functional>
#include <map>
#include <memory>
#include <string>
#include <vector>
#include <zvec/db/doc.h>
#include <zvec/db/schema.h>
#include <zvec/db/type.h>
@ -32,16 +32,16 @@ class Reranker {
Reranker() = default;
virtual ~Reranker() = default;
virtual void bind_schema(CollectionSchema::Ptr) {}
virtual void bind_schema(CollectionSchema::Ptr /*schema*/,
const std::vector<std::string> & /*field_names*/) {}
//! Re-rank documents from one or more vector queries.
//! \param query_results Mapping from vector field name to list of retrieved
//! documents (sorted by relevance).
//! \param query_results Per-query lists of retrieved documents (sorted by
//! relevance), in the same order as the sub-queries supplied by the caller.
//! \param topn Maximum number of documents to return.
//! \return Re-ranked list of documents (length <= topn), with updated scores.
virtual Result<DocPtrList> rerank(
const std::map<std::string, DocPtrList> &query_results,
int topn = 10) const = 0;
const std::vector<DocPtrList> &query_results, int topn = 10) const = 0;
};
//! Intermediate base for rerankers that compute per-document scores.
@ -51,17 +51,18 @@ class Reranker {
//! Subclasses only need to implement rescore().
class ScoreBasedReranker : public Reranker {
public:
//! Compute the contribution score for a single document.
//! \param score The document's raw relevance score from the vector field.
//! \param rank The document's position (0-based) in the per-field result
//! list. \param field_name The name of the vector field this result came
//! from. \return The score contribution to be accumulated for this document.
virtual Result<double> rescore(double score, int rank,
const std::string &field_name) const = 0;
Result<DocPtrList> rerank(const std::vector<DocPtrList> &query_results,
int topn = 10) const override;
Result<DocPtrList> rerank(
const std::map<std::string, DocPtrList> &query_results,
int topn = 10) const override;
private:
//! Compute the contribution score for a single document.
//! \param score The document's raw relevance score from the vector query.
//! \param rank The document's position (0-based) in the per-query result
//! list. \param query_index The index (0-based) of the sub-query this result
//! came from. \return The score contribution to be accumulated for this
//! document.
virtual Result<double> rescore(double score, int rank,
int query_index) const = 0;
};
//! Re-ranker using Reciprocal Rank Fusion (RRF) for multi-vector search.
@ -79,10 +80,10 @@ class RrfReranker : public ScoreBasedReranker {
return rank_constant_;
}
Result<double> rescore(double score, int rank,
const std::string &field_name) const override;
private:
Result<double> rescore(double score, int rank,
int query_index) const override;
int rank_constant_;
};
@ -91,24 +92,30 @@ class RrfReranker : public ScoreBasedReranker {
//! Each vector field's relevance score is normalized based on its own metric
//! type, then scaled by a user-provided weight. Final scores are summed across
//! fields. Supported metrics: L2, IP, COSINE.
//!
//! @note NOT thread-safe. The bind_schema() and rerank() calls share mutable
//! state. Each concurrent query must use its own WeightedReranker instance or
//! serialize access externally.
class WeightedReranker : public ScoreBasedReranker {
public:
explicit WeightedReranker(const std::map<std::string, double> &weights = {});
explicit WeightedReranker(const std::vector<double> &weights = {});
void bind_schema(CollectionSchema::Ptr schema) override;
void bind_schema(CollectionSchema::Ptr schema,
const std::vector<std::string> &field_names) override;
const std::map<std::string, double> &weights() const {
const std::vector<double> &weights() const {
return weights_;
}
Result<double> rescore(double score, int rank,
const std::string &field_name) const override;
private:
Result<double> rescore(double score, int rank,
int query_index) const override;
static Result<double> normalize_score(double score, const FieldSchema &field);
CollectionSchema::Ptr schema_;
std::map<std::string, double> weights_;
std::vector<std::string> field_names_;
std::vector<double> weights_;
};
//! Callback-based re-ranker for cross-language bridging.
@ -118,13 +125,16 @@ class WeightedReranker : public ScoreBasedReranker {
class CallbackReranker : public Reranker {
public:
using Callback =
std::function<DocPtrList(const std::map<std::string, DocPtrList> &, int)>;
std::function<DocPtrList(const std::vector<DocPtrList> &, int)>;
explicit CallbackReranker(Callback fn) : callback_(std::move(fn)) {}
Result<DocPtrList> rerank(
const std::map<std::string, DocPtrList> &query_results,
int topn = 10) const override {
Result<DocPtrList> rerank(const std::vector<DocPtrList> &query_results,
int topn = 10) const override {
if (!callback_) {
return tl::make_unexpected(
Status::InvalidArgument("CallbackReranker: callback is empty"));
}
return callback_(query_results, topn);
}

View File

@ -4370,41 +4370,40 @@ void test_reranker_functions(void) {
TEST_START();
// Test 1: Create RRF reranker
zvec_reranker_t *rrf = zvec_reranker_create_rrf(60);
zvec_reranker_t *rrf = zvec_create_rrf_reranker(60);
TEST_ASSERT(rrf != NULL);
if (rrf) {
TEST_ASSERT(zvec_reranker_get_rank_constant(rrf) == 60);
zvec_reranker_destroy(rrf);
TEST_ASSERT(zvec_get_reranker_rank_constant(rrf) == 60);
zvec_destroy_reranker(rrf);
}
// Test 2: Create RRF reranker with different rank constant
zvec_reranker_t *rrf2 = zvec_reranker_create_rrf(100);
zvec_reranker_t *rrf2 = zvec_create_rrf_reranker(100);
TEST_ASSERT(rrf2 != NULL);
if (rrf2) {
TEST_ASSERT(zvec_reranker_get_rank_constant(rrf2) == 100);
zvec_reranker_destroy(rrf2);
TEST_ASSERT(zvec_get_reranker_rank_constant(rrf2) == 100);
zvec_destroy_reranker(rrf2);
}
// Test 3: Create Weighted reranker
const char *fields[] = {"embedding1", "embedding2"};
double weights[] = {0.7, 0.3};
zvec_reranker_t *weighted = zvec_reranker_create_weighted(fields, weights, 2);
zvec_reranker_t *weighted = zvec_create_weighted_reranker(weights, 2);
TEST_ASSERT(weighted != NULL);
if (weighted) {
TEST_ASSERT(zvec_reranker_get_rank_constant(weighted) == -1);
zvec_reranker_destroy(weighted);
TEST_ASSERT(zvec_get_reranker_rank_constant(weighted) == -1);
zvec_destroy_reranker(weighted);
}
// Test 4: Create Weighted reranker with no fields
zvec_reranker_t *weighted2 = zvec_reranker_create_weighted(NULL, NULL, 0);
// Test 4: Create Weighted reranker with no weights
zvec_reranker_t *weighted2 = zvec_create_weighted_reranker(NULL, 0);
TEST_ASSERT(weighted2 != NULL);
if (weighted2) {
zvec_reranker_destroy(weighted2);
zvec_destroy_reranker(weighted2);
}
// Test 5: NULL reranker operations
TEST_ASSERT(zvec_reranker_get_rank_constant(NULL) == -1);
zvec_reranker_destroy(NULL); // Should not crash
TEST_ASSERT(zvec_get_reranker_rank_constant(NULL) == -1);
zvec_destroy_reranker(NULL); // Should not crash
TEST_END();
}
@ -4528,14 +4527,14 @@ void test_multi_vector_query_with_rrf_reranker(void) {
multi_query_fixture_t f;
TEST_ASSERT(setup_multi_query_fixture(&f, "zvec_test_mq_rrf", "mq_rrf"));
zvec_reranker_t *rrf = zvec_reranker_create_rrf(60);
zvec_reranker_t *rrf = zvec_create_rrf_reranker(60);
TEST_ASSERT(rrf != NULL);
int count = execute_multi_query_with_reranker(&f, rrf, 3, 3);
TEST_ASSERT(count > 0);
TEST_ASSERT(count <= 3);
zvec_reranker_destroy(rrf);
zvec_destroy_reranker(rrf);
// MultiQuery property setters/getters
zvec_multi_query_t *mvq2 = zvec_multi_query_create();
@ -4589,16 +4588,15 @@ void test_multi_vector_query_with_weighted_reranker(void) {
TEST_ASSERT(
setup_multi_query_fixture(&f, "zvec_test_mq_weighted", "mq_weighted"));
const char *fields[] = {"embedding1", "embedding2"};
double weights[] = {0.7, 0.3};
zvec_reranker_t *weighted = zvec_reranker_create_weighted(fields, weights, 2);
zvec_reranker_t *weighted = zvec_create_weighted_reranker(weights, 2);
TEST_ASSERT(weighted != NULL);
int count = execute_multi_query_with_reranker(&f, weighted, 3, 3);
TEST_ASSERT(count > 0);
TEST_ASSERT(count <= 3);
zvec_reranker_destroy(weighted);
zvec_destroy_reranker(weighted);
teardown_multi_query_fixture(&f);
TEST_END();

View File

@ -3804,30 +3804,6 @@ TEST_F(CollectionTest, Feature_MultiQuery_Validate) {
EXPECT_EQ(result.error().code(), StatusCode::INVALID_ARGUMENT);
}
// Test 4: Duplicate field names should fail
{
MultiQuery mvq;
mvq.topk = 10;
mvq.reranker = std::make_shared<RrfReranker>(60);
SubQuery vq1;
vq1.num_candidates_ = 10;
vq1.target_.field_name_ = "dense_fp32";
std::get<VectorClause>(vq1.target_.clause_)
.query_vector_.assign(128 * sizeof(float), '\0');
mvq.queries.push_back(vq1);
SubQuery vq2;
vq2.num_candidates_ = 10;
vq2.target_.field_name_ = "dense_fp32";
std::get<VectorClause>(vq2.target_.clause_)
.query_vector_.assign(128 * sizeof(float), '\0');
mvq.queries.push_back(vq2);
auto result = collection->Query(mvq);
ASSERT_FALSE(result.has_value());
EXPECT_EQ(result.error().code(), StatusCode::INVALID_ARGUMENT);
}
}
TEST_F(CollectionTest, Feature_MultiQuery_SingleFieldWithReranker) {
@ -3936,9 +3912,8 @@ TEST_F(CollectionTest, Feature_MultiQuery_MultiFieldWeighted) {
MultiQuery mvq;
mvq.topk = 10;
std::map<std::string, double> weights = {{"dense_fp32", 0.7},
{"sparse_fp32", 0.3}};
mvq.reranker = std::make_shared<WeightedReranker>(weights);
mvq.reranker =
std::make_shared<WeightedReranker>(std::vector<double>{0.7, 0.3});
// Query dense_fp32 field
{
@ -4090,11 +4065,11 @@ TEST_F(CollectionTest, Feature_MultiQuery_CallbackReranker) {
// Use CallbackReranker with a lambda that merges and sorts by score
bool callback_invoked = false;
auto callback_fn = [&callback_invoked](
const std::map<std::string, DocPtrList> &query_results,
const std::vector<DocPtrList> &query_results,
int topn) -> DocPtrList {
callback_invoked = true;
DocPtrList all_docs;
for (const auto &[_, docs] : query_results) {
for (const auto &docs : query_results) {
for (const auto &doc : docs) {
all_docs.push_back(doc);
}

View File

@ -14,8 +14,8 @@
#define _USE_MATH_DEFINES
#include <cmath>
#include <map>
#include <memory>
#include <set>
#include <string>
#include <vector>
#include <gtest/gtest.h>
@ -55,11 +55,11 @@ TEST(RrfRerankerTest, BasicRRF) {
RrfReranker reranker(/*rank_constant=*/60);
// Two vector fields, each returning 3 documents with some overlap
std::map<std::string, DocPtrList> query_results;
query_results["vec1"] = {MakeDoc("a", 0.9f), MakeDoc("b", 0.8f),
MakeDoc("c", 0.7f)};
query_results["vec2"] = {MakeDoc("b", 0.95f), MakeDoc("a", 0.85f),
MakeDoc("d", 0.75f)};
std::vector<DocPtrList> query_results;
query_results.push_back(
{MakeDoc("a", 0.9f), MakeDoc("b", 0.8f), MakeDoc("c", 0.7f)});
query_results.push_back(
{MakeDoc("b", 0.95f), MakeDoc("a", 0.85f), MakeDoc("d", 0.75f)});
auto result = reranker.rerank(query_results, /*topn=*/10);
ASSERT_TRUE(result.has_value());
@ -72,9 +72,9 @@ TEST(RrfRerankerTest, BasicRRF) {
// So a and b should have equal scores and be at the top
ASSERT_GE(results.size(), 3u);
// "a" and "b" should have the highest RRF scores
EXPECT_EQ(results[0]->pk(), "a");
EXPECT_EQ(results[1]->pk(), "b");
// "a" and "b" should have the highest RRF scores (equal, order unspecified)
std::set<std::string> top2{results[0]->pk(), results[1]->pk()};
EXPECT_EQ(top2, (std::set<std::string>{"a", "b"}));
// Verify scores are close (a and b have same RRF score)
EXPECT_NEAR(results[0]->score(), results[1]->score(), 1e-10);
}
@ -82,9 +82,9 @@ TEST(RrfRerankerTest, BasicRRF) {
TEST(RrfRerankerTest, Topn) {
RrfReranker reranker(/*rank_constant=*/60);
std::map<std::string, DocPtrList> query_results;
query_results["vec1"] = {MakeDoc("a", 0.9f), MakeDoc("b", 0.8f),
MakeDoc("c", 0.7f)};
std::vector<DocPtrList> query_results;
query_results.push_back(
{MakeDoc("a", 0.9f), MakeDoc("b", 0.8f), MakeDoc("c", 0.7f)});
auto result = reranker.rerank(query_results, /*topn=*/2);
ASSERT_TRUE(result.has_value());
@ -94,8 +94,8 @@ TEST(RrfRerankerTest, Topn) {
TEST(RrfRerankerTest, SingleField) {
RrfReranker reranker(/*rank_constant=*/60);
std::map<std::string, DocPtrList> query_results;
query_results["vec1"] = {MakeDoc("a", 0.9f), MakeDoc("b", 0.8f)};
std::vector<DocPtrList> query_results;
query_results.push_back({MakeDoc("a", 0.9f), MakeDoc("b", 0.8f)});
auto result = reranker.rerank(query_results);
ASSERT_TRUE(result.has_value());
@ -108,7 +108,7 @@ TEST(RrfRerankerTest, SingleField) {
TEST(RrfRerankerTest, EmptyResults) {
RrfReranker reranker(/*rank_constant=*/60);
std::map<std::string, DocPtrList> query_results;
std::vector<DocPtrList> query_results;
auto result = reranker.rerank(query_results);
ASSERT_TRUE(result.has_value());
EXPECT_TRUE(result.value().empty());
@ -119,12 +119,12 @@ TEST(RrfRerankerTest, EmptyResults) {
TEST(WeightedRerankerTest, BasicWeighted) {
auto schema =
MakeSchema({{"vec1", MetricType::L2}, {"vec2", MetricType::L2}});
WeightedReranker reranker({{"vec1", 0.7}, {"vec2", 0.3}});
reranker.bind_schema(schema);
WeightedReranker reranker({0.7, 0.3});
reranker.bind_schema(schema, {"vec1", "vec2"});
std::map<std::string, DocPtrList> query_results;
query_results["vec1"] = {MakeDoc("a", 0.5f), MakeDoc("b", 0.3f)};
query_results["vec2"] = {MakeDoc("a", 0.8f), MakeDoc("c", 0.6f)};
std::vector<DocPtrList> query_results;
query_results.push_back({MakeDoc("a", 0.5f), MakeDoc("b", 0.3f)});
query_results.push_back({MakeDoc("a", 0.8f), MakeDoc("c", 0.6f)});
auto result = reranker.rerank(query_results);
ASSERT_TRUE(result.has_value());
@ -137,12 +137,12 @@ TEST(WeightedRerankerTest, BasicWeighted) {
TEST(WeightedRerankerTest, MixedMetrics) {
auto schema =
MakeSchema({{"vec1", MetricType::L2}, {"vec2", MetricType::COSINE}});
WeightedReranker reranker({{"vec1", 0.5}, {"vec2", 0.5}});
reranker.bind_schema(schema);
WeightedReranker reranker({0.5, 0.5});
reranker.bind_schema(schema, {"vec1", "vec2"});
std::map<std::string, DocPtrList> query_results;
query_results["vec1"] = {MakeDoc("a", 0.5f)};
query_results["vec2"] = {MakeDoc("a", 0.4f)};
std::vector<DocPtrList> query_results;
query_results.push_back({MakeDoc("a", 0.5f)});
query_results.push_back({MakeDoc("a", 0.4f)});
auto result = reranker.rerank(query_results);
ASSERT_TRUE(result.has_value());
@ -161,12 +161,12 @@ TEST(WeightedRerankerTest, MixedMetrics) {
TEST(WeightedRerankerTest, MissingMetricError) {
auto schema = MakeSchema({{"vec1", MetricType::L2}});
WeightedReranker reranker;
reranker.bind_schema(schema);
std::map<std::string, DocPtrList> query_results;
query_results["vec1"] = {MakeDoc("a", 0.5f)};
query_results["vec2"] = {MakeDoc("b", 0.3f)};
// Binding a field that is absent from the schema should fail at rerank time.
reranker.bind_schema(schema, {"vec1", "vec2"});
std::vector<DocPtrList> query_results;
query_results.push_back({MakeDoc("a", 0.5f)});
query_results.push_back({MakeDoc("b", 0.3f)});
auto result = reranker.rerank(query_results);
ASSERT_FALSE(result.has_value());
}
@ -174,10 +174,10 @@ TEST(WeightedRerankerTest, MissingMetricError) {
TEST(WeightedRerankerTest, NormalizeL2) {
auto schema = MakeSchema({{"vec1", MetricType::L2}});
WeightedReranker reranker;
reranker.bind_schema(schema);
reranker.bind_schema(schema, {"vec1"});
std::map<std::string, DocPtrList> query_results;
query_results["vec1"] = {MakeDoc("a", 0.0f), MakeDoc("b", 1.0f)};
std::vector<DocPtrList> query_results;
query_results.push_back({MakeDoc("a", 0.0f), MakeDoc("b", 1.0f)});
auto result = reranker.rerank(query_results);
ASSERT_TRUE(result.has_value());
@ -193,10 +193,10 @@ TEST(WeightedRerankerTest, NormalizeL2) {
TEST(WeightedRerankerTest, NormalizeIP) {
auto schema = MakeSchema({{"vec1", MetricType::IP}});
WeightedReranker reranker;
reranker.bind_schema(schema);
reranker.bind_schema(schema, {"vec1"});
std::map<std::string, DocPtrList> query_results;
query_results["vec1"] = {MakeDoc("a", 0.0f), MakeDoc("b", 1.0f)};
std::vector<DocPtrList> query_results;
query_results.push_back({MakeDoc("a", 0.0f), MakeDoc("b", 1.0f)});
auto result = reranker.rerank(query_results);
ASSERT_TRUE(result.has_value());
@ -211,11 +211,11 @@ TEST(WeightedRerankerTest, NormalizeIP) {
TEST(WeightedRerankerTest, NormalizeCosine) {
auto schema = MakeSchema({{"vec1", MetricType::COSINE}});
WeightedReranker reranker;
reranker.bind_schema(schema);
reranker.bind_schema(schema, {"vec1"});
std::map<std::string, DocPtrList> query_results;
query_results["vec1"] = {MakeDoc("a", 0.0f), MakeDoc("b", 1.0f),
MakeDoc("c", 2.0f)};
std::vector<DocPtrList> query_results;
query_results.push_back(
{MakeDoc("a", 0.0f), MakeDoc("b", 1.0f), MakeDoc("c", 2.0f)});
auto result = reranker.rerank(query_results);
ASSERT_TRUE(result.has_value());
@ -230,11 +230,11 @@ TEST(WeightedRerankerTest, NormalizeCosine) {
TEST(WeightedRerankerTest, Topn) {
auto schema = MakeSchema({{"vec1", MetricType::L2}});
WeightedReranker reranker;
reranker.bind_schema(schema);
reranker.bind_schema(schema, {"vec1"});
std::map<std::string, DocPtrList> query_results;
query_results["vec1"] = {MakeDoc("a", 0.1f), MakeDoc("b", 0.2f),
MakeDoc("c", 0.3f)};
std::vector<DocPtrList> query_results;
query_results.push_back(
{MakeDoc("a", 0.1f), MakeDoc("b", 0.2f), MakeDoc("c", 0.3f)});
auto result = reranker.rerank(query_results, /*topn=*/2);
ASSERT_TRUE(result.has_value());
@ -248,10 +248,9 @@ TEST(CallbackRerankerTest, BasicCallback) {
// Simple callback that returns docs sorted by score descending, limited to
// topn
CallbackReranker::Callback cb =
[](const std::map<std::string, DocPtrList> &query_results,
int topn) -> DocPtrList {
[](const std::vector<DocPtrList> &query_results, int topn) -> DocPtrList {
DocPtrList all_docs;
for (const auto &[_, docs] : query_results) {
for (const auto &docs : query_results) {
for (const auto &doc : docs) {
all_docs.push_back(doc);
}
@ -268,9 +267,9 @@ TEST(CallbackRerankerTest, BasicCallback) {
CallbackReranker reranker(cb);
std::map<std::string, DocPtrList> query_results;
query_results["vec1"] = {MakeDoc("a", 0.5f), MakeDoc("b", 0.9f)};
query_results["vec2"] = {MakeDoc("c", 0.7f)};
std::vector<DocPtrList> query_results;
query_results.push_back({MakeDoc("a", 0.5f), MakeDoc("b", 0.9f)});
query_results.push_back({MakeDoc("c", 0.7f)});
auto result = reranker.rerank(query_results, /*topn=*/10);
ASSERT_TRUE(result.has_value());