diff --git a/redisvl/utils/rerank/hf_cross_encoder.py b/redisvl/utils/rerank/hf_cross_encoder.py index fd6b3325..d28f8b94 100644 --- a/redisvl/utils/rerank/hf_cross_encoder.py +++ b/redisvl/utils/rerank/hf_cross_encoder.py @@ -120,8 +120,8 @@ def rank( scores = [float(score) for score in scores] docs_with_scores = list(zip(doc_subset, scores)) docs_with_scores.sort(key=lambda x: x[1], reverse=True) - reranked_docs = [doc for doc, _ in docs_with_scores[:limit]] - scores = scores[:limit] + reranked_docs_tuple, scores_tuple = zip(*docs_with_scores[:limit]) + reranked_docs, scores = list(reranked_docs_tuple), list(scores_tuple) if return_score: return reranked_docs, scores # type: ignore diff --git a/tests/integration/test_cross_encoder_reranker.py b/tests/integration/test_cross_encoder_reranker.py index a4311544..528b264e 100644 --- a/tests/integration/test_cross_encoder_reranker.py +++ b/tests/integration/test_cross_encoder_reranker.py @@ -33,6 +33,23 @@ async def test_async_rank_documents(reranker): assert all(isinstance(score, float) for score in scores) +def test_rank_scores_align_with_reranked_docs(reranker): + # https://github.com/redis/redis-vl-python/issues/693 + docs = [ + "irrelevant document", + "highly relevant document", + "somewhat relevant document", + ] + query = "relevant document" + + reranked_docs, scores = reranker.rank(query, docs, limit=len(docs)) + + # reranked_docs is sorted by score descending, so scores must be too, + # this catches the case where scores are taken from the original + # unsorted list instead of the sorted one. + assert list(scores) == sorted(scores, reverse=True) + + def test_bad_input(reranker): with pytest.raises(ValueError): reranker.rank("", []) # Empty query