RAG integration tests

This commit is contained in:
Your Name
2020-09-13 16:32:23 -07:00
parent df5eec9f14
commit 64796004dc
4 changed files with 113 additions and 10 deletions
+4
View File
@@ -62,6 +62,8 @@ RAG_CONFIG_DOC = r"""
The path to the serialized faiss index on disk.
passages_path: (:obj:`str`, optional, defaults to ``None``):
A path to text passages compatible with the faiss index. Required if using :class:`~transformers.retrieval_rag.LegacyIndex`
dummy (:obj:`bool`, optional, defaults to ``False``)
Whether to load a ``dummy`` variant of the dataset specified by ``dataset`` argument.
pretrained_question_encoder_tokenizer_name_or_path: (:obj:`str`, optional, defaults to ``facebook/dpr-question_encoder-single-nq-base``):
A string specifying the ``question_encoder`` tokenizer to be loaded.
pretrained_question_encoder_name_or_path: (:obj:`str`, optional, defaults to ``facebook/dpr-question_encoder-single-nq-base``):
@@ -110,6 +112,7 @@ class RagConfig(PretrainedConfig):
index_name="embeddings",
index_path=None,
passages_path=None,
dummy=False,
pretrained_question_encoder_tokenizer_name_or_path="facebook/dpr-question_encoder-single-nq-base",
pretrained_question_encoder_name_or_path="facebook/dpr-question_encoder-single-nq-base",
pretrained_generator_tokenizer_name_or_path="facebook/bart-large",
@@ -140,6 +143,7 @@ class RagConfig(PretrainedConfig):
self.retrieval_batch_size = retrieval_batch_size
self.passages_path = passages_path
self.index_path = index_path
self.dummy = dummy
self.pretrained_question_encoder_tokenizer_name_or_path = pretrained_question_encoder_tokenizer_name_or_path
self.pretrained_question_encoder_name_or_path = pretrained_question_encoder_name_or_path
+1 -1
View File
@@ -638,7 +638,7 @@ class RagSequence(PreTrainedRagModel):
)
top_cand_inds = (-outputs["loss"]).topk(rag_num_return_sequences)[1]
if logger.get_verbosity() == logging.DEBUG:
if logging.get_verbosity() == logging.DEBUG:
output_strings = self.model.generator_tokenizer.batch_decode(output_sequences)
logger.debug("Hypos with scores:")
for score, hypo in zip(outputs.loss, output_strings):
+17 -8
View File
@@ -150,12 +150,12 @@ class LegacyIndex(Index):
class HFIndex(Index):
"""
A wrapper around an instance of :class:`~nlp.Datasets`. If ``index_path`` is set to ``None``,
we load the pre-computed index available with the :class:`~nlp.arrow_dataset.Dataset`, otherwise, we load the index from the indicated path on disk.
A wrapper around an instance of :class:`~datasets.Datasets`. If ``index_path`` is set to ``None``,
we load the pre-computed index available with the :class:`~datasets.arrow_dataset.Dataset`, otherwise, we load the index from the indicated path on disk.
Args:
dataset (:obj:`str`, optional, defaults to ``wiki_dpr``):
A datatset identifier of the indexed dataset on HuggingFace AWS bucket (list all available datasets and ids with ``nlp.list_datasets()``).
A datatset identifier of the indexed dataset on HuggingFace AWS bucket (list all available datasets and ids with ``datasets.list_datasets()``).
dataset_split (:obj:`str`, optional, defaults to ``train``)
Which split of the ``dataset`` to load.
index_name (:obj:`str`, optional, defaults to ``train``)
@@ -170,13 +170,15 @@ class HFIndex(Index):
dataset_split,
index_name,
index_path,
dummy,
):
super().__init__()
self.dataset = dataset
self.dataset_split = dataset_split
self.index_name = index_name
self.index_path = index_path
self.index = load_dataset(self.dataset, with_index=False, split=self.dataset_split)
self.dummy = dummy
self.index = load_dataset(self.dataset, with_index=False, split=self.dataset_split, dummy=self.dummy)
self._index_initialize = False
def is_initialized(self):
@@ -186,16 +188,23 @@ class HFIndex(Index):
if self.index_path is not None:
self.index.load_faiss_index(index_name=self.index_name, file=self.index_path)
else:
self.index = load_dataset(self.dataset, with_embeddings=True, with_index=True, split=self.dataset_split)
self.index = load_dataset(
self.dataset,
with_embeddings=True,
with_index=True,
split=self.dataset_split,
index_name=self.index_name,
dummy=self.dummy,
)
self._index_initialize = True
def get_doc_dicts(self, doc_ids):
return [self.index[doc_ids[i].tolist()] for i in range(doc_ids.shape[0])]
def get_top_docs(self, query_vectors, n_docs=5):
_, docs = self.index.get_nearest_examples_batch(self.index_name, query_vectors, n_docs)
_, docs = self.index.get_nearest_examples_batch("embeddings", query_vectors, n_docs)
ids = [[int(i) for i in doc["id"]] for doc in docs]
vectors = [doc[self.index_name] for doc in docs]
vectors = [doc["embeddings"] for doc in docs]
return torch.tensor(ids), torch.tensor(vectors)
@@ -217,7 +226,7 @@ class RagRetriever(object):
), "invalid retirever type"
self.retriever = (
HFIndex(config.dataset, config.dataset_split, config.index_name, config.index_path)
HFIndex(config.dataset, config.dataset_split, config.index_name, config.index_path, config.dummy)
if config.retriever_type == "hf_retriever"
else LegacyIndex(config.retrieval_vector_size, config.index_path, config.passages_path)
)
+91 -1
View File
@@ -19,9 +19,10 @@ import unittest
from unittest.mock import patch
from transformers.file_utils import is_datasets_available, is_faiss_available, is_psutil_available, is_torch_available
from transformers.testing_utils import require_torch, torch_device
from transformers.testing_utils import require_torch, slow, torch_device
from .test_configuration_common import ConfigTester
from .test_modeling_bart import TOLERANCE, _assert_tensors_equal
from .test_modeling_common import ids_tensor
@@ -34,6 +35,7 @@ if is_torch_available() and is_datasets_available() and is_faiss_available() and
DPRConfig,
DPRQuestionEncoder,
RagConfig,
RagRetriever,
RagSequence,
RagToken,
)
@@ -358,3 +360,91 @@ class RagModelTest(unittest.TestCase):
result.doc_scores.shape, (self.model_tester.batch_size, self.model_tester.rag_config.n_docs)
)
self.assertEqual(result.loss.shape, torch.Size([]))
@require_torch
@require_retrieval
class RagModelIntegrationTests(unittest.TestCase):
def get_rag_config(self):
return RagConfig(
bos_token_id=0,
decoder_start_token_id=2,
eos_token_id=2,
is_encoder_decoder=True,
pad_token_id=1,
vocab_size=50264,
title_sep=" / ",
doc_sep=" // ",
n_docs=5,
max_combined_length=300,
retriever_type="hf_retriever",
dataset="wiki_dpr",
dataset_split="train",
index_name="exact",
index_path=None,
dummy=True,
retrieval_vector_size=768,
retrieval_batch_size=8,
pretrained_question_encoder_name_or_path="facebook/dpr-question_encoder-single-nq-base",
pretrained_generator_tokenizer_name_or_path="facebook/bart-large-cnn",
pretrained_generator_name_or_path="facebook/bart-large-cnn",
)
@slow
def test_rag_sequence_inference(self):
rag_config = self.get_rag_config()
rag_retriever = RagRetriever(rag_config)
input_ids = torch.tensor([[0, 8155, 22707, 473, 37, 657, 162, 19, 5898, 102, 2]])
attention_mask = torch.tensor([[1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1]])
decoder_input_ids = torch.tensor([[0, 574, 8865, 2505, 2]])
rag_sequence = RagSequence.from_pretrained(config=rag_config).to(torch_device)
with torch.no_grad():
output = rag_sequence(
input_ids,
retriever=rag_retriever,
attention_mask=attention_mask,
decoder_input_ids=decoder_input_ids,
return_loss=True,
print_docs=True,
)
expected_shape = torch.Size([5, 5, 50264])
self.assertEqual(output.logits.shape, expected_shape)
expected_loss = torch.tensor([38.7446])
_assert_tensors_equal(expected_loss, output.loss, atol=TOLERANCE)
expected_doc_scores = torch.tensor([[75.0286, 74.4998, 74.0804, 74.0306, 73.9504]])
_assert_tensors_equal(expected_doc_scores, output.doc_scores, atol=TOLERANCE)
@slow
def test_rag_token_inference(self):
rag_config = self.get_rag_config()
rag_retriever = RagRetriever(rag_config)
input_ids = torch.tensor([[0, 8155, 22707, 473, 37, 657, 162, 19, 5898, 102, 2]])
attention_mask = torch.tensor([[1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1]])
decoder_input_ids = torch.tensor([[0, 574, 8865, 2505, 2]])
rag_token = RagToken.from_pretrained(config=rag_config).to(torch_device)
with torch.no_grad():
output = rag_token(
input_ids,
retriever=rag_retriever,
attention_mask=attention_mask,
decoder_input_ids=decoder_input_ids,
return_loss=True,
)
expected_shape = torch.Size([5, 5, 50264])
self.assertEqual(output.logits.shape, expected_shape)
expected_loss = torch.tensor([38.7045])
_assert_tensors_equal(expected_loss, output.loss, atol=TOLERANCE)
expected_doc_scores = torch.tensor([[75.0286, 74.4998, 74.0804, 74.0306, 73.9504]])
_assert_tensors_equal(expected_doc_scores, output.doc_scores, atol=TOLERANCE)