From 5265a8b241d342fe499f6747d896639c697b1e4d Mon Sep 17 00:00:00 2001 From: maximedb Date: Wed, 9 Feb 2022 18:52:57 +0000 Subject: [PATCH 1/8] rm batch size step in range to pass over all q. --- beir/retrieval/search/sparse/sparse_search.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/beir/retrieval/search/sparse/sparse_search.py b/beir/retrieval/search/sparse/sparse_search.py index c81fb469..2d665d67 100644 --- a/beir/retrieval/search/sparse/sparse_search.py +++ b/beir/retrieval/search/sparse/sparse_search.py @@ -25,7 +25,7 @@ def search(self, self.sparse_matrix = self.model.encode_corpus(documents, batch_size=self.batch_size) logging.info("Starting to Retrieve...") - for start_idx in trange(0, len(queries), self.batch_size, desc='query'): + for start_idx in trange(0, len(queries), desc='query'): qid = query_ids[start_idx] query_tokens = self.model.encode_query(queries[qid]) #Get the candidate passages From c1b5903e3d369088ca7757792f79fc4ca345bc03 Mon Sep 17 00:00:00 2001 From: maximedb Date: Wed, 9 Feb 2022 18:53:49 +0000 Subject: [PATCH 2/8] to be consistent with dense search --- beir/retrieval/search/sparse/sparse_search.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/beir/retrieval/search/sparse/sparse_search.py b/beir/retrieval/search/sparse/sparse_search.py index 2d665d67..8c7949dc 100644 --- a/beir/retrieval/search/sparse/sparse_search.py +++ b/beir/retrieval/search/sparse/sparse_search.py @@ -31,7 +31,7 @@ def search(self, #Get the candidate passages scores = np.asarray(self.sparse_matrix[query_tokens, :].sum(axis=0)).squeeze(0) top_k_ind = np.argpartition(scores, -top_k)[-top_k:] - self.results[qid] = {doc_ids[pid]: float(scores[pid]) for pid in top_k_ind} + self.results[qid] = {doc_ids[pid]: float(scores[pid]) for pid in top_k_ind if doc_ids[pid] != qid} return self.results From 6cc296a7749c0be26cb83cb1ef581fe6fbe7947d Mon Sep 17 00:00:00 2001 From: maximedb Date: Wed, 9 Feb 2022 19:37:30 +0000 Subject: [PATCH 3/8] change the query to be a sparse vector --- beir/retrieval/models/sparta.py | 13 ++++++++----- beir/retrieval/search/sparse/sparse_search.py | 16 +++++++++------- 2 files changed, 17 insertions(+), 12 deletions(-) diff --git a/beir/retrieval/models/sparta.py b/beir/retrieval/models/sparta.py index 972362ec..10579e00 100644 --- a/beir/retrieval/models/sparta.py +++ b/beir/retrieval/models/sparta.py @@ -54,8 +54,11 @@ def _compute_sparse_embeddings(self, documents): return sparse_embeddings def encode_query(self, query: str, **kwargs): - return self.tokenizer(query, add_special_tokens=False)['input_ids'] - + col = self.tokenizer(query, add_special_tokens=False)['input_ids'] + row = [0]*len(col) + data = [1]*len(col) + return csr_matrix((data, (row, col)), shape=(1, len(self.bert_input_embeddings)), dtype=np.float) + def encode_corpus(self, corpus: List[Dict[str, str]], batch_size: int = 16, **kwargs): sentences = [(doc["title"] + self.sep + doc["text"]).strip() for doc in corpus] @@ -69,9 +72,9 @@ def encode_corpus(self, corpus: List[Dict[str, str]], batch_size: int = 16, **kw doc_embs = self._compute_sparse_embeddings(sentences[start_idx: start_idx + batch_size]) for doc_id, emb in enumerate(doc_embs): for tid, score in emb: - col[sparse_idx] = start_idx+doc_id - row[sparse_idx] = tid + col[sparse_idx] = tid + row[sparse_idx] = start_idx+doc_id values[sparse_idx] = score sparse_idx += 1 - return csr_matrix((values, (row, col)), shape=(len(self.bert_input_embeddings), len(sentences)), dtype=np.float) \ No newline at end of file + return csr_matrix((values, (row, col)), shape=(len(sentences), len(self.bert_input_embeddings)), dtype=np.float) \ No newline at end of file diff --git a/beir/retrieval/search/sparse/sparse_search.py b/beir/retrieval/search/sparse/sparse_search.py index 8c7949dc..526d30e4 100644 --- a/beir/retrieval/search/sparse/sparse_search.py +++ b/beir/retrieval/search/sparse/sparse_search.py @@ -2,6 +2,7 @@ from typing import List, Dict, Union, Tuple import logging import numpy as np +import torch logger = logging.getLogger(__name__) @@ -22,16 +23,17 @@ def search(self, query_ids = list(queries.keys()) documents = [corpus[doc_id] for doc_id in doc_ids] logging.info("Computing document embeddings and creating sparse matrix") - self.sparse_matrix = self.model.encode_corpus(documents, batch_size=self.batch_size) - + self.sparse_matrix = self.model.encode_corpus(documents, batch_size=self.batch_size) # [n_doc, n_voc] logging.info("Starting to Retrieve...") for start_idx in trange(0, len(queries), desc='query'): qid = query_ids[start_idx] - query_tokens = self.model.encode_query(queries[qid]) - #Get the candidate passages - scores = np.asarray(self.sparse_matrix[query_tokens, :].sum(axis=0)).squeeze(0) - top_k_ind = np.argpartition(scores, -top_k)[-top_k:] - self.results[qid] = {doc_ids[pid]: float(scores[pid]) for pid in top_k_ind if doc_ids[pid] != qid} + query_vector = self.model.encode_query(queries[qid]) # [1, n_voc] + scores = self.sparse_matrix.dot(query_vector.transpose()).todense() + scores = torch.from_numpy(scores).squeeze() + top_k_values, top_k_indices = torch.topk(scores, top_k, sorted=False) + top_k_values = top_k_values.squeeze().tolist() + top_k_indices = top_k_indices.squeeze().tolist() + self.results[qid] = {doc_ids[pid]: score for pid, score in zip(top_k_indices, top_k_values) if doc_ids[pid] != qid} return self.results From f2c6a2254269b69b16225143f9500aaf97b63ddf Mon Sep 17 00:00:00 2001 From: maximedb Date: Wed, 9 Feb 2022 21:07:15 +0000 Subject: [PATCH 4/8] add splade model and eval script --- beir/retrieval/models/__init__.py | 3 +- beir/retrieval/models/splade.py | 46 +++++++++++++ .../evaluation/sparse/evaluate_splade.py | 64 +++++++++++++++++++ 3 files changed, 112 insertions(+), 1 deletion(-) create mode 100644 beir/retrieval/models/splade.py create mode 100644 examples/retrieval/evaluation/sparse/evaluate_splade.py diff --git a/beir/retrieval/models/__init__.py b/beir/retrieval/models/__init__.py index bba9d4ad..2f2a93d9 100644 --- a/beir/retrieval/models/__init__.py +++ b/beir/retrieval/models/__init__.py @@ -2,4 +2,5 @@ from .use_qa import UseQA from .sparta import SPARTA from .dpr import DPR -from .bpr import BinarySentenceBERT \ No newline at end of file +from .bpr import BinarySentenceBERT +from .splade import SPLADE \ No newline at end of file diff --git a/beir/retrieval/models/splade.py b/beir/retrieval/models/splade.py new file mode 100644 index 00000000..251be27b --- /dev/null +++ b/beir/retrieval/models/splade.py @@ -0,0 +1,46 @@ +from typing import List, Dict + + +import tqdm +import torch +import numpy as np +import transformers +from scipy import sparse + + +class SPLADE: + def __init__(self, model_name_or_path, max_length=256): + self.model = transformers.AutoModelForMaskedLM.from_pretrained(model_name_or_path) + self.tokenizer = transformers.AutoTokenizer.from_pretrained(model_name_or_path) + self.device = "cuda" if torch.cuda.is_available() else "cpu" + self.model.to(self.device) + self.max_length = max_length + + def encode(self, text): + inputs = self.tokenizer(text, max_length=self.max_length, padding=True, truncation=True, return_tensors="pt").to(self.device) + with torch.no_grad(): + outputs = self.model(**inputs) + token_embeddings = outputs[0] + attention_mask = inputs["attention_mask"] + sentence_embedding = torch.max(torch.log(1 + torch.relu(token_embeddings)) * attention_mask.unsqueeze(-1), dim=1).values + return sentence_embedding.cpu().numpy() + + def encode_query(self, query: str, **kwargs) -> sparse.csr_matrix: + """ returns a csr_matrix of shape [1, n_vocab] """ + output = self.encode(query) + return sparse.csr_matrix(output) + + def encode_corpus(self, corpus: List[Dict[str, str]], batch_size: int, **kwargs) -> sparse.csr_matrix: + """ returns a csr_matrix of shape [n_documents, n_vocab] """ + data, row, col = [], [], [] + sentences = [(doc["title"] + " " + doc["text"]).strip() for doc in corpus] + for i in tqdm.tqdm(range(0, len(sentences), batch_size), desc="encode_corpus"): + batch = sentences[i:i+batch_size] + dense = self.encode(batch) + sparse_mat = sparse.coo_matrix(dense) + data.extend(sparse_mat.data) + row.extend(sparse_mat.row + i) + col.extend(sparse_mat.col) + shape = (max(row)+1, self.model.config.vocab_size) + results = sparse.csr_matrix((data, (row, col)), shape=shape, dtype=np.float) + return results diff --git a/examples/retrieval/evaluation/sparse/evaluate_splade.py b/examples/retrieval/evaluation/sparse/evaluate_splade.py new file mode 100644 index 00000000..b0f14679 --- /dev/null +++ b/examples/retrieval/evaluation/sparse/evaluate_splade.py @@ -0,0 +1,64 @@ +from beir import util, LoggingHandler +from beir.retrieval import models +from beir.datasets.data_loader import GenericDataLoader +from beir.retrieval.evaluation import EvaluateRetrieval +from beir.retrieval.search.sparse import SparseSearch + +import logging +import pathlib, os +import random +import shutil + +#### Just some code to print debug information to stdout +logging.basicConfig(format='%(asctime)s - %(message)s', + datefmt='%Y-%m-%d %H:%M:%S', + level=logging.INFO, + handlers=[LoggingHandler()]) +#### /print debug information to stdout + +dataset = "scifact" + +#### Download scifact dataset and unzip the dataset +url = "https://public.ukp.informatik.tu-darmstadt.de/thakur/BEIR/datasets/{}.zip".format(dataset) +out_dir = os.path.join(pathlib.Path(__file__).parent.absolute(), "datasets") +data_path = util.download_and_unzip(url, out_dir) + +#### Provide the data path where scifact has been downloaded and unzipped to the data loader +# data folder would contain these files: +# (1) scifact/corpus.jsonl (format: jsonlines) +# (2) scifact/queries.jsonl (format: jsonlines) +# (3) scifact/qrels/test.tsv (format: tsv ("\t")) + +corpus, queries, qrels = GenericDataLoader(data_folder=data_path).load(split="test") + +#### Sparse Retrieval using SPLADE #### +url = "https://download-de.europe.naverlabs.com/Splade_Release_Jan22/splade_distil_CoCodenser_large.tar.gz" +out_dir = os.path.join(pathlib.Path(__file__).parent.absolute(), "weights") +os.makedirs(out_dir, exist_ok=True) +filename = os.path.join(out_dir, "splade.tar.gz") +model_dir = os.path.join(out_dir, "splade_distil_CoCodenser_large") +if not os.path.exists(model_dir): + util.download_url(url, filename) + shutil.unpack_archive(filename, out_dir) +sparse_model = SparseSearch(models.SPLADE(model_dir, max_length=256), batch_size=24) +retriever = EvaluateRetrieval(sparse_model) + +#### Retrieve dense results (format of results is identical to qrels) +results = retriever.retrieve(corpus, queries) + +#### Evaluate your retrieval using NDCG@k, MAP@K ... + +logging.info("Retriever evaluation for k in: {}".format(retriever.k_values)) +ndcg, _map, recall, precision = retriever.evaluate(qrels, results, retriever.k_values) + +#### Print top-k documents retrieved #### +top_k = 10 + +query_id, ranking_scores = random.choice(list(results.items())) +scores_sorted = sorted(ranking_scores.items(), key=lambda item: item[1], reverse=True) +logging.info("Query : %s\n" % queries[query_id]) + +for rank in range(top_k): + doc_id = scores_sorted[rank][0] + # Format: Rank x: ID [Title] Body + logging.info("Rank %d: %s [%s] - %s\n" % (rank+1, doc_id, corpus[doc_id].get("title"), corpus[doc_id].get("text"))) \ No newline at end of file From 25a7d0475fe6aa531df91df53a6649b28744e66c Mon Sep 17 00:00:00 2001 From: maximedb Date: Fri, 11 Feb 2022 10:24:46 +0000 Subject: [PATCH 5/8] update model URL to match the leaderboard --- examples/retrieval/evaluation/sparse/evaluate_splade.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/examples/retrieval/evaluation/sparse/evaluate_splade.py b/examples/retrieval/evaluation/sparse/evaluate_splade.py index b0f14679..784ff027 100644 --- a/examples/retrieval/evaluation/sparse/evaluate_splade.py +++ b/examples/retrieval/evaluation/sparse/evaluate_splade.py @@ -16,7 +16,7 @@ handlers=[LoggingHandler()]) #### /print debug information to stdout -dataset = "scifact" +dataset = "arguana" #### Download scifact dataset and unzip the dataset url = "https://public.ukp.informatik.tu-darmstadt.de/thakur/BEIR/datasets/{}.zip".format(dataset) @@ -32,11 +32,11 @@ corpus, queries, qrels = GenericDataLoader(data_folder=data_path).load(split="test") #### Sparse Retrieval using SPLADE #### -url = "https://download-de.europe.naverlabs.com/Splade_Release_Jan22/splade_distil_CoCodenser_large.tar.gz" +url = "https://download-de.europe.naverlabs.com/Splade_Release_Jan22/distilsplade_max.tar.gz" out_dir = os.path.join(pathlib.Path(__file__).parent.absolute(), "weights") os.makedirs(out_dir, exist_ok=True) filename = os.path.join(out_dir, "splade.tar.gz") -model_dir = os.path.join(out_dir, "splade_distil_CoCodenser_large") +model_dir = os.path.join(out_dir, "distilsplade_max") if not os.path.exists(model_dir): util.download_url(url, filename) shutil.unpack_archive(filename, out_dir) From 836aaac551a1a696beaf3772c4b68772e5712a28 Mon Sep 17 00:00:00 2001 From: maximedb Date: Sat, 12 Feb 2022 08:31:55 +0000 Subject: [PATCH 6/8] batched query --- beir/retrieval/search/sparse/sparse_search.py | 28 +++++++++++-------- 1 file changed, 16 insertions(+), 12 deletions(-) diff --git a/beir/retrieval/search/sparse/sparse_search.py b/beir/retrieval/search/sparse/sparse_search.py index 526d30e4..0ca3bbee 100644 --- a/beir/retrieval/search/sparse/sparse_search.py +++ b/beir/retrieval/search/sparse/sparse_search.py @@ -17,23 +17,27 @@ def __init__(self, model, batch_size: int = 16, **kwargs): def search(self, corpus: Dict[str, Dict[str, str]], queries: Dict[str, str], - top_k: int, *args, **kwargs) -> Dict[str, Dict[str, float]]: + top_k: int, *args, **kwargs + ) -> Dict[str, Dict[str, float]]: doc_ids = list(corpus.keys()) query_ids = list(queries.keys()) documents = [corpus[doc_id] for doc_id in doc_ids] logging.info("Computing document embeddings and creating sparse matrix") - self.sparse_matrix = self.model.encode_corpus(documents, batch_size=self.batch_size) # [n_doc, n_voc] + self.sparse_matrix_doc = self.model.encode_corpus(documents, batch_size=self.batch_size) # [n_doc, n_voc] logging.info("Starting to Retrieve...") - for start_idx in trange(0, len(queries), desc='query'): - qid = query_ids[start_idx] - query_vector = self.model.encode_query(queries[qid]) # [1, n_voc] - scores = self.sparse_matrix.dot(query_vector.transpose()).todense() - scores = torch.from_numpy(scores).squeeze() - top_k_values, top_k_indices = torch.topk(scores, top_k, sorted=False) - top_k_values = top_k_values.squeeze().tolist() - top_k_indices = top_k_indices.squeeze().tolist() - self.results[qid] = {doc_ids[pid]: score for pid, score in zip(top_k_indices, top_k_values) if doc_ids[pid] != qid} - + for start_idx in trange(0, len(queries), self.batch_size, desc='query'): + local_query_ids = query_ids[start_idx:start_idx+self.batch_size] + local_queries = [queries[qid] for qid in local_query_ids] + qry_matrix = self.model.encode_query(local_queries) + scores = self.sparse_matrix_doc.dot(qry_matrix.transpose()).todense() # [n_doc, vocab]x[vocab, n_qry] -> [n_doc, n_qry] + scores = torch.from_numpy(scores) # [n_qry, n_doc] + top_k_values, top_k_indices = torch.topk(scores, top_k, dim=0, sorted=False) + top_k_values = top_k_values.transpose(0, 1).tolist() # [n_qry, top_k] + top_k_indices = top_k_indices.transpose(0, 1).tolist() # [n_qry, top_k] + for i, qid in enumerate(local_query_ids): + k_ind = top_k_indices[i] + k_val = top_k_values[i] + self.results[qid] = {doc_ids[pid]: score for pid, score in zip(k_ind, k_val) if doc_ids[pid] != qid} return self.results From a9b7ea6d40e6bc1dab9103fcf587be6dad522924 Mon Sep 17 00:00:00 2001 From: maximedb Date: Sat, 12 Feb 2022 08:32:32 +0000 Subject: [PATCH 7/8] CSR construction optimization --- beir/retrieval/models/splade.py | 28 ++++++++++++++++++---------- 1 file changed, 18 insertions(+), 10 deletions(-) diff --git a/beir/retrieval/models/splade.py b/beir/retrieval/models/splade.py index 251be27b..7a569a04 100644 --- a/beir/retrieval/models/splade.py +++ b/beir/retrieval/models/splade.py @@ -1,6 +1,6 @@ from typing import List, Dict - +import array import tqdm import torch import numpy as np @@ -29,18 +29,26 @@ def encode_query(self, query: str, **kwargs) -> sparse.csr_matrix: """ returns a csr_matrix of shape [1, n_vocab] """ output = self.encode(query) return sparse.csr_matrix(output) - - def encode_corpus(self, corpus: List[Dict[str, str]], batch_size: int, **kwargs) -> sparse.csr_matrix: + + def encode_corpus(self, corpus: List[Dict[str, str]], batch_size: int, is_queries=False, **kwargs) -> sparse.csr_matrix: """ returns a csr_matrix of shape [n_documents, n_vocab] """ - data, row, col = [], [], [] + # https://maciejkula.github.io/2015/02/22/incremental-construction-of-sparse-matrices/ + indices = array.array("i") + indptr = array.array("i") + data = array.array("f") sentences = [(doc["title"] + " " + doc["text"]).strip() for doc in corpus] + indptr.append(0) + last_indptr = 0 for i in tqdm.tqdm(range(0, len(sentences), batch_size), desc="encode_corpus"): batch = sentences[i:i+batch_size] dense = self.encode(batch) - sparse_mat = sparse.coo_matrix(dense) - data.extend(sparse_mat.data) - row.extend(sparse_mat.row + i) - col.extend(sparse_mat.col) - shape = (max(row)+1, self.model.config.vocab_size) - results = sparse.csr_matrix((data, (row, col)), shape=shape, dtype=np.float) + nz_rows, nz_cols = np.nonzero(dense) + nz_values = dense[(nz_rows, nz_cols)] + data.extend(nz_values) + local_indptr = np.bincount(nz_rows).cumsum() + last_indptr + indptr.extend(local_indptr) + indices.extend(nz_cols) + last_indptr = local_indptr[-1] + shape = (len(corpus), self.model.config.vocab_size) + results = sparse.csr_matrix((data, indices, indptr), shape=shape, dtype=np.float) return results From 14f1f387e41903c7cb0f8b16685b6620945ae938 Mon Sep 17 00:00:00 2001 From: maximedb Date: Sat, 12 Feb 2022 08:34:36 +0000 Subject: [PATCH 8/8] CSR construction optimization --- .../retrieval/evaluation/sparse/evaluate_splade.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/examples/retrieval/evaluation/sparse/evaluate_splade.py b/examples/retrieval/evaluation/sparse/evaluate_splade.py index 784ff027..d95631fa 100644 --- a/examples/retrieval/evaluation/sparse/evaluate_splade.py +++ b/examples/retrieval/evaluation/sparse/evaluate_splade.py @@ -40,7 +40,7 @@ if not os.path.exists(model_dir): util.download_url(url, filename) shutil.unpack_archive(filename, out_dir) -sparse_model = SparseSearch(models.SPLADE(model_dir, max_length=256), batch_size=24) +sparse_model = SparseSearch(models.SPLADE(model_dir, max_length=256), batch_size=48) retriever = EvaluateRetrieval(sparse_model) #### Retrieve dense results (format of results is identical to qrels) @@ -58,7 +58,7 @@ scores_sorted = sorted(ranking_scores.items(), key=lambda item: item[1], reverse=True) logging.info("Query : %s\n" % queries[query_id]) -for rank in range(top_k): - doc_id = scores_sorted[rank][0] - # Format: Rank x: ID [Title] Body - logging.info("Rank %d: %s [%s] - %s\n" % (rank+1, doc_id, corpus[doc_id].get("title"), corpus[doc_id].get("text"))) \ No newline at end of file +# for rank in range(top_k): +# doc_id = scores_sorted[rank][0] +# # Format: Rank x: ID [Title] Body +# logging.info("Rank %d: %s [%s] - %s\n" % (rank+1, doc_id, corpus[doc_id].get("title"), corpus[doc_id].get("text"))) \ No newline at end of file