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:
parent
f562bdd636
commit
c46efe1241
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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")),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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")),
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -40,7 +40,7 @@ from zvec import (
|
|||
VectorSchema,
|
||||
)
|
||||
|
||||
from _zvec.param import _VectorQuery
|
||||
from _zvec.param import _SearchQuery
|
||||
|
||||
# ----------------------------
|
||||
# Invert Index Param Test Case
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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]: ...
|
||||
|
|
|
|||
|
|
@ -16,11 +16,9 @@ from __future__ import annotations
|
|||
from .query_executor import (
|
||||
QueryContext,
|
||||
QueryExecutor,
|
||||
QueryExecutorFactory,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"QueryContext",
|
||||
"QueryExecutor",
|
||||
"QueryExecutorFactory",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
"""
|
||||
...
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
|
|
|
|||
|
|
@ -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>();
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
|
|
|
|||
Loading…
Reference in New Issue