From 9ea43d28945c5efa76eddcc08a0152bbb65141c1 Mon Sep 17 00:00:00 2001 From: Gioia Zheng Date: Mon, 6 Jul 2026 14:00:08 +0200 Subject: [PATCH] fix: support HF datasets in dense exact search --- beir/retrieval/search/dense/exact_search.py | 78 +++++++++++++------ .../dense/test_exact_search_hf_dataset.py | 71 +++++++++++++++++ 2 files changed, 125 insertions(+), 24 deletions(-) create mode 100644 tests/retrieval/search/dense/test_exact_search_hf_dataset.py diff --git a/beir/retrieval/search/dense/exact_search.py b/beir/retrieval/search/dense/exact_search.py index 070c786..98b680e 100644 --- a/beir/retrieval/search/dense/exact_search.py +++ b/beir/retrieval/search/dense/exact_search.py @@ -20,6 +20,46 @@ logger = logging.getLogger(__name__) +def _has_columns(data, columns: set[str]) -> bool: + column_names = getattr(data, "column_names", None) + return column_names is not None and columns.issubset(column_names) + + +def _prepare_queries(queries): + if _has_columns(queries, {"id", "text"}): + return [str(query_id) for query_id in queries["id"]], list(queries["text"]) + + query_ids = list(queries.keys()) + return query_ids, [queries[qid] for qid in queries] + + +def _prepare_corpus(corpus): + if _has_columns(corpus, {"id", "text"}): + def dataset_text_length(index): + row = corpus[index] + return len((row.get("title") or "") + (row.get("text") or "")) + + row_indices = sorted(range(len(corpus)), key=dataset_text_length, reverse=True) + corpus_ids = [str(corpus[index]["id"]) for index in row_indices] + + def dataset_batch(start, end): + return [corpus[index] for index in row_indices[start:end]] + + return corpus_ids, dataset_batch + + corpus_ids = sorted( + corpus, + key=lambda k: len(corpus[k].get("title", "") + corpus[k].get("text", "")), + reverse=True, + ) + corpus_docs = [corpus[cid] for cid in corpus_ids] + + def dict_batch(start, end): + return corpus_docs[start:end] + + return corpus_ids, dict_batch + + # DenseRetrievalExactSearch is parent class for any dense model that can be used for retrieval # Abstract class is BaseSearch class DenseRetrievalExactSearch(BaseSearch): @@ -55,11 +95,10 @@ def search( ) logger.info("Encoding Queries...") - query_ids = list(queries.keys()) + query_ids, query_texts = _prepare_queries(queries) self.results = {qid: {} for qid in query_ids} - queries = [queries[qid] for qid in queries] query_embeddings = self.model.encode_queries( - queries, + query_texts, batch_size=self.batch_size, show_progress_bar=self.show_progress_bar, convert_to_tensor=self.convert_to_tensor, @@ -67,26 +106,22 @@ def search( logger.info("Sorting Corpus by document length (Longest first)...") - corpus_ids = sorted( - corpus, - key=lambda k: len(corpus[k].get("title", "") + corpus[k].get("text", "")), - reverse=True, - ) - corpus = [corpus[cid] for cid in corpus_ids] + corpus_ids, get_corpus_batch = _prepare_corpus(corpus) logger.info("Encoding Corpus in batches... Warning: This might take a while!") logger.info(f"Scoring Function: {self.score_function_desc[score_function]} ({score_function})") - itr = range(0, len(corpus), self.corpus_chunk_size) + itr = range(0, len(corpus_ids), self.corpus_chunk_size) result_heaps = {qid: [] for qid in query_ids} # Keep only the top-k docs for each query for batch_num, corpus_start_idx in enumerate(itr): logger.info(f"Encoding Batch {batch_num + 1}/{len(itr)}...") - corpus_end_idx = min(corpus_start_idx + self.corpus_chunk_size, len(corpus)) + corpus_end_idx = min(corpus_start_idx + self.corpus_chunk_size, len(corpus_ids)) # Encode chunk of corpus + sub_corpus = get_corpus_batch(corpus_start_idx, corpus_end_idx) sub_corpus_embeddings = self.model.encode_corpus( - corpus[corpus_start_idx:corpus_end_idx], + sub_corpus, batch_size=self.batch_size, show_progress_bar=self.show_progress_bar, convert_to_tensor=self.convert_to_tensor, @@ -136,14 +171,13 @@ def encode( **kwargs, ): logger.info("Encoding Queries...") - query_ids = list(queries.keys()) + query_ids, query_texts = _prepare_queries(queries) self.results = {qid: {} for qid in query_ids} - queries = [queries[qid] for qid in queries] query_embeddings_file = os.path.join(encode_output_path, query_filename) if not os.path.exists(query_embeddings_file) or overwrite: query_embeddings = self.model.encode_queries( - queries, + query_texts, batch_size=self.batch_size, show_progress_bar=self.show_progress_bar, convert_to_tensor=self.convert_to_tensor, @@ -155,27 +189,23 @@ def encode( logger.info("Sorting Corpus by document length (Longest first)...") - corpus_ids = sorted( - corpus, - key=lambda k: len(corpus[k].get("title", "") + corpus[k].get("text", "")), - reverse=True, - ) - corpus = [corpus[cid] for cid in corpus_ids] + corpus_ids, get_corpus_batch = _prepare_corpus(corpus) logger.info("Encoding Corpus in batches... Warning: This might take a while!") - itr = range(0, len(corpus), self.corpus_chunk_size) + itr = range(0, len(corpus_ids), self.corpus_chunk_size) for batch_num, corpus_start_idx in enumerate(itr): batch_corpus_filename = corpus_filename.replace("*", str(batch_num)) corpus_embeddings_file = os.path.join(encode_output_path, batch_corpus_filename) if not os.path.exists(corpus_embeddings_file) or overwrite: logger.info(f"Encoding Batch {batch_num + 1}/{len(itr)}...") - corpus_end_idx = min(corpus_start_idx + self.corpus_chunk_size, len(corpus)) + corpus_end_idx = min(corpus_start_idx + self.corpus_chunk_size, len(corpus_ids)) # Encode chunk of corpus + sub_corpus = get_corpus_batch(corpus_start_idx, corpus_end_idx) sub_corpus_embeddings = self.model.encode_corpus( - corpus[corpus_start_idx:corpus_end_idx], + sub_corpus, batch_size=self.batch_size, show_progress_bar=self.show_progress_bar, convert_to_tensor=self.convert_to_tensor, diff --git a/tests/retrieval/search/dense/test_exact_search_hf_dataset.py b/tests/retrieval/search/dense/test_exact_search_hf_dataset.py new file mode 100644 index 0000000..c0f392f --- /dev/null +++ b/tests/retrieval/search/dense/test_exact_search_hf_dataset.py @@ -0,0 +1,71 @@ +from __future__ import annotations + +import torch +from datasets import Dataset + +from beir.retrieval.search.dense import DenseRetrievalExactSearch +from beir.retrieval.search.dense.util import pickle_load + + +class DatasetModel: + def encode_queries(self, queries, **kwargs): + assert queries == ["alpha", "gamma"] + return torch.tensor([[1.0, 0.0], [0.0, 1.0]]) + + def encode_corpus(self, corpus, **kwargs): + vectors = [] + for doc in corpus: + if doc["text"] == "alpha": + vectors.append([1.0, 0.0]) + elif doc["text"] == "gamma": + vectors.append([0.0, 1.0]) + else: + vectors.append([0.0, 0.0]) + return torch.tensor(vectors) + + +def _hf_inputs(): + corpus = Dataset.from_dict( + { + "id": ["d1", "d2", "d3"], + "title": ["", "", ""], + "text": ["alpha", "beta", "gamma"], + } + ) + queries = Dataset.from_dict({"id": ["q1", "q2"], "text": ["alpha", "gamma"]}) + return corpus, queries + + +def test_exact_search_accepts_hf_dataset_inputs(): + corpus, queries = _hf_inputs() + retriever = DenseRetrievalExactSearch(DatasetModel(), show_progress_bar=False) + + results = retriever.search( + corpus=corpus, + queries=queries, + top_k=1, + score_function="dot", + ) + + assert results == {"q1": {"d1": 1.0}, "q2": {"d3": 1.0}} + + +def test_exact_search_encode_accepts_hf_dataset_inputs(tmp_path): + corpus, queries = _hf_inputs() + retriever = DenseRetrievalExactSearch( + DatasetModel(), corpus_chunk_size=2, show_progress_bar=False + ) + + retriever.encode( + corpus=corpus, + queries=queries, + encode_output_path=str(tmp_path), + overwrite=True, + ) + + _, query_ids = pickle_load(tmp_path / "queries.pkl") + _, corpus_ids_0 = pickle_load(tmp_path / "corpus.0.pkl") + _, corpus_ids_1 = pickle_load(tmp_path / "corpus.1.pkl") + + assert query_ids == ["q1", "q2"] + assert set(corpus_ids_0 + corpus_ids_1) == {"d1", "d2", "d3"}