Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
78 changes: 54 additions & 24 deletions beir/retrieval/search/dense/exact_search.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -55,38 +95,33 @@ 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,
)

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,
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand Down
71 changes: 71 additions & 0 deletions tests/retrieval/search/dense/test_exact_search_hf_dataset.py
Original file line number Diff line number Diff line change
@@ -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"}