RAG integration tests
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user