From 98f46bd77694d14e592d665fcd4bfbb6f1a1b493 Mon Sep 17 00:00:00 2001 From: Gioia Zheng Date: Wed, 1 Jul 2026 10:45:16 +0200 Subject: [PATCH] fix: handle single-query exact search --- beir/retrieval/search/dense/exact_search.py | 2 +- .../search/dense/test_exact_search.py | 31 +++++++++++++++++++ 2 files changed, 32 insertions(+), 1 deletion(-) create mode 100644 tests/retrieval/search/dense/test_exact_search.py diff --git a/beir/retrieval/search/dense/exact_search.py b/beir/retrieval/search/dense/exact_search.py index 070c786..e8f632e 100644 --- a/beir/retrieval/search/dense/exact_search.py +++ b/beir/retrieval/search/dense/exact_search.py @@ -99,7 +99,7 @@ def search( # Get top-k values cos_scores_top_k_values, cos_scores_top_k_idx = torch.topk( cos_scores, - min(top_k + 1, len(cos_scores[1])), + min(top_k + 1, cos_scores.shape[1]), dim=1, largest=True, sorted=return_sorted, diff --git a/tests/retrieval/search/dense/test_exact_search.py b/tests/retrieval/search/dense/test_exact_search.py new file mode 100644 index 0000000..91c7a00 --- /dev/null +++ b/tests/retrieval/search/dense/test_exact_search.py @@ -0,0 +1,31 @@ +from __future__ import annotations + +import torch + +from beir.retrieval.search.dense import DenseRetrievalExactSearch + + +class DummyModel: + def encode_queries(self, queries, **kwargs): + assert queries == ["alpha"] + return torch.tensor([[1.0, 0.0]]) + + def encode_corpus(self, corpus, **kwargs): + assert [doc["text"] for doc in corpus] == ["alpha", "beta"] + return torch.tensor([[1.0, 0.0], [0.0, 1.0]]) + + +def test_exact_search_handles_single_query(): + retriever = DenseRetrievalExactSearch(DummyModel(), show_progress_bar=False) + + results = retriever.search( + corpus={ + "d1": {"title": "", "text": "alpha"}, + "d2": {"title": "", "text": "beta"}, + }, + queries={"q1": "alpha"}, + top_k=1, + score_function="dot", + ) + + assert results == {"q1": {"d1": 1.0}}