Compare commits

..
Author SHA1 Message Date
patrickvonplaten a2f830d5a1 delete rag_api 2020-09-17 10:03:41 +02:00
patrickvonplaten 8f5fd79b8f big refactor generate 2020-09-17 09:51:17 +02:00
patrickvonplaten c1be41f452 set generate to default 2020-09-17 09:35:23 +02:00
patrickvonplaten 135689bba3 save intermediate 2020-09-17 09:33:30 +02:00
patrickvonplaten 64141bab07 fix some tests 2020-09-17 02:06:31 +02:00
Patrick von Platen 3cd4a574c2 fix generate for other modles 2020-09-16 19:45:05 +02:00
Patrick von Platen 237f27f724 fix generate problem for rag 2020-09-16 19:42:50 +02:00
Patrick von Platen 4274e9c223 finalize model api 2020-09-16 19:26:57 +02:00
Patrick von Platen 47b137e175 fix conflicts 2020-09-16 19:07:14 +02:00
Patrick von Platen 82afc4b93e finish model outputs 2020-09-16 19:05:22 +02:00
Quentin Lhoest 59ce19cde4 fix retrieval tests 2020-09-16 18:34:17 +02:00
Quentin Lhoest 3abfc19ae0 add doc_ids to retriever's outputs 2020-09-16 18:28:30 +02:00
Patrick von Platen 5b47f0bc3b fix conflict 2020-09-16 17:13:37 +02:00
Patrick von Platen 897101cfce clean model api 2020-09-16 17:11:55 +02:00
Quentin Lhoest 60c8defa01 docstrings + simple retrieval test for distributed 2020-09-16 16:02:06 +02:00
Quentin Lhoest d7e169b3d9 add legacy index URL 2020-09-16 15:52:10 +02:00
Quentin Lhoest 1cf7dcbe71 style 2020-09-16 15:50:10 +02:00
Quentin Lhoest 82a22c20a8 naming 2020-09-16 15:49:51 +02:00
Quentin Lhoest d8fb4c836b make retriever platform agnostic 2020-09-16 15:29:04 +02:00
Patrick von Platen 6bfa18e1a7 make first tests work 2020-09-16 14:52:34 +02:00
Patrick von Platen aca8e30ddf add first version of test 2020-09-16 13:44:46 +02:00
Patrick von Platen 349f85f241 implement thoms suggestions 2020-09-16 11:43:36 +02:00
Patrick von Platen 3860f3144a finalize tests 2020-09-15 22:07:53 +02:00
Patrick von Platen f69a9d32fa add tests 2020-09-15 21:38:19 +02:00
Patrick von Platen c8c5ce0fd3 align test with previous version and make all tests pass 2020-09-15 21:12:54 +02:00
Patrick von Platen 6ab7a4584b Merge branch 'finalize_rag' of https://github.com/huggingface/transformers into finalize_rag 2020-09-15 20:08:53 +02:00
Patrick von Platen e210739bef finish token generate 2020-09-15 20:08:41 +02:00
Quentin Lhoest 8977533c8d add retrieval tests 2020-09-15 18:22:53 +02:00
Patrick von Platen cf9561a4bb change default index name 2020-09-15 17:46:37 +02:00
Patrick von Platen f64b6c1dc8 add labels to models 2020-09-15 16:54:44 +02:00
Patrick von Platen b378005edf small name changes in tokenizer 2020-09-15 16:39:11 +02:00
Patrick von Platen eaf68afffe fix ragfortokengeneration 2020-09-15 16:23:20 +02:00
Patrick von Platen 4f546ad160 make from pretrained more flexible 2020-09-15 16:19:31 +02:00
Patrick von Platen 2884cd7bdb Merge branch 'finalize_rag' of https://github.com/huggingface/transformers into finalize_rag 2020-09-15 16:06:40 +02:00
Patrick von Platen 0f32ad8319 add dpr to autotokenizer 2020-09-15 16:06:07 +02:00
Quentin Lhoest 2c83e1bfd6 LegacyIndex index download refactor 2020-09-15 16:06:07 +02:00
Patrick von Platen 15af641996 finalize retriver api and config api 2020-09-15 15:46:21 +02:00
Patrick von Platen 95cd16275c fix conflict 2020-09-15 15:41:37 +02:00
Patrick von Platen 044fa94285 finalize config 2020-09-15 15:35:10 +02:00
Quentin Lhoest f780b9f415 hardcode legacy index and passages paths (todo: add the right urls) 2020-09-15 14:59:15 +02:00
Quentin Lhoest 2f1211bbb5 init wiki_dpr only once 2020-09-15 13:00:00 +02:00
Quentin Lhoest 00a1fc9ae4 pass config to AutoTokenizer.from_pretrained for Rag tokenizers 2020-09-15 13:00:00 +02:00
Patrick von Platen 2faaa4ad3c Merge branch 'finalize_rag' of https://github.com/huggingface/transformers into finalize_rag 2020-09-15 12:21:54 +02:00
Patrick von Platen c1bc9fe05d delete unnecessary imports 2020-09-15 12:21:25 +02:00
Quentin Lhoest 6e9f30748f don't save paths 2020-09-15 12:14:36 +02:00
Quentin Lhoest bc440f3e7c add RagTokenizer.save/from_pretrained and RagRetriever.save/from_pretrained 2020-09-15 11:44:54 +02:00
Patrick von Platen 2094d37888 make all integration tests pass 2020-09-15 01:23:05 +02:00
Patrick von Platen b3f9e986d9 save working solution 2020-09-15 00:36:34 +02:00
Patrick von Platen b4a094dd97 add from encoder generator 2020-09-14 21:23:25 +02:00
Patrick von Platen 3688823d19 change structure 2020-09-14 21:00:26 +02:00
Patrick von Platen e0a37450e1 change structure 2020-09-14 18:32:14 +02:00
Patrick von Platen 720ee41342 add tests; fix position ids warning 2020-09-14 13:21:55 +02:00
19 changed files with 2070 additions and 1047 deletions
+4 -4
View File
@@ -2,7 +2,7 @@
RAG (for Retrieval Augmented Generation) is a seq2seq model which encapsulates two core components: a question encoder and a generator. During a forward pass, we encode the input with the question encoder and pass it
to the retriever to extract relevant context documents. The documents are then prepended to the input. Such contextualized input is passed to the generator. See [the paper](https://arxiv.org/pdf/2005.11401.pdf) for mored details.
We implement two variants of the model, both presented in the paper - `RagSequence` and `RagToken`. In both cases we use `DPRQuestionEncoder` as the question encoder. As for the generator, two compatible architectures have been tested: `BartForConditionalGeneration` and `T5ForConditionalGeneration`.
We implement two variants of the model, both presented in the paper - `RagSequenceForGeneration. and `RagTokenForGeneration`. In both cases we use `DPRQuestionEncoder` as the question encoder. As for the generator, two compatible architectures have been tested: `BartForConditionalGeneration` and `T5ForConditionalGeneration`.
Key files:
- `modeling_rag.py`, `tokenization_rag.py`, `configuration_rag.py` the core model implementation
@@ -65,9 +65,9 @@ python examples/rag/eval_rag.py \
--model_type rag_sequence \ # RAG model type (rag_token or rag_sequence)
--evaluation_set path/to/output/biencoder-nq-dev.questions \ # an input dataset for evaluation
--gold_data_path path/to/output/biencoder-nq-dev.pages \ # a dataset containing ground truth answers for samples from the evaluation_set
--predictions_path path/to/retrieval_preds.tsv \ # path to a file in which predictions will be stored
--predictions_filename retrieval_preds.tsv \ # name of file in which predictions will be stored
--eval_mode retrieval \ # indicates whether we're performing retrieval evaluation or e2e evaluation
--recalculate # if predictions_path already exists, and this option is set - we regenerate the answers, otherwise we reuse the predicsion file to calculate metrics.
--recalculate # if predictions_filename already exists, and this option is set - we regenerate the answers, otherwise we reuse the predicsion file to calculate metrics.
```
@@ -78,7 +78,7 @@ python examples/rag/eval_rag.py \
--model_type rag_sequence \
--evaluation_set path/to/test.source \
--gold_data_path path/to/gold_data \
--predictions_path path/to/e2e_preds.txt \
--predictions_filename e2e_preds.txt \
--eval_mode e2e \ # indicates whether we're performing retrieval evaluation or e2e evaluation (default)
--n_docs 5 \ # You can experiment with retrieving different number of documents at evaluation time
--print_predictions
+19 -14
View File
@@ -10,7 +10,13 @@ import pandas as pd
import torch
from tqdm import tqdm
from transformers import BartForConditionalGeneration, BartTokenizer, RagRetriever, RagSequence, RagToken
from transformers import (
BartForConditionalGeneration,
BartTokenizer,
RagRetriever,
RagSequenceForGeneration,
RagTokenForGeneration,
)
from transformers import logging as transformers_logging
@@ -100,7 +106,7 @@ def evaluate_batch_retrieval(args, rag_model, tokenizer, retriever, questions):
)
retriever_input_embs = rag_model.model.question_encoder(retriever_inputs["input_ids"].to(args.device))[0]
_, all_docs = retriever.retrieve(retriever_input_embs, rag_model.config.n_docs)
_, all_docs = retriever.retrieve(retriever_input_embs.numpy(), rag_model.config.n_docs)
provenance_strings = []
for docs in all_docs:
@@ -196,10 +202,10 @@ def get_args():
"ans - a single line of the gold file contains the expected answer string",
)
parser.add_argument(
"--predictions_path",
"--predictions_filename",
type=str,
default="predictions.txt",
help="Path under which to store prediction files. The base dir needs to exists, the file will be generated.",
help="Name of the predictions file, to be stored in the checkpoints directry",
)
parser.add_argument(
"--eval_all_checkpoints",
@@ -247,7 +253,7 @@ def main(args):
args.model_type = infer_model_type(args.model_name_or_path)
assert args.model_type is not None
if args.model_type.startswith("rag"):
model_class = RagToken if args.model_type == "rag_token" else RagSequence
model_class = RagTokenForGeneration if args.model_type == "rag_token" else RagSequenceForGeneration
model_kwargs["n_docs"] = args.n_docs
if args.retriever_type is not None:
model_kwargs["retriever_type"] = args.retriever_type
@@ -268,18 +274,19 @@ def main(args):
evaluate_batch_fn = evaluate_batch_e2e if args.eval_mode == "e2e" else evaluate_batch_retrieval
for checkpoint in checkpoints:
if os.path.exists(args.predictions_path) and (not args.recalculate):
logger.info("Calculating metrics based on an existing predictions file: {}".format(args.predictions_path))
score_fn(args, args.predictions_path, args.gold_data_path)
predictions_path = os.path.join(checkpoint, args.predictions_filename)
if os.path.exists(predictions_path) and (not args.recalculate):
logger.info("Calculating metrics based on an existing predictions file: {}".format(predictions_path))
score_fn(args, predictions_path, args.gold_data_path)
continue
logger.info("***** Running evaluation for {} *****".format(checkpoint))
logger.info(" Batch size = %d", args.eval_batch_size)
logger.info(" Predictions will be stored under {}".format(args.predictions_path))
logger.info(" Predictions will be stored under {}".format(predictions_path))
model = model_class.from_pretrained(checkpoint, **model_kwargs)
model.to(args.device)
retriever = RagRetriever(model.config)
retriever = RagRetriever(model.config) # TODO: add tokenizers
tokenizer = (
retriever.generator_tokenizer
if args.model_type != "bart" and args.eval_mode == "e2e"
@@ -287,10 +294,8 @@ def main(args):
if args.model_type != "bart" and args.eval_mode == "retrieval"
else BartTokenizer.from_pretrained("facebook/bart-large")
)
if args.model_type != "bart":
retriever.init_retrieval(distributed_port=12345)
with open(args.evaluation_set, "r") as eval_file, open(args.predictions_path, "w") as preds_file:
with open(args.evaluation_set, "r") as eval_file, open(predictions_path, "w") as preds_file:
questions = []
for line in tqdm(eval_file):
questions.append(line.strip())
@@ -304,7 +309,7 @@ def main(args):
preds_file.write("\n".join(answers))
preds_file.flush()
score_fn(args, args.predictions_path, args.gold_data_path)
score_fn(args, predictions_path, args.gold_data_path)
if __name__ == "__main__":
+13 -12
View File
@@ -22,9 +22,9 @@ from transformers import (
AutoTokenizer,
BartForConditionalGeneration,
RagConfig,
RagRetriever,
RagSequence,
RagToken,
RagPyTorchDistributedRetriever,
RagSequenceForGeneration,
RagTokenForGeneration,
T5ForConditionalGeneration,
get_linear_schedule_with_warmup,
)
@@ -74,9 +74,9 @@ class GenerativeQAModule(BaseTransformer):
if isinstance(hparams, dict):
hparams = AttrDict(hparams)
if hparams.model_type == "rag_sequence":
self.model_class = RagSequence
self.model_class = RagSequenceForGeneration
elif hparams.model_type == "rag_token":
self.model_class = RagToken
self.model_class = RagTokenForGeneration
elif hparams.model_type == "bart":
self.model_class = BartForConditionalGeneration
else:
@@ -89,14 +89,13 @@ class GenerativeQAModule(BaseTransformer):
# set extra_model_params for generator configs and load_model
extra_model_params = ("encoder_layerdrop", "decoder_layerdrop", "attention_dropout", "dropout")
if self.is_rag_model:
pretrained_generator_name_or_path = (
config.pretrained_generator_name_or_path
if config.pretrained_generator_name_or_path is not None
else os.path.join(hparams.model_name_or_path, "generator")
generator_config = AutoConfig.from_pretrained(
config.pretrained_generator_name_or_path, prefix=config.prefix
)
generator_config = AutoConfig.from_pretrained(pretrained_generator_name_or_path, prefix=config.prefix)
hparams, generator_config = set_extra_model_params(extra_model_params, hparams, generator_config)
model = self.model_class.from_pretrained(hparams.model_name_or_path, generator_config=generator_config)
model = self.model_class.from_pretrained(
hparams.model_name_or_path, config=config, generator_config=generator_config
)
else:
if args.prefix is not None:
setattr(config, "prefix", args.prefix)
@@ -112,7 +111,9 @@ class GenerativeQAModule(BaseTransformer):
super().__init__(hparams, config=config, tokenizer=tokenizer, model=model)
self.retriever = RagRetriever(self.model.config) if self.is_rag_model else None
self.retriever = (
RagPyTorchDistributedRetriever(self.model.config) if self.is_rag_model else None
) # TODO add tokenizers
save_git_info(self.hparams.output_dir)
self.output_dir = Path(self.hparams.output_dir)
+2 -2
View File
@@ -715,9 +715,9 @@ if is_tf_available():
if is_torch_available() and is_datasets_available() and is_faiss_available() and is_psutil_available():
from .configuration_rag import RagConfig
from .modeling_rag import RagModel, RagSequence, RagToken
from .modeling_rag import RagModel, RagSequenceForGeneration, RagTokenForGeneration
from .retrieval_rag import RagRetriever
from .tokenization_rag import RagDefaultTokenizer
from .tokenization_rag import RagTokenizer
if not is_tf_available() and not is_torch_available():
+4
View File
@@ -24,6 +24,7 @@ from .configuration_bert_generation import BertGenerationConfig
from .configuration_camembert import CAMEMBERT_PRETRAINED_CONFIG_ARCHIVE_MAP, CamembertConfig
from .configuration_ctrl import CTRL_PRETRAINED_CONFIG_ARCHIVE_MAP, CTRLConfig
from .configuration_distilbert import DISTILBERT_PRETRAINED_CONFIG_ARCHIVE_MAP, DistilBertConfig
from .configuration_dpr import DPR_PRETRAINED_CONFIG_ARCHIVE_MAP, DPRConfig
from .configuration_electra import ELECTRA_PRETRAINED_CONFIG_ARCHIVE_MAP, ElectraConfig
from .configuration_encoder_decoder import EncoderDecoderConfig
from .configuration_flaubert import FLAUBERT_PRETRAINED_CONFIG_ARCHIVE_MAP, FlaubertConfig
@@ -71,6 +72,7 @@ ALL_PRETRAINED_CONFIG_ARCHIVE_MAP = dict(
RETRIBERT_PRETRAINED_CONFIG_ARCHIVE_MAP,
FUNNEL_PRETRAINED_CONFIG_ARCHIVE_MAP,
LXMERT_PRETRAINED_CONFIG_ARCHIVE_MAP,
DPR_PRETRAINED_CONFIG_ARCHIVE_MAP,
]
for key, value, in pretrained_map.items()
)
@@ -105,6 +107,7 @@ CONFIG_MAPPING = OrderedDict(
("encoder-decoder", EncoderDecoderConfig),
("funnel", FunnelConfig),
("lxmert", LxmertConfig),
("dpr", DPRConfig),
]
)
@@ -137,6 +140,7 @@ MODEL_NAMES_MAPPING = OrderedDict(
("encoder-decoder", "Encoder decoder"),
("funnel", "Funnel Transformer"),
("lxmert", "LXMERT"),
("dpr", "DPR"),
]
)
+65 -4
View File
@@ -14,7 +14,7 @@
# limitations under the License.
""" DPR model configuration """
from .configuration_bert import BertConfig
from .configuration_utils import PretrainedConfig
from .utils import logging
@@ -27,7 +27,7 @@ DPR_PRETRAINED_CONFIG_ARCHIVE_MAP = {
}
class DPRConfig(BertConfig):
class DPRConfig(PretrainedConfig):
r"""
:class:`~transformers.DPRConfig` is the configuration class to store the configuration of a
`DPRModel`.
@@ -36,12 +36,73 @@ class DPRConfig(BertConfig):
It is used to instantiate the components of the DPR model.
Args:
vocab_size (:obj:`int`, optional, defaults to 30522):
Vocabulary size of the DPR model. Defines the different tokens that
can be represented by the `inputs_ids` passed to the forward method of :class:`~transformers.BertModel`.
hidden_size (:obj:`int`, optional, defaults to 768):
Dimensionality of the encoder layers and the pooler layer.
num_hidden_layers (:obj:`int`, optional, defaults to 12):
Number of hidden layers in the Transformer encoder.
num_attention_heads (:obj:`int`, optional, defaults to 12):
Number of attention heads for each attention layer in the Transformer encoder.
intermediate_size (:obj:`int`, optional, defaults to 3072):
Dimensionality of the "intermediate" (i.e., feed-forward) layer in the Transformer encoder.
hidden_act (:obj:`str` or :obj:`function`, optional, defaults to "gelu"):
The non-linear activation function (function or string) in the encoder and pooler.
If string, "gelu", "relu", "swish" and "gelu_new" are supported.
hidden_dropout_prob (:obj:`float`, optional, defaults to 0.1):
The dropout probabilitiy for all fully connected layers in the embeddings, encoder, and pooler.
attention_probs_dropout_prob (:obj:`float`, optional, defaults to 0.1):
The dropout ratio for the attention probabilities.
max_position_embeddings (:obj:`int`, optional, defaults to 512):
The maximum sequence length that this model might ever be used with.
Typically set this to something large just in case (e.g., 512 or 1024 or 2048).
type_vocab_size (:obj:`int`, optional, defaults to 2):
The vocabulary size of the `token_type_ids` passed into :class:`~transformers.BertModel`.
initializer_range (:obj:`float`, optional, defaults to 0.02):
The standard deviation of the truncated_normal_initializer for initializing all weight matrices.
layer_norm_eps (:obj:`float`, optional, defaults to 1e-12):
The epsilon used by the layer normalization layers.
gradient_checkpointing (:obj:`bool`, optional, defaults to :obj:`False`):
If True, use gradient checkpointing to save memory at the expense of slower backward pass.
projection_dim (:obj:`int`, optional, defaults to 0):
Dimension of the projection for the context and question encoders.
If it is set to zero (default), then no projection is done.
"""
model_type = "dpr"
def __init__(self, projection_dim: int = 0, **kwargs): # projection of the encoders, 0 for no projection
super().__init__(**kwargs)
def __init__(
self,
vocab_size=30522,
hidden_size=768,
num_hidden_layers=12,
num_attention_heads=12,
intermediate_size=3072,
hidden_act="gelu",
hidden_dropout_prob=0.1,
attention_probs_dropout_prob=0.1,
max_position_embeddings=512,
type_vocab_size=2,
initializer_range=0.02,
layer_norm_eps=1e-12,
pad_token_id=0,
gradient_checkpointing=False,
projection_dim: int = 0,
**kwargs
):
super().__init__(pad_token_id=pad_token_id, **kwargs)
self.vocab_size = vocab_size
self.hidden_size = hidden_size
self.num_hidden_layers = num_hidden_layers
self.num_attention_heads = num_attention_heads
self.hidden_act = hidden_act
self.intermediate_size = intermediate_size
self.hidden_dropout_prob = hidden_dropout_prob
self.attention_probs_dropout_prob = attention_probs_dropout_prob
self.max_position_embeddings = max_position_embeddings
self.type_vocab_size = type_vocab_size
self.initializer_range = initializer_range
self.layer_norm_eps = layer_norm_eps
self.gradient_checkpointing = gradient_checkpointing
self.projection_dim = projection_dim
+73 -29
View File
@@ -14,6 +14,7 @@
# limitations under the License.
""" RAG model configuration """
import copy
from .configuration_utils import PretrainedConfig
from .file_utils import add_start_docstrings_to_callable
@@ -50,7 +51,7 @@ RAG_CONFIG_DOC = r"""
retriever_type (:obj:`str`, optional, defaults to ``hf_retriever``):
A type of index encapsulated by the ``retriever``. Possible options include:
- ``hf_retriever`` - and index build for an instance of :class:`~datasets.Datasets`
- ``hf_retriever`` - and index build for an instance of :class:`~nlp.Datasets`
- ``legacy_retriever`` - an index build with the native DPR implementation (see https://github.com/facebookresearch/DPR for details).
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()``).
@@ -59,28 +60,19 @@ RAG_CONFIG_DOC = r"""
index_name (:obj:`str`, optional, defaults to ``train``)
The index_name of the index associated with the ``dataset``.
index_path (:obj:`str`, optional, defaults to ``None``)
Can be either:
- A path to a serialized faiss index on disk, compatible with :class:`~transformers.retrieval_rag.HFIndex`
- A string with the `shortcut name` of a pretrained index compatible with
:class:`~transformers.retrieval_rag.LegacyIndex` to load from cache or download,
e.g. ``facebook/rag-index``.
- A path to a `directory` containing index files compatible with
:class:`~transformers.retrieval_rag.LegacyIndex`
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``):
A string specifying the ``question_encoder`` model to be loaded. If a RAG model is loaded from ``pretrained_model_name_or_path``
and ``pretrained_question_encoder_name_or_path`` is not ``None``, ``pretrained_question_encoder_name_or_path`` takes precedence
over the question encoder model specified by the ``pretrained_model_name_or_path``.
A string specifying the ``question_encoder`` model to be loaded.
pretrained_generator_tokenizer_name_or_path: (:obj:`str`, optional, defaults to ``facebook/bart-large``):
A string specifying the ``generator`` tokenizer to be loaded.
pretrained_generator_name_or_path: (:obj:`str`, optional, defaults to ``facebook/bart-large``):
A string specifying the ``generator`` model to be loaded. If a RAG model is loaded from ``pretrained_model_name_or_path``
and ``pretrained_generator_name_or_path`` is not ``None``, ``pretrained_generator_name_or_path`` takes precedence
over the generator model specified by the ``pretrained_model_name_or_path``.
A string specifying the ``generator`` model to be loaded.
Args linked to the tokenizer - they have to be compatible with equivalent parameters of the ``generator``:
prefix (:obj:`str`, `optional`):
@@ -93,6 +85,12 @@ RAG_CONFIG_DOC = r"""
The id of the `end-of-stream` token.
decoder_start_token_id** (:obj:`int`, `optional`):
If an encoder-decoder model starts decoding with a different token than `bos`, the id of that token.
marginalize (:obj:`bool`, `optional`, defaults to :obj:`False`):
If :obj:`True`, `logits`, returned as part of :class:`~transformers.file_utils.Seq2SeqLMOutputWithDocs` are marginalized, yielding
the shape of :obj:`(batch_size, sequence_length, hidden_size)`. Otherwise we return raw, non-marginalized logits of shape
:obj:`(batch_size * n_docs, sequence_length, hidden_size)`. ``marginalize`` is set to :obj:`True` during generation. The parameter is
ignored if ``return_loss`` is set to :obj:`True`.
"""
@@ -115,19 +113,40 @@ class RagConfig(PretrainedConfig):
max_combined_length=300,
retrieval_vector_size=768,
retrieval_batch_size=8,
retriever_type="hf_retriever",
dataset="wiki_dpr",
dataset_split="train",
index_name="embeddings",
index_name="compressed",
index_path=None,
dummy=False,
pretrained_question_encoder_tokenizer_name_or_path="facebook/dpr-question_encoder-single-nq-base",
pretrained_question_encoder_name_or_path=None,
pretrained_generator_tokenizer_name_or_path="facebook/bart-large",
pretrained_generator_name_or_path=None,
passages_path=None,
use_dummy_dataset=False,
reduce_loss=False,
label_smoothing=0.0,
deduplicate=True, # defaults to True
num_doc_return_sequences=1, # defaults to 1
num_doc_beams=1, # defaults to 1
exclude_bos_score=False,
do_marginalize=False,
output_retrieved=False,
**kwargs
):
super().__init__(**kwargs)
assert (
"question_encoder" in kwargs and "generator" in kwargs
), "Config has to be initialized with question_encoder and generator config"
question_encoder_config = kwargs.pop("question_encoder")
question_encoder_model_type = question_encoder_config.pop("model_type")
decoder_config = kwargs.pop("generator")
decoder_model_type = decoder_config.pop("model_type")
from .configuration_auto import AutoConfig
self.question_encoder = AutoConfig.for_model(question_encoder_model_type, **question_encoder_config)
self.generator = AutoConfig.for_model(decoder_model_type, **decoder_config)
self.reduce_loss = reduce_loss
self.label_smoothing = label_smoothing
self.exclude_bos_score = exclude_bos_score
self.do_marginalize = do_marginalize
self.vocab_size = vocab_size
self.is_encoder_decoder = is_encoder_decoder
self.prefix = prefix
@@ -141,18 +160,43 @@ class RagConfig(PretrainedConfig):
self.n_docs = n_docs
self.max_combined_length = max_combined_length
self.retriever_type = retriever_type
self.dataset = dataset
self.dataset_split = dataset_split
self.index_name = index_name
self.retrieval_vector_size = retrieval_vector_size
self.retrieval_batch_size = retrieval_batch_size
self.passages_path = passages_path
self.index_path = index_path
self.dummy = dummy
self.use_dummy_dataset = use_dummy_dataset
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
self.pretrained_generator_tokenizer_name_or_path = pretrained_generator_tokenizer_name_or_path
self.pretrained_generator_name_or_path = pretrained_generator_name_or_path
self.output_retrieved = output_retrieved
self.deduplicate = deduplicate
self.num_doc_return_sequences = num_doc_return_sequences
self.num_doc_beams = num_doc_beams
@classmethod
def from_question_encoder_generator_configs(
cls, question_encoder_config: PretrainedConfig, generator_config: PretrainedConfig, **kwargs
) -> PretrainedConfig:
r"""
Instantiate a :class:`~transformers.EncoderDecoderConfig` (or a derived class) from a pre-trained encoder model configuration and decoder model configuration.
Returns:
:class:`EncoderDecoderConfig`: An instance of a configuration object
"""
return cls(question_encoder=question_encoder_config.to_dict(), generator=generator_config.to_dict(), **kwargs)
def to_dict(self):
"""
Serializes this instance to a Python dictionary. Override the default `to_dict()` from `PretrainedConfig`.
Returns:
:obj:`Dict[str, any]`: Dictionary of all the attributes that make up this configuration instance,
"""
output = copy.deepcopy(self.__dict__)
output["question_encoder"] = self.question_encoder.to_dict()
output["generator"] = self.generator.to_dict()
output["model_type"] = self.__class__.model_type
return output
+2
View File
@@ -69,6 +69,7 @@ try:
import datasets # noqa: F401
_datasets_available = True
logger.debug(f"Succesfully imported datasets version {datasets.__version__}")
except ImportError:
_datasets_available = False
@@ -124,6 +125,7 @@ try:
import faiss # noqa: F401
_faiss_available = True
logger.debug(f"Succesfully imported faiss version {faiss.__version__}")
except ImportError:
_faiss_available = False
+1 -9
View File
@@ -399,15 +399,7 @@ class GenerationMixin:
# get encoder and store encoder outputs
encoder = self.get_encoder()
if "retriever" in model_kwargs:
encoder_outputs: ModelOutput = encoder(
input_ids,
retriever=model_kwargs["retriever"],
attention_mask=attention_mask,
return_dict=True,
)
else:
encoder_outputs: ModelOutput = encoder(input_ids, attention_mask=attention_mask, return_dict=True)
encoder_outputs: ModelOutput = encoder(input_ids, attention_mask=attention_mask, return_dict=True)
# Expand input ids if num_beams > 1 or num_return_sequences > 1
if num_return_sequences > 1 or num_beams > 1:
+3
View File
@@ -27,6 +27,7 @@ from .configuration_auto import (
CamembertConfig,
CTRLConfig,
DistilBertConfig,
DPRConfig,
ElectraConfig,
EncoderDecoderConfig,
FlaubertConfig,
@@ -94,6 +95,7 @@ from .modeling_distilbert import (
DistilBertForTokenClassification,
DistilBertModel,
)
from .modeling_dpr import DPRQuestionEncoder
from .modeling_electra import (
ElectraForMaskedLM,
ElectraForMultipleChoice,
@@ -217,6 +219,7 @@ MODEL_MAPPING = OrderedDict(
(FunnelConfig, FunnelModel),
(LxmertConfig, LxmertModel),
(BertGenerationConfig, BertGenerationEncoder),
(DPRConfig, DPRQuestionEncoder),
]
)
+3
View File
@@ -272,6 +272,7 @@ class DPRPretrainedContextEncoder(PreTrainedModel):
config_class = DPRConfig
load_tf_weights = None
base_model_prefix = "ctx_encoder"
authorized_missing_keys = [r"position_ids"]
def init_weights(self):
self.ctx_encoder.init_weights()
@@ -285,6 +286,7 @@ class DPRPretrainedQuestionEncoder(PreTrainedModel):
config_class = DPRConfig
load_tf_weights = None
base_model_prefix = "question_encoder"
authorized_missing_keys = [r"position_ids"]
def init_weights(self):
self.question_encoder.init_weights()
@@ -298,6 +300,7 @@ class DPRPretrainedReader(PreTrainedModel):
config_class = DPRConfig
load_tf_weights = None
base_model_prefix = "span_predictor"
authorized_missing_keys = [r"position_ids"]
def init_weights(self):
self.span_predictor.encoder.init_weights()
+1 -1
View File
@@ -250,7 +250,7 @@ class EncoderDecoderModel(PreTrainedModel):
encoder_config.is_decoder = False
encoder_config.add_cross_attention = False
kwargs_encoder["config"] = encoder_config
kwargs_encoder["config"] = encoder_config
encoder = AutoModel.from_pretrained(encoder_pretrained_model_name_or_path, *model_args, **kwargs_encoder)
File diff suppressed because it is too large Load Diff
+281 -198
View File
@@ -17,6 +17,7 @@
import os
import pickle
import time
from typing import Iterable, List, Optional, Tuple
import faiss
import numpy as np
@@ -25,16 +26,20 @@ import torch
import torch.distributed as dist
from datasets import load_dataset
from .configuration_rag import RagConfig
from .file_utils import cached_path, is_remote_url
from .tokenization_auto import AutoTokenizer
from .tokenization_dpr import DPRQuestionEncoderTokenizer
from .tokenization_rag import RagTokenizer
from .tokenization_t5 import T5Tokenizer
from .tokenization_utils_base import BatchEncoding
from .utils import logging
logger = logging.get_logger(__name__)
LEGACY_INDEX_PATH = "https://storage.googleapis.com/huggingface-nlp/datasets/wiki_dpr/"
class Index(object):
"""
A base class for the Indices encapsulated by the :class:`~transformers.RagRetriever`.
@@ -43,29 +48,29 @@ class Index(object):
def __init__(self, *args, **kwargs):
pass
def get_doc_dicts(self, doc_ids):
def get_doc_dicts(self, doc_ids: np.ndarray) -> List[dict]:
"""
Returns a list of dictionaries, containing titles and text of the retrieved documents.
Args:
doc_ids (:obj:`torch.Tensor` of shape :obj:`(batch_size, n_docs)`):
doc_ids (:obj:`np.ndarray` of shape :obj:`(batch_size, n_docs)`):
A tensor of document indices.
"""
pass
def get_top_docs(self, query_vectors, n_docs):
def get_top_docs(self, question_hidden_states: np.ndarray, n_docs=5) -> Tuple[np.ndarray, np.ndarray]:
"""
For each query in the batch, retrieves ``n_docs`` documents.
Args:
query_vectors (:obj:`np.array` of shape :obj:`(batch_size, vector_size):
question_hidden_states (:obj:`np.ndarray` of shape :obj:`(batch_size, vector_size):
An array of query vectors.
n_docs (:obj:`int`):
The number of docs retrieved per query.
Returns:
:obj:`torch.Tensor` of shape :obj:`(batch_size, n_docs)`: A tensor of indices of retrieved documents.
:obj:`torch.Tensor` of shape :obj:`(batch_size, vector_size)`: A tensor of vector representations of retrieved documents.
:obj:`np.ndarray` of shape :obj:`(batch_size, n_docs)`: A tensor of indices of retrieved documents.
:obj:`np.ndarray` of shape :obj:`(batch_size, vector_size)`: A tensor of vector representations of retrieved documents.
"""
raise NotImplementedError
@@ -92,13 +97,8 @@ class LegacyIndex(Index):
vector_size (:obj:`int`):
The dimension of indexed vectors.
index_path (:obj:`str`):
Can be either
- A string with the `identifier name` of a pretrained index compatible with
:class:`~transformers.retrieval_rag.LegacyIndex` to load from cache or download,
e.g. ``facebook/rag-index``.
- A path to a `directory` containing index files compatible with
:class:`~transformers.retrieval_rag.LegacyIndex`
A path to a `directory` containing index files compatible with
:class:`~transformers.retrieval_rag.LegacyIndex`
"""
INDEX_FILENAME = "hf_bert_base.hnswSQ8_correct_phi_128.c_index"
@@ -123,7 +123,7 @@ class LegacyIndex(Index):
except EnvironmentError:
msg = (
f"Can't load '{archive_file}'. Make sure that:\n\n"
f"- '{index_path}' is a correct identifier listed on 'https://huggingface.co/models'\n\n"
f"- '{index_path}' is a correct remote path to a directory containing a file named {filename}"
f"- or '{index_path}' is the correct path to a directory containing a file named {filename}.\n\n"
)
raise EnvironmentError(msg)
@@ -134,6 +134,7 @@ class LegacyIndex(Index):
return resolved_archive_file
def _load_passages(self):
logger.info("Loading passages from {}".format(self.index_path))
passages_path = self._resolve_path(self.index_path, self.PASSAGE_FILENAME)
with open(passages_path, "rb") as passages_file:
passages = pickle.load(passages_file)
@@ -175,13 +176,13 @@ class LegacyIndex(Index):
doc_dicts.append(doc_dict)
return doc_dicts
def get_top_docs(self, query_vectors: np.array, n_docs: int = 5):
aux_dim = np.zeros(len(query_vectors), dtype="float32").reshape(-1, 1)
query_nhsw_vectors = np.hstack((query_vectors, aux_dim))
def get_top_docs(self, question_hidden_states: np.ndarray, n_docs=5) -> Tuple[np.ndarray, np.ndarray]:
aux_dim = np.zeros(len(question_hidden_states), dtype="float32").reshape(-1, 1)
query_nhsw_vectors = np.hstack((question_hidden_states, aux_dim))
_, docs_ids = self.index.search(query_nhsw_vectors, n_docs)
vectors = [[self.index.reconstruct(int(doc_id))[:-1] for doc_id in doc_ids] for doc_ids in docs_ids]
ids = [[int(self.index_id_to_db_id[doc_id]) for doc_id in doc_ids] for doc_ids in docs_ids]
return torch.tensor(ids), torch.tensor(vectors)
return np.array(ids), np.array(vectors)
class HFIndex(Index):
@@ -202,46 +203,52 @@ class HFIndex(Index):
def __init__(
self,
dataset,
dataset_split,
index_name,
index_path,
dummy,
dataset_name: str,
dataset_split: str,
index_name: str,
index_path: Optional[str] = None,
use_dummy_dataset=False,
):
super().__init__()
self.dataset = dataset
self.dataset_name = dataset_name
self.dataset_split = dataset_split
self.index_name = index_name
self.index_path = index_path
self.dummy = dummy
self.index = load_dataset(self.dataset, with_index=False, split=self.dataset_split, dummy=self.dummy)
self.use_dummy_dataset = use_dummy_dataset
self._index_initialize = False
logger.info("Loading passages from {}".format(self.dataset_name))
self.dataset = load_dataset(
self.dataset_name, with_index=False, split=self.dataset_split, dummy=self.use_dummy_dataset
)
def is_initialized(self):
return self._index_initialize
def init_index(self):
if self.index_path is not None:
logger.info("Loading index from {}".format(self.index_path))
self.index.load_faiss_index(index_name=self.index_name, file=self.index_path)
else:
self.index = load_dataset(
self.dataset,
logger.info("Loading index from {}".format(self.dataset_name + " with index name " + self.index_name))
self.dataset = load_dataset(
self.dataset_name,
with_embeddings=True,
with_index=True,
split=self.dataset_split,
index_name=self.index_name,
dummy=self.dummy,
dummy=self.use_dummy_dataset,
)
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_doc_dicts(self, doc_ids: np.ndarray) -> List[dict]:
return [self.dataset[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("embeddings", query_vectors, n_docs)
def get_top_docs(self, question_hidden_states: np.ndarray, n_docs=5) -> Tuple[np.ndarray, np.ndarray]:
_, docs = self.dataset.get_nearest_examples_batch("embeddings", question_hidden_states, n_docs)
ids = [[int(i) for i in doc["id"]] for doc in docs]
vectors = [doc["embeddings"] for doc in docs]
return torch.tensor(ids), torch.tensor(vectors)
return np.array(ids), np.array(vectors) # shapes (batch_size, n_docs) and (batch_size, n_docs, d)
class RagRetriever(object):
@@ -255,23 +262,23 @@ class RagRetriever(object):
The configuration of the RAG model this Retriever is used with. Contains parameters indicating which ``Index`` to build.
"""
def __init__(self, config):
super().__init__()
assert (
config.retriever_type == "hf_retriever" or config.retriever_type == "legacy_retriever"
), "invalid retirever type"
_init_retrieval = True
self.retriever = (
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)
def __init__(self, config, question_encoder_tokenizer, generator_tokenizer):
super().__init__()
self.index = (
LegacyIndex(
config.retrieval_vector_size,
config.index_path or LEGACY_INDEX_PATH,
)
if config.index_name == "legacy"
else HFIndex(
config.dataset, config.dataset_split, config.index_name, config.index_path, config.use_dummy_dataset
)
)
self.generator_tokenizer = AutoTokenizer.from_pretrained(config.pretrained_generator_tokenizer_name_or_path)
# TODO(piktus): To be replaced with AutoTokenizer once it supports DPRQuestionEncoderTokenizer
self.question_encoder_tokenizer = DPRQuestionEncoderTokenizer.from_pretrained(
config.pretrained_question_encoder_tokenizer_name_or_path
)
self.process_group = None
self.generator_tokenizer = generator_tokenizer
self.question_encoder_tokenizer = question_encoder_tokenizer
self.n_docs = config.n_docs
self.batch_size = config.retrieval_batch_size
@@ -279,14 +286,212 @@ class RagRetriever(object):
self.batch_size *= torch.cuda.device_count()
self.config = config
if self._init_retrieval:
self.init_retrieval()
def init_retrieval(self, distributed_port):
@classmethod
def from_pretrained(cls, retriever_name_or_path, **kwargs):
config = RagConfig.from_pretrained(retriever_name_or_path, **kwargs)
rag_tokenizer = RagTokenizer.from_pretrained(retriever_name_or_path, config=config)
question_encoder_tokenizer = rag_tokenizer.question_encoder
generator_tokenizer = rag_tokenizer.generator
return cls(
config, question_encoder_tokenizer=question_encoder_tokenizer, generator_tokenizer=generator_tokenizer
)
def save_pretrained(self, save_directory):
self.config.save_pretrained(save_directory)
rag_tokenizer = RagTokenizer(
question_encoder_tokenizer=self.question_encoder_tokenizer,
generator_tokenizer=self.generator_tokenizer,
)
rag_tokenizer.save_pretrained(save_directory)
def init_retrieval(self):
"""
Retriever initalization function. It loads the index into memory.
"""
Retrirever initalization function, needs to be called from the training process. The function sets some common parameters
and environment variables. On top of that, (only) the main process in the process group loads the index into memory.
If this functin doesn't get called, we assume we're operating in a non-distributed environment and the index gets loaded
at first query.
logger.info("initializing retrieval")
self.index.init_index()
def postprocess_docs(self, docs, input_strings, prefix, n_docs, return_tensors=None):
r"""
Postprocessing retrieved ``docs`` and combining them with ``input_strings``.
Args:
doc_scores (:obj:`np.ndarray` of shape :obj:`(batch_size, n_docs)`):
Retrieval scores of respective docs - passed for logging.
docs (:obj:`dict`):
Retrieved documents.
input_strings (:obj:`str`):
Input strings decoded by ``preprocess_query``.
prefix (:obj:`str`):
Prefix added at the beginning of each input, typically used with T5-based models.
Return:
:obj:`tuple(tensors)`:
a tuple consisting of two elements: contextualized ``input_ids`` and a compatible ``attention_mask``.
"""
def cat_input_and_doc(doc_title, doc_text, input_string, prefix):
# TODO(Patrick): if we train more RAG models, I want to put the input first to take advantage of effortless truncation
# TODO(piktus): better handling of truncation
if doc_title.startswith('"'):
doc_title = doc_title[1:]
if doc_title.endswith('"'):
doc_title = doc_title[:-1]
if prefix is None:
prefix = ""
# TODO(Patrick, piktus, quention) with current master, T5Tokenizer should add eos token => so `add_eos` is not needed anymore
# suffix = self.generator_tokenizer.eos_token if add_eos else ""
out = (prefix + doc_title + self.config.title_sep + doc_text + self.config.doc_sep + input_string).replace(
" ", " "
)
return out
rag_input_strings = [
cat_input_and_doc(
docs[i]["title"][j],
docs[i]["text"][j],
input_strings[i],
prefix,
)
for i in range(len(docs))
for j in range(n_docs)
]
contextualized_inputs = self.generator_tokenizer.batch_encode_plus(
rag_input_strings,
max_length=self.config.max_combined_length,
return_tensors=return_tensors,
padding="max_length",
truncation=True,
)
return contextualized_inputs["input_ids"], contextualized_inputs["attention_mask"]
def _chunk_tensor(self, t: Iterable, chunk_size: int) -> List[Iterable]:
return [t[i : i + chunk_size] for i in range(0, len(t), chunk_size)]
def _main_retrieve(self, question_hidden_states: np.ndarray, n_docs: int) -> Tuple[np.ndarray, np.ndarray]:
question_hidden_states_batched = self._chunk_tensor(question_hidden_states, self.batch_size)
ids_batched = []
vectors_batched = []
for question_hidden_states in question_hidden_states_batched:
start_time = time.time()
ids, vectors = self.index.get_top_docs(question_hidden_states, n_docs)
logger.debug(
"index search time: {} sec, batch size {}".format(
time.time() - start_time, question_hidden_states.shape
)
)
ids_batched.extend(ids)
vectors_batched.extend(vectors)
return np.array(ids_batched), np.array(
vectors_batched
) # shapes (batch_size, n_docs) and (batch_size, n_docs, d)
def retrieve(self, question_hidden_states: np.ndarray, n_docs: int) -> Tuple[np.ndarray, List[dict]]:
"""
Retrieves documents for specified ``question_hidden_states``.
Args:
question_hidden_states (:obj:`np.ndarray` of shape :obj:`(batch_size, vector_size)`:
A batch of query vectors to retrieve with.
n_docs (:obj:`int`):
The number of docs retrieved per query.
Ouput:
retrieved_doc_embeds (:obj:`np.ndarray` of shape :obj:`(batch_size, n_docs, dim)`
The retrieval embeddings of the retrieved docs per query.
doc_ids (:obj:`np.ndarray` of shape :obj:`batch_size, n_docs`)
The ids of the documents in the index
doc_dicts (:obj:`List[dict]`):
The retrieved_doc_embeds examples per query.
"""
doc_ids, retrieved_doc_embeds = self._main_retrieve(question_hidden_states, n_docs)
return retrieved_doc_embeds, doc_ids, self.index.get_doc_dicts(doc_ids)
def __call__(
self,
question_input_ids: List[List[int]],
question_hidden_states: np.ndarray,
prefix: str,
n_docs=None,
return_tensors=None,
) -> BatchEncoding:
"""
Retrieves documents for specified ``question_hidden_states``.
Args:
question_input_ids: (:obj:`List[List[int]]`) batch of input ids
question_hidden_states (:obj:`np.ndarray` of shape :obj:`(batch_size, vector_size)`:
A batch of query vectors to retrieve with.
n_docs (:obj:`int`):
The number of docs retrieved per query.
Output:
:class:`~transformers.BatchEncoding`: A :class:`~transformers.BatchEncoding` with the following fields:
- **context_input_ids** -- List of token ids to be fed to a model.
`What are input IDs? <../glossary.html#input-ids>`__
- **context_attention_mask** -- List of indices specifying which tokens should be attended to by the model (when
:obj:`return_attention_mask=True` or if `"attention_mask"` is in :obj:`self.model_input_names`).
`What are attention masks? <../glossary.html#attention-mask>`__
- **retrieved_doc_embeds** -- List of embeddings of the retrieved documents
- **doc_ids** -- List of ids of the retrieved documents
"""
n_docs = n_docs if n_docs is not None else self.n_docs
retrieved_doc_embeds, doc_ids, docs = self.retrieve(question_hidden_states, n_docs)
input_strings = self.question_encoder_tokenizer.batch_decode(question_input_ids, skip_special_tokens=True)
context_input_ids, context_attention_mask = self.postprocess_docs(
docs, input_strings, prefix, n_docs, return_tensors=return_tensors
)
return BatchEncoding(
{
"context_input_ids": context_input_ids,
"context_attention_mask": context_attention_mask,
"retrieved_doc_embeds": retrieved_doc_embeds,
"doc_ids": doc_ids,
},
tensor_type=return_tensors,
)
class RagPyTorchDistributedRetriever(RagRetriever):
"""
A distributed retriever built on top of the ``torch.distributed`` communication package. During training all workers
initalize their own instance of the retriever, however, only the main worker loads the index into memory. The index is stored
in cpu memory. The index will also work well in a non-distributed setup.
Args:
config (:class:`~transformers.RagConfig`):
The configuration of the RAG model this Retriever is used with. Contains parameters indicating which ``Index`` to build.
distributed_port (:obj:`int`):
The port on which the main communication of the training run is carried out. We set the port for retrieval-related
communication as ``distributed_port + 1``.
"""
_init_retrieval = False
def __init__(self, config, question_encoder_tokenizer, generator_tokenizer):
super().__init__(
config, question_encoder_tokenizer=question_encoder_tokenizer, generator_tokenizer=generator_tokenizer
)
self.process_group = None
def init_retrieval(self, distributed_port: int):
"""
Retriever initalization function, needs to be called from the training process. The function sets some common parameters
and environment variables. On top of that, (only) the main process in the process group loads the index into memory.
Args:
distributed_port (:obj:`int`):
@@ -310,120 +515,15 @@ class RagRetriever(object):
# initialize retriever only on the main worker
if not dist.is_initialized() or self._is_main():
logger.info("dist not initialized / main")
self.retriever.init_index()
self.index.init_index()
# all processes wait untill the retriever is initialized by the main process
if dist.is_initialized():
torch.distributed.barrier(group=self.process_group)
def preprocess_query(self, input_ids, prefix):
r"""
Preprocesses the ``input_id`` by first converting it to string using the ``generator_tokenizer`` and
then tokenizing it using the ``question_encoder_tokenizer``.
Args:
input_ids (:obj:`torch.LongTensor` of shape :obj:`(batch_size, sequence_length)`):
Indices of input sequence tokens in the vocabulary.
Return:
:obj:`torch.LongTensor`:
Tokenized input.
:obj:`str`:
Decoded input strings.
"""
input_strings = self.generator_tokenizer.batch_decode(input_ids, skip_special_tokens=True)
# handle prefix for T5
if isinstance(self.generator_tokenizer, T5Tokenizer):
for i, s in enumerate(input_strings):
if not s.startswith(prefix):
logger.warning("T5 prefix mismatch in {}".format(s))
if len(input_strings[i]) <= len(prefix):
input_strings[i] = ""
else:
input_strings[i] = input_strings[i][len(prefix) :]
retriever_inputs = self.question_encoder_tokenizer.batch_encode_plus(
input_strings,
return_tensors="pt",
padding=True,
truncation=True,
)
return retriever_inputs["input_ids"].to(input_ids.device), input_strings
def postprocess_docs(self, doc_scores, docs, input_strings, add_eos, prefix, print_docs=False):
r"""
Postprocessing retrieved ``docs`` and combining them with ``input_strings``.
Args:
doc_scores (:obj:`torch.FloatTensor` of shape :obj:`(batch_size, n_docs)`):
Retrieval scores of respective docs - passed for logging.
docs (:obj:`dict`):
Retrieved documents.
input_strings (:obj:`str`):
Input strings decoded by ``preprocess_query``.
add_eos (:obj:`bool`):
A boolean flag signalling that eos token needs to be added to the contextualized input.
prefix (:obj:`str`):
Prefix added at the beginning of each input, typically used with T5-based models.
print_docs (:obj:`bool`, `optional`, defaults to :obj:`False`):
If :obj:`True`, documents retrieved during the forward pass will be printed out. Intended for debugging purposes.
Return:
:obj:`tuple(tuple(torch.FloatTensor)`:
a tuple consisting od two elements: contextualized ``input_ids`` and a compatible ``attention_mask``.
"""
def cat_input_and_doc(doc_score, doc_title, doc_text, input_string, add_eos, prefix, print_docs=False):
# TODO(Patrick): if we train more RAG models, I want to put the input first to take advantage of effortless truncation
# TODO(piktus): better handling of truncation
if doc_title.startswith('"'):
doc_title = doc_title[1:]
if doc_title.endswith('"'):
doc_title = doc_title[:-1]
if prefix is None:
prefix = ""
suffix = self.generator_tokenizer.eos_token if add_eos else ""
out = (
prefix + doc_title + self.config.title_sep + doc_text + self.config.doc_sep + input_string + suffix
).replace(" ", " ")
if print_docs:
logger.info("{} {}".format(doc_score, out))
return out
rag_input_strings = [
cat_input_and_doc(
doc_scores[i][j],
docs[i]["title"][j],
docs[i]["text"][j],
input_strings[i],
add_eos,
prefix,
print_docs,
)
for i in range(len(docs))
for j in range(self.n_docs)
]
contextualized_inputs = self.generator_tokenizer.batch_encode_plus(
rag_input_strings,
max_length=self.config.max_combined_length,
return_tensors="pt",
padding="max_length",
truncation=True,
).to(doc_scores.device)
return contextualized_inputs["input_ids"], contextualized_inputs["attention_mask"]
def _is_main(self):
return dist.get_rank(group=self.process_group) == 0
def _chunk_tensor(self, t, chunk_size):
n_chunks = t.shape[0] // chunk_size + int(t.shape[0] % chunk_size > 0)
return list(torch.chunk(t, n_chunks, dim=0))
def _scattered(self, scatter_list, target_shape, target_type=torch.float32):
target_tensor = torch.empty(target_shape, dtype=target_type)
dist.scatter(target_tensor, src=0, scatter_list=scatter_list, group=self.process_group)
@@ -435,48 +535,30 @@ class RagRetriever(object):
ifname = next((addr for addr in addrs if addr.startswith("e")), None)
return ifname
def _main_retrieve(self, query_vectors):
query_vectors_batched = self._chunk_tensor(query_vectors, self.batch_size)
ids_batched = []
vectors_batched = []
for query_vectors in query_vectors_batched:
start_time = time.time()
ids, vectors = self.retriever.get_top_docs(query_vectors.numpy(), self.n_docs)
logger.debug(
"index search time: {} sec, batch size {}".format(time.time() - start_time, query_vectors.shape)
)
ids_batched.append(ids)
vectors_batched.append(vectors)
return torch.cat(ids_batched), torch.cat(vectors_batched)
def retrieve(self, query_vectors, n_docs):
def retrieve(self, question_hidden_states: np.ndarray, n_docs: int) -> Tuple[np.ndarray, List[dict]]:
"""
Retrieves documents for specified ``query_vectors``. The main process, which has the access to the index stored in memory, gathers queries
Retrieves documents for specified ``question_hidden_states``. The main process, which has the access to the index stored in memory, gathers queries
from all the processes in the main training process group, performs the retrieval and scatters back the results.
Args:
query_vectors (:obj:`torch.Tensor` of shape :obj:`(batch_size, vector_size)`:
question_hidden_states (:obj:`np.ndarray` of shape :obj:`(batch_size, vector_size)`:
A batch of query vectors to retrieve with.
n_docs (:obj:`int`):
The number of docs retrieved per query.
Ouput:
total_scores (:obj:`torch.Tensor` of shape :obj:`(batch_size, n_docs)`
The retrieval scores of the retrieved docs per query.
total_examples (:obj:`List[dict]`):
The retrieved examples per query.
retrieved_doc_embeds (:obj:`np.ndarray` of shape :obj:`(batch_size, n_docs, dim)`
The retrieval embeddings of the retrieved docs per query.
doc_ids (:obj:`np.ndarray` of shape :obj:`batch_size, n_docs`)
The ids of the documents in the index
doc_dicts (:obj:`List[dict]`):
The retrieved_doc_embeds examples per query.
"""
# non-ddp initialization (init_retrieval() is called at ddp initialization, if no ddp, then it's never called,
# so it has to be initalized separately.
if not dist.is_initialized() and not self.retriever.is_initialized():
logger.info("Initializing index at first query")
self.retriever.init_index()
# single GPU training
if not dist.is_initialized():
doc_ids, doc_vectors = self._main_retrieve(query_vectors)
return doc_vectors, self.retriever.get_doc_dicts(doc_ids)
doc_ids, retrieved_doc_embeds = self._main_retrieve(question_hidden_states, n_docs)
return retrieved_doc_embeds, doc_ids, self.index.get_doc_dicts(doc_ids)
# distributed training
world_size = dist.get_world_size(group=self.process_group)
@@ -484,19 +566,20 @@ class RagRetriever(object):
# gather logic
gather_list = None
if self._is_main():
gather_list = [torch.empty(query_vectors.shape, dtype=torch.float32) for _ in range(world_size)]
dist.gather(query_vectors, dst=0, gather_list=gather_list, group=self.process_group)
gather_list = [torch.empty(question_hidden_states.shape, dtype=torch.float32) for _ in range(world_size)]
dist.gather(torch.tensor(question_hidden_states), dst=0, gather_list=gather_list, group=self.process_group)
# scatter logic
n_queries = query_vectors.shape[0]
n_queries = question_hidden_states.shape[0]
scatter_ids = []
scatter_vectors = []
if self._is_main():
assert len(gather_list) == world_size
ids, vectors = self._main_retrieve(torch.cat(gather_list))
ids, vectors = self._main_retrieve(torch.cat(gather_list).numpy(), n_docs)
ids, vectors = torch.tensor(ids), torch.tensor(vectors)
scatter_ids = self._chunk_tensor(ids, n_queries)
scatter_vectors = self._chunk_tensor(vectors, n_queries)
doc_ids = self._scattered(scatter_ids, [n_queries, self.n_docs], target_type=torch.int64)
doc_vectors = self._scattered(scatter_vectors, [n_queries, self.n_docs, query_vectors.shape[1]])
doc_ids = self._scattered(scatter_ids, [n_queries, n_docs], target_type=torch.int64)
retrieved_doc_embeds = self._scattered(scatter_vectors, [n_queries, n_docs, question_hidden_states.shape[1]])
return doc_vectors, self.retriever.get_doc_dicts(doc_ids)
return retrieved_doc_embeds.numpy(), doc_ids.numpy(), self.index.get_doc_dicts(doc_ids)
+25 -1
View File
@@ -10,7 +10,7 @@ from distutils.util import strtobool
from io import StringIO
from pathlib import Path
from .file_utils import _tf_available, _torch_available, _torch_tpu_available
from .file_utils import _datasets_available, _faiss_available, _tf_available, _torch_available, _torch_tpu_available
SMALL_MODEL_IDENTIFIER = "julien-c/bert-xsmall-dummy"
@@ -161,6 +161,30 @@ def require_torch_and_cuda(test_case):
return test_case
def require_datasets(test_case):
"""
Decorator marking a test that requires TensorFlow.
These tests are skipped when TensorFlow isn't installed.
"""
if not _datasets_available:
test_case = unittest.skip("test requires Datasets")(test_case)
return test_case
def require_faiss(test_case):
"""
Decorator marking a test that requires TensorFlow.
These tests are skipped when TensorFlow isn't installed.
"""
if not _faiss_available:
test_case = unittest.skip("test requires Faiss")(test_case)
return test_case
def get_tests_dir():
"""
returns the full path to the `tests` dir, so that the tests can be invoked from anywhere
+3
View File
@@ -26,6 +26,7 @@ from .configuration_auto import (
CamembertConfig,
CTRLConfig,
DistilBertConfig,
DPRConfig,
ElectraConfig,
EncoderDecoderConfig,
FlaubertConfig,
@@ -57,6 +58,7 @@ from .tokenization_bert_japanese import BertJapaneseTokenizer
from .tokenization_camembert import CamembertTokenizer
from .tokenization_ctrl import CTRLTokenizer
from .tokenization_distilbert import DistilBertTokenizer, DistilBertTokenizerFast
from .tokenization_dpr import DPRQuestionEncoderTokenizer, DPRQuestionEncoderTokenizerFast
from .tokenization_electra import ElectraTokenizer, ElectraTokenizerFast
from .tokenization_flaubert import FlaubertTokenizer
from .tokenization_funnel import FunnelTokenizer, FunnelTokenizerFast
@@ -110,6 +112,7 @@ TOKENIZER_MAPPING = OrderedDict(
(XLMConfig, (XLMTokenizer, None)),
(CTRLConfig, (CTRLTokenizer, None)),
(BertGenerationConfig, (BertGenerationTokenizer, None)),
(DPRConfig, (DPRQuestionEncoderTokenizer, DPRQuestionEncoderTokenizerFast)),
]
)
+34 -40
View File
@@ -13,54 +13,48 @@
# See the License for the specific language governing permissions and
# limitations under the License.
"""Tokenization classes for RAG."""
import os
from .configuration_rag import RagConfig
from .tokenization_auto import AutoTokenizer
from .utils import logging
from .tokenization_bart import BartTokenizer, BartTokenizerFast
logger = logging.get_logger(__name__)
VOCAB_FILES_NAMES = {
"vocab_file": "vocab.json",
"merges_file": "merges.txt",
}
class RagTokenizer:
def __init__(self, question_encoder, generator):
self.question_encoder = question_encoder
self.generator = generator
def save_pretrained(self, save_directory):
if os.path.isfile(save_directory):
logger.error("Provided path ({}) should be a directory, not a file".format(save_directory))
return
os.makedirs(save_directory, exist_ok=True)
question_encoder_path = os.path.join(save_directory, "question_encoder")
generator_path = os.path.join(save_directory, "generator")
self.question_encoder.save_pretrained(question_encoder_path)
self.generator.save_pretrained(generator_path)
RAG_PRETRAINED_VOCAB_FILES_MAP = {
"vocab_file": {
"facebook/rag-sequence-nq": "https://s3.amazonaws.com/models.huggingface.co/bert/roberta-large-vocab.json",
"facebook/rag-token-nq": "https://s3.amazonaws.com/models.huggingface.co/bert/roberta-large-vocab.json",
},
"merges_file": {
"facebook/rag-sequence-nq": "https://s3.amazonaws.com/models.huggingface.co/bert/roberta-large-merges.txt",
"facebook/rag-token-nq": "https://s3.amazonaws.com/models.huggingface.co/bert/roberta-large-merges.txt",
},
}
@classmethod
def from_pretrained(cls, pretrained_model_name_or_path, **kwargs):
config = kwargs.pop("config", None)
if config is None:
config = RagConfig.from_pretrained(pretrained_model_name_or_path)
class RagDefaultTokenizer(BartTokenizer):
r"""
Constructs a RagDefaultTokenizer.
question_encoder_path = os.path.join(pretrained_model_name_or_path, "question_encoder_tokenizer")
generator_path = os.path.join(pretrained_model_name_or_path, "generator_tokenizer")
question_encoder = AutoTokenizer.from_pretrained(question_encoder_path, config=config.question_encoder)
generator = AutoTokenizer.from_pretrained(generator_path, config=config.generator)
return cls(question_encoder=question_encoder, generator=generator)
:class:`~transformers.RagDefaultTokenizer` is identical to :class:`~transformers.BertTokenizer` and runs end-to-end
tokenization: punctuation splitting + wordpiece.
def __call__(self, *args, **kwargs):
return self.question_encoder(*args, **kwargs)
Refer to superclass :class:`~transformers.BertTokenizer` for usage examples and documentation concerning
parameters.
"""
def batch_decode(self, *args, **kwargs):
return self.generator.batch_decode(*args, **kwargs)
vocab_files_names = VOCAB_FILES_NAMES
pretrained_vocab_files_map = RAG_PRETRAINED_VOCAB_FILES_MAP
class RagDefaultTokenizerFast(BartTokenizerFast):
r"""
Constructs a RagDefaultTokenizerFast.
:class:`~transformers.RagDefaultTokenizerFast` is identical to :class:`~transformers.BertTokenizer` and runs end-to-end
tokenization: punctuation splitting + wordpiece.
Refer to superclass :class:`~transformers.BertTokenizer` for usage examples and documentation concerning
parameters.
"""
vocab_files_names = VOCAB_FILES_NAMES
pretrained_vocab_files_map = RAG_PRETRAINED_VOCAB_FILES_MAP
# TODO(Patrick) add prepare_seq2seq_batch function
+567 -216
View File
@@ -14,33 +14,53 @@
# limitations under the License.
import copy
import json
import os
import shutil
import tempfile
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, slow, torch_device
import numpy as np
from .test_configuration_common import ConfigTester
from transformers.file_utils import (
cached_property,
is_datasets_available,
is_faiss_available,
is_psutil_available,
is_torch_available,
)
from transformers.modeling_outputs import BaseModelOutput
from transformers.testing_utils import require_torch, slow, torch_device
from transformers.tokenization_bart import BartTokenizer
from transformers.tokenization_bert import VOCAB_FILES_NAMES as DPR_VOCAB_FILES_NAMES
from transformers.tokenization_dpr import DPRQuestionEncoderTokenizer
from transformers.tokenization_roberta import VOCAB_FILES_NAMES as BART_VOCAB_FILES_NAMES
from .test_modeling_bart import ModelTester as BartModelTester
from .test_modeling_common import ids_tensor
from .test_modeling_dpr import DPRModelTester
TOLERANCE = 1e-4
if is_torch_available() and is_datasets_available() and is_faiss_available() and is_psutil_available():
import faiss
import torch
from datasets import Dataset
from transformers import (
AutoConfig,
AutoModel,
AutoModelForSeq2SeqLM,
BartConfig,
BartForConditionalGeneration,
BartTokenizer,
DPRConfig,
DPRQuestionEncoder,
RagConfig,
RagModel,
RagRetriever,
RagSequence,
RagToken,
RagSequenceForGeneration,
RagTokenForGeneration,
)
@@ -130,7 +150,6 @@ class RagModelTester:
bos_token_id=self.bos_token_id,
pad_token_id=self.pad_token_id,
decoder_start_token_id=self.decoder_start_token_id,
return_dict=True,
)
# DPR params
@@ -161,7 +180,6 @@ class RagModelTester:
type_vocab_size=self.type_vocab_size,
is_decoder=False,
initializer_range=self.initializer_range,
return_dict=True,
)
def prepare_inputs(self):
@@ -176,215 +194,333 @@ class RagModelTester:
@require_torch
@require_retrieval
class RagModelTest(unittest.TestCase):
class RagTestMixin:
all_model_classes = (
(RagSequence, RagToken)
(RagModel, RagTokenForGeneration, RagSequenceForGeneration)
if is_torch_available() and is_datasets_available() and is_faiss_available() and is_psutil_available()
else ()
)
retrieval_vector_size = 32
n_docs = 2
max_combined_length = 16
def setUp(self):
self.model_tester = RagModelTester(self)
self.config_tester = ConfigTester(self, config_class=RagConfig, hidden_size=37)
self.tmpdirname = tempfile.mkdtemp()
def test_config(self):
self.config_tester.create_and_test_config_to_json_string()
self.config_tester.create_and_test_config_to_json_file()
self.config_tester.create_and_test_config_from_and_save_pretrained()
self.config_tester.create_and_test_config_with_num_labels()
# DPR tok
vocab_tokens = [
"[UNK]",
"[CLS]",
"[SEP]",
"[PAD]",
"[MASK]",
"want",
"##want",
"##ed",
"wa",
"un",
"runn",
"##ing",
",",
"low",
"lowest",
]
dpr_tokenizer_path = os.path.join(self.tmpdirname, "dpr_tokenizer")
os.makedirs(dpr_tokenizer_path, exist_ok=True)
self.vocab_file = os.path.join(dpr_tokenizer_path, DPR_VOCAB_FILES_NAMES["vocab_file"])
with open(self.vocab_file, "w", encoding="utf-8") as vocab_writer:
vocab_writer.write("".join([x + "\n" for x in vocab_tokens]))
def test_constructor_from_config(self):
for model_class in self.all_model_classes:
model = model_class(config=self.model_tester.rag_config)
self.assertEqual(model.n_docs, self.model_tester.rag_config.n_docs)
self.assertEqual(model.model.n_docs, self.model_tester.rag_config.n_docs)
self.assertIsNotNone(model.model)
self.assertIsNotNone(model.model.question_encoder)
self.assertIsNotNone(model.model.generator)
self.assertTrue(model.config.is_encoder_decoder)
# BART tok
vocab = [
"l",
"o",
"w",
"e",
"r",
"s",
"t",
"i",
"d",
"n",
"\u0120",
"\u0120l",
"\u0120n",
"\u0120lo",
"\u0120low",
"er",
"\u0120lowest",
"\u0120newer",
"\u0120wider",
"<unk>",
]
vocab_tokens = dict(zip(vocab, range(len(vocab))))
merges = ["#version: 0.2", "\u0120 l", "\u0120l o", "\u0120lo w", "e r", ""]
self.special_tokens_map = {"unk_token": "<unk>"}
def test_constructor_from_object(self):
for model_class in self.all_model_classes:
model = model_class(
config=self.model_tester.rag_config,
question_encoder=DPRQuestionEncoder(self.model_tester.dpr_config),
generator=BartForConditionalGeneration(self.model_tester.bart_config),
bart_tokenizer_path = os.path.join(self.tmpdirname, "bart_tokenizer")
os.makedirs(bart_tokenizer_path, exist_ok=True)
self.vocab_file = os.path.join(bart_tokenizer_path, BART_VOCAB_FILES_NAMES["vocab_file"])
self.merges_file = os.path.join(bart_tokenizer_path, BART_VOCAB_FILES_NAMES["merges_file"])
with open(self.vocab_file, "w", encoding="utf-8") as fp:
fp.write(json.dumps(vocab_tokens) + "\n")
with open(self.merges_file, "w", encoding="utf-8") as fp:
fp.write("\n".join(merges))
@cached_property
def dpr_tokenizer(self) -> DPRQuestionEncoderTokenizer:
return DPRQuestionEncoderTokenizer.from_pretrained(os.path.join(self.tmpdirname, "dpr_tokenizer"))
@cached_property
def bart_tokenizer(self) -> BartTokenizer:
return BartTokenizer.from_pretrained(os.path.join(self.tmpdirname, "bart_tokenizer"))
def tearDown(self):
shutil.rmtree(self.tmpdirname)
def get_retriever(self, config):
dataset = Dataset.from_dict(
{
"id": ["0", "1"],
"text": ["foo", "bar"],
"title": ["Foo", "Bar"],
"embeddings": [np.ones(self.retrieval_vector_size), 2 * np.ones(self.retrieval_vector_size)],
}
)
dataset.add_faiss_index("embeddings", string_factory="Flat", metric_type=faiss.METRIC_INNER_PRODUCT)
with patch("transformers.retrieval_rag.load_dataset") as mock_load_dataset:
mock_load_dataset.return_value = dataset
retriever = RagRetriever(
config,
question_encoder_tokenizer=self.dpr_tokenizer,
generator_tokenizer=self.bart_tokenizer,
)
self.assertEqual(model.n_docs, self.model_tester.rag_config.n_docs)
self.assertEqual(model.model.n_docs, self.model_tester.rag_config.n_docs)
self.assertIsNotNone(model.model)
self.assertIsNotNone(model.model.question_encoder)
self.assertIsNotNone(model.model.generator)
self.assertTrue(model.config.is_encoder_decoder)
return retriever
def check_model_with_retriever(
self, config, input_ids, attention_mask, decoder_input_ids, decoder_attention_mask, **kwargs
):
self.assertIsNotNone(config.question_encoder)
self.assertIsNotNone(config.generator)
def test_constructor_from_pretrained(self):
for model_class in self.all_model_classes:
model = model_class.from_pretrained(config=self.model_tester.rag_config)
self.assertEqual(model.n_docs, self.model_tester.rag_config.n_docs)
self.assertEqual(model.model.n_docs, self.model_tester.rag_config.n_docs)
self.assertIsNotNone(model.model)
self.assertIsNotNone(model.model.question_encoder)
self.assertIsNotNone(model.model.generator)
self.assertTrue(model.config.is_encoder_decoder)
def test_constructor_mismatch(self):
mismatched_bart_config = copy.deepcopy(self.model_tester.bart_config)
def test_mismatch():
for model_class in self.all_model_classes:
with self.assertRaises(
AssertionError,
):
model_class(
config=self.model_tester.rag_config,
question_encoder=DPRQuestionEncoder(self.model_tester.dpr_config),
generator=BartForConditionalGeneration(mismatched_bart_config),
)
mismatched_bart_config.eos_token_id = self.model_tester.bart_config.eos_token_id + 1
test_mismatch()
mismatched_bart_config.eos_token_id = self.model_tester.bart_config.eos_token_id
mismatched_bart_config.bos_token_id = self.model_tester.bart_config.bos_token_id + 1
test_mismatch()
mismatched_bart_config.bos_token_id = self.model_tester.bart_config.bos_token_id
mismatched_bart_config.pad_token_id = self.model_tester.bart_config.pad_token_id + 1
test_mismatch()
mismatched_bart_config.pad_token_id = self.model_tester.bart_config.pad_token_id
mismatched_bart_config.decoder_start_token_id = self.model_tester.bart_config.decoder_start_token_id + 1
test_mismatch()
mismatched_bart_config.decoder_start_token_id = self.model_tester.bart_config.decoder_start_token_id
mismatched_bart_config.is_encoder_decoder = not self.model_tester.bart_config.is_encoder_decoder
test_mismatch()
mismatched_bart_config.is_encoder_decoder = self.model_tester.bart_config.is_encoder_decoder
mismatched_bart_config.vocab_size = not self.model_tester.bart_config.vocab_size + 1
test_mismatch()
def mock_contextualize(*args, **kwargs):
input_ids = torch.tensor([[0, 31414, 232, 328, 2]] * 3 * 13)
attention_mask = torch.tensor([[1, 1, 1, 1, 1]] * 3 * 13)
doc_scores = torch.tensor([[0.111, 0.222, 0.333]] * 13)
return input_ids, attention_mask, doc_scores
@patch("transformers.RagModel.contextualize", mock_contextualize)
def test_forward_pass(self):
input_ids, attention_mask = self.model_tester.prepare_inputs()
decoder_input_ids = torch.tensor([[0, 31414, 232, 328, 2]] * self.model_tester.batch_size)
tgt_len = decoder_input_ids.shape[1]
for model_class in self.all_model_classes:
model = model_class(
config=self.model_tester.rag_config,
question_encoder=DPRQuestionEncoder(self.model_tester.dpr_config),
generator=BartForConditionalGeneration(self.model_tester.bart_config),
)
model.to(torch_device)
model = model_class(config, retriever=self.get_retriever(config)).to(torch_device)
model.eval()
# use cache
result = model(
input_ids,
retriever=None,
decoder_input_ids=decoder_input_ids,
attention_mask=attention_mask,
marginalize=False,
use_cache=True,
)
self.assertEqual(
result.logits.shape,
(self.model_tester.rag_config.n_docs * self.model_tester.batch_size, 1, self.model_tester.vocab_size),
)
self.assertEqual(
result.doc_scores.shape, (self.model_tester.batch_size, self.model_tester.rag_config.n_docs)
)
self.assertIsNone(result.loss)
self.assertTrue(model.config.is_encoder_decoder)
# no cache
result = model(
input_ids,
retriever=None,
decoder_input_ids=decoder_input_ids,
outputs = model(
input_ids=input_ids,
attention_mask=attention_mask,
marginalize=False,
use_cache=False,
)
self.assertEqual(
result.logits.shape,
(
self.model_tester.rag_config.n_docs * self.model_tester.batch_size,
tgt_len,
self.model_tester.vocab_size,
),
)
self.assertEqual(
result.doc_scores.shape, (self.model_tester.batch_size, self.model_tester.rag_config.n_docs)
)
self.assertIsNone(result.loss)
# marginalization in RagToken + no cache
if isinstance(model_class, RagToken):
result = model(
input_ids,
decoder_input_ids=decoder_input_ids,
attention_mask=attention_mask,
marginalize=True,
use_cache=False,
)
self.assertEqual(
result.logits.shape, (self.model_tester.batch_size, tgt_len, self.model_tester.vocab_size)
)
self.assertEqual(
result.doc_scores.shape, (self.model_tester.batch_size, self.model_tester.rag_config.n_docs)
)
self.assertIsNone(result.loss)
# return_loss, no reduce
result = model(
input_ids,
retriever=None,
decoder_input_ids=decoder_input_ids,
attention_mask=attention_mask,
return_loss=True,
decoder_attention_mask=decoder_attention_mask,
)
self.assertEqual(
result.logits.shape,
(
self.model_tester.rag_config.n_docs * self.model_tester.batch_size,
tgt_len,
self.model_tester.vocab_size,
),
)
self.assertEqual(
result.doc_scores.shape, (self.model_tester.batch_size, self.model_tester.rag_config.n_docs)
)
self.assertEqual(result.loss.shape, (self.model_tester.batch_size,))
# return_loss, reduce
result = model(
# logits
self.assertEqual(
outputs.logits.shape,
(self.n_docs * decoder_input_ids.shape[0], decoder_input_ids.shape[1], config.generator.vocab_size),
)
# generator encoder last hidden states
self.assertEqual(
outputs.generator_enc_last_hidden_state.shape,
(self.n_docs * decoder_input_ids.shape[0], self.max_combined_length, config.generator.hidden_size),
)
# doc scores
self.assertEqual(outputs.doc_scores.shape, (input_ids.shape[0], self.n_docs))
def check_model_generate(
self, config, input_ids, attention_mask, decoder_input_ids, decoder_attention_mask, **kwargs
):
self.assertIsNotNone(config.question_encoder)
self.assertIsNotNone(config.generator)
for model_class in self.all_model_classes[1:]:
model = model_class(config, retriever=self.get_retriever(config)).to(torch_device)
model.eval()
self.assertTrue(model.config.is_encoder_decoder)
outputs = model.generate(
input_ids=input_ids,
num_beams=2,
num_return_sequences=2,
decoder_start_token_id=config.generator.eos_token_id,
)
def check_model_without_retriever(
self, config, input_ids, attention_mask, decoder_input_ids, decoder_attention_mask, **kwargs
):
self.assertIsNotNone(config.question_encoder)
self.assertIsNotNone(config.generator)
retriever = self.get_retriever(config)
for model_class in self.all_model_classes:
model = model_class(config).to(torch_device)
model.eval()
self.assertTrue(model.config.is_encoder_decoder)
question_hidden_states = model.question_encoder(input_ids, attention_mask=attention_mask)[0]
out = retriever(
input_ids,
retriever=None,
question_hidden_states.cpu().detach().to(torch.float32).numpy(),
prefix=config.generator.prefix,
return_tensors="pt",
)
context_input_ids, context_attention_mask, retrieved_doc_embeds = (
out["context_input_ids"],
out["context_attention_mask"],
out["retrieved_doc_embeds"],
)
# cast
retrieved_doc_embeds = retrieved_doc_embeds.to(question_hidden_states)
context_input_ids = context_input_ids.to(input_ids)
context_attention_mask = context_attention_mask.to(input_ids)
# compute doc_scores
doc_scores = torch.bmm(question_hidden_states.unsqueeze(1), retrieved_doc_embeds.transpose(1, 2)).squeeze(
1
)
outputs = model(
context_input_ids=context_input_ids,
context_attention_mask=context_attention_mask,
doc_scores=doc_scores,
decoder_input_ids=decoder_input_ids,
decoder_attention_mask=decoder_attention_mask,
)
# logits
self.assertEqual(
outputs.logits.shape,
(self.n_docs * decoder_input_ids.shape[0], decoder_input_ids.shape[1], config.generator.vocab_size),
)
# generator encoder last hidden states
self.assertEqual(
outputs.generator_enc_last_hidden_state.shape,
(self.n_docs * decoder_input_ids.shape[0], self.max_combined_length, config.generator.hidden_size),
)
# doc scores
self.assertEqual(outputs.doc_scores.shape, (input_ids.shape[0], self.n_docs))
def check_model_with_encoder_outputs(
self, config, input_ids, attention_mask, decoder_input_ids, decoder_attention_mask, **kwargs
):
self.assertIsNotNone(config.question_encoder)
self.assertIsNotNone(config.generator)
for model_class in self.all_model_classes:
model = model_class(config, retriever=self.get_retriever(config)).to(torch_device)
model.eval()
self.assertTrue(model.config.is_encoder_decoder)
outputs = model(
input_ids=input_ids,
attention_mask=attention_mask,
return_loss=True,
reduce=True,
decoder_input_ids=decoder_input_ids,
decoder_attention_mask=decoder_attention_mask,
)
encoder_outputs = BaseModelOutput(outputs.generator_enc_last_hidden_state)
# run only generator
outputs = model(
encoder_outputs=encoder_outputs,
doc_scores=outputs.doc_scores,
decoder_input_ids=decoder_input_ids,
decoder_attention_mask=decoder_attention_mask,
)
# logits
self.assertEqual(
result.logits.shape,
(
self.model_tester.rag_config.n_docs * self.model_tester.batch_size,
tgt_len,
self.model_tester.vocab_size,
),
outputs.logits.shape,
(self.n_docs * decoder_input_ids.shape[0], decoder_input_ids.shape[1], config.generator.vocab_size),
)
# generator encoder last hidden states
self.assertEqual(
result.doc_scores.shape, (self.model_tester.batch_size, self.model_tester.rag_config.n_docs)
outputs.generator_enc_last_hidden_state.shape,
(self.n_docs * decoder_input_ids.shape[0], self.max_combined_length, config.generator.hidden_size),
)
self.assertEqual(result.loss.shape, torch.Size([]))
# doc scores
self.assertEqual(outputs.doc_scores.shape, (input_ids.shape[0], self.n_docs))
def test_model_with_retriever(self):
inputs_dict = self.config_and_inputs
self.check_model_with_retriever(**inputs_dict)
def test_model_without_retriever(self):
inputs_dict = self.config_and_inputs
self.check_model_without_retriever(**inputs_dict)
def test_model_with_encoder_outputs(self):
inputs_dict = self.config_and_inputs
self.check_model_with_encoder_outputs(**inputs_dict)
def test_model_generate(self):
inputs_dict = self.config_and_inputs
self.check_model_generate(**inputs_dict)
@require_torch
@require_retrieval
class RagDPRBartTest(RagTestMixin, unittest.TestCase):
@cached_property
def config_and_inputs(self):
question_encoder_tester = DPRModelTester(self)
dpr_config_and_inputs = question_encoder_tester.prepare_config_and_inputs()
generator_tester = BartModelTester(self)
bart_config_and_inputs = generator_tester.prepare_config_and_inputs_for_common()
(question_encoder_config, input_ids, _, input_mask, _, _, _) = dpr_config_and_inputs
(generator_config, bart_inputs_dict) = bart_config_and_inputs
decoder_input_ids, decoder_attention_mask = bart_inputs_dict["input_ids"], bart_inputs_dict["attention_mask"]
config = RagConfig.from_question_encoder_generator_configs(
question_encoder_config,
generator_config,
n_docs=self.n_docs,
retrieval_vector_size=self.retrieval_vector_size,
max_combined_length=self.max_combined_length,
use_cache=False,
)
return {
"config": config,
"input_ids": input_ids,
"attention_mask": input_mask,
"decoder_input_ids": decoder_input_ids,
"decoder_attention_mask": decoder_attention_mask,
}
@require_torch
@require_retrieval
class RagModelIntegrationTests(unittest.TestCase):
@cached_property
def sequence_model(self):
return RagSequenceForGeneration.from_pretrained_question_encoder_generator(
"facebook/dpr-question_encoder-single-nq-base", "facebook/bart-large-cnn"
).to(torch_device)
@cached_property
def token_model(self):
return RagTokenForGeneration.from_pretrained_question_encoder_generator(
"facebook/dpr-question_encoder-single-nq-base", "facebook/bart-large-cnn"
).to(torch_device)
def get_rag_config(self):
return RagConfig(
question_encoder_config = AutoConfig.from_pretrained("facebook/dpr-question_encoder-single-nq-base")
generator_config = AutoConfig.from_pretrained("facebook/bart-large-cnn")
return RagConfig.from_question_encoder_generator_configs(
question_encoder_config,
generator_config,
bos_token_id=0,
decoder_start_token_id=2,
eos_token_id=2,
@@ -395,78 +531,293 @@ class RagModelIntegrationTests(unittest.TestCase):
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,
use_dummy_dataset=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)
rag_tokenizer = BartTokenizer.from_pretrained("facebook/bart-large-cnn")
rag_decoder_tokenizer = BartTokenizer.from_pretrained("facebook/bart-large-cnn")
rag_question_encoder_tokenizer = DPRQuestionEncoderTokenizer.from_pretrained(
"facebook/dpr-question_encoder-single-nq-base"
)
rag_retriever = RagRetriever(
rag_config,
question_encoder_tokenizer=rag_question_encoder_tokenizer,
generator_tokenizer=rag_decoder_tokenizer,
)
input_ids = rag_tokenizer("who sings does he love me with reba", return_tensors="pt").input_ids
decoder_input_ids = rag_tokenizer("Linda Davis", return_tensors="pt").input_ids
rag_sequence = self.sequence_model
rag_sequence.set_retriever(rag_retriever)
input_ids = rag_question_encoder_tokenizer(
"who sings does he love me with reba", return_tensors="pt"
).input_ids
decoder_input_ids = rag_decoder_tokenizer("Linda Davis", return_tensors="pt").input_ids
input_ids = input_ids.to(torch_device)
decoder_input_ids = decoder_input_ids.to(torch_device)
rag_sequence = RagSequence.from_pretrained(config=rag_config).to(torch_device)
with torch.no_grad():
output = rag_sequence(
input_ids,
retriever=rag_retriever,
decoder_input_ids=decoder_input_ids,
return_loss=True,
print_docs=True,
labels=decoder_input_ids,
)
expected_shape = torch.Size([5, 5, 50264])
self.assertEqual(output.logits.shape, expected_shape)
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)
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)
rag_tokenizer = BartTokenizer.from_pretrained("facebook/bart-large-cnn")
rag_decoder_tokenizer = BartTokenizer.from_pretrained("facebook/bart-large-cnn")
rag_question_encoder_tokenizer = DPRQuestionEncoderTokenizer.from_pretrained(
"facebook/dpr-question_encoder-single-nq-base"
)
rag_retriever = RagRetriever(
rag_config,
question_encoder_tokenizer=rag_question_encoder_tokenizer,
generator_tokenizer=rag_decoder_tokenizer,
)
input_ids = rag_tokenizer("who sings does he love me with reba", return_tensors="pt").input_ids
decoder_input_ids = rag_tokenizer("Linda Davis", return_tensors="pt").input_ids
rag_token = self.token_model
rag_token.set_retriever(rag_retriever)
input_ids = rag_question_encoder_tokenizer(
"who sings does he love me with reba", return_tensors="pt"
).input_ids
decoder_input_ids = rag_decoder_tokenizer("Linda Davis", return_tensors="pt").input_ids
input_ids = input_ids.to(torch_device)
decoder_input_ids = decoder_input_ids.to(torch_device)
rag_token = RagToken.from_pretrained(config=rag_config).to(torch_device)
with torch.no_grad():
output = rag_token(
input_ids,
retriever=rag_retriever,
decoder_input_ids=decoder_input_ids,
return_loss=True,
labels=decoder_input_ids,
)
expected_shape = torch.Size([5, 5, 50264])
self.assertEqual(output.logits.shape, expected_shape)
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)
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)
@slow
def test_rag_sequence_generate(self):
rag_config = self.get_rag_config()
rag_decoder_tokenizer = BartTokenizer.from_pretrained("facebook/bart-large-cnn")
rag_question_encoder_tokenizer = DPRQuestionEncoderTokenizer.from_pretrained(
"facebook/dpr-question_encoder-single-nq-base"
)
rag_retriever = RagRetriever(
rag_config,
question_encoder_tokenizer=rag_question_encoder_tokenizer,
generator_tokenizer=rag_decoder_tokenizer,
)
rag_sequence = self.sequence_model
rag_sequence.set_retriever(rag_retriever)
input_ids = rag_question_encoder_tokenizer(
"who sings does he love me with reba", return_tensors="pt"
).input_ids
input_ids = input_ids.to(torch_device)
output_ids = rag_sequence.generate(
input_ids,
)
# sequence generate test
output_text = rag_decoder_tokenizer.decode(output_ids[0], skip_special_tokens=True)
EXPECTED_OUTPUT_TEXT = """The album showed a songwriting maturity and depth of feeling distinctly lacking from their earlier recordings. The album\'s title track refers to secret meetings held against the approval of totalitarian governments in Soviet-dominated states. The only major single release, "One of Us", proved to be the last of ABBA\'s nine number-one singles in Germany."""
self.assertEqual(output_text, EXPECTED_OUTPUT_TEXT)
@slow
def test_rag_token_generate(self):
rag_config = self.get_rag_config()
rag_decoder_tokenizer = BartTokenizer.from_pretrained("facebook/bart-large-cnn")
rag_question_encoder_tokenizer = DPRQuestionEncoderTokenizer.from_pretrained(
"facebook/dpr-question_encoder-single-nq-base"
)
rag_retriever = RagRetriever(
rag_config,
question_encoder_tokenizer=rag_question_encoder_tokenizer,
generator_tokenizer=rag_decoder_tokenizer,
)
rag_token = self.token_model
rag_token.set_retriever(rag_retriever)
input_ids = rag_question_encoder_tokenizer(
"who sings does he love me with reba", return_tensors="pt"
).input_ids
input_ids = input_ids.to(torch_device)
output_ids = rag_token.generate(
input_ids, decoder_start_token_id=rag_token.generator.config.decoder_start_token_id
)
# sequence generate test
output_text = rag_decoder_tokenizer.decode(output_ids[0], skip_special_tokens=True)
EXPECTED_OUTPUT_TEXT = """. The song peaked at"""
self.assertEqual(output_text, EXPECTED_OUTPUT_TEXT)
@require_torch
@require_retrieval
class RagModelSaveLoadTests(unittest.TestCase):
def get_rag_config(self):
question_encoder_config = AutoConfig.from_pretrained("facebook/dpr-question_encoder-single-nq-base")
generator_config = AutoConfig.from_pretrained("facebook/bart-large-cnn")
return RagConfig.from_question_encoder_generator_configs(
question_encoder_config,
generator_config,
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,
dataset="wiki_dpr",
dataset_split="train",
index_name="exact",
index_path=None,
use_dummy_dataset=True,
retrieval_vector_size=768,
retrieval_batch_size=8,
)
@slow
def test_rag_sequence_from_pretrained(self):
rag_config = self.get_rag_config()
rag_decoder_tokenizer = BartTokenizer.from_pretrained("facebook/bart-large-cnn")
rag_question_encoder_tokenizer = DPRQuestionEncoderTokenizer.from_pretrained(
"facebook/dpr-question_encoder-single-nq-base"
)
rag_retriever = RagRetriever(
rag_config,
question_encoder_tokenizer=rag_question_encoder_tokenizer,
generator_tokenizer=rag_decoder_tokenizer,
)
input_ids = rag_question_encoder_tokenizer(
"who sings does he love me with reba", return_tensors="pt"
).input_ids
decoder_input_ids = rag_decoder_tokenizer("Linda Davis", return_tensors="pt").input_ids
input_ids = input_ids.to(torch_device)
decoder_input_ids = decoder_input_ids.to(torch_device)
with tempfile.TemporaryDirectory() as tmp_dirname:
rag_sequence = RagSequenceForGeneration.from_pretrained_question_encoder_generator(
"facebook/dpr-question_encoder-single-nq-base",
"facebook/bart-large-cnn",
retriever=rag_retriever,
config=rag_config,
)
# check that the from pretrained methods work
rag_sequence.save_pretrained(tmp_dirname)
rag_sequence.from_pretrained(tmp_dirname, retriever=rag_retriever)
with torch.no_grad():
output = rag_sequence(
input_ids,
labels=decoder_input_ids,
)
loss_pretrained = output.loss
del rag_sequence
question_encoder = AutoModel.from_pretrained("facebook/dpr-question_encoder-single-nq-base")
generator = AutoModelForSeq2SeqLM.from_pretrained("facebook/bart-large-cnn")
rag_sequence = RagSequenceForGeneration(
config=rag_config, question_encoder=question_encoder, generator=generator, retriever=rag_retriever
)
with torch.no_grad():
output = rag_sequence(
input_ids,
labels=decoder_input_ids,
)
loss_init = output.loss
self.assertAlmostEqual(loss_pretrained.item(), loss_init.item(), places=4)
@slow
def test_rag_token_from_pretrained(self):
rag_config = self.get_rag_config()
rag_decoder_tokenizer = BartTokenizer.from_pretrained("facebook/bart-large-cnn")
rag_question_encoder_tokenizer = DPRQuestionEncoderTokenizer.from_pretrained(
"facebook/dpr-question_encoder-single-nq-base"
)
rag_retriever = RagRetriever(
rag_config,
question_encoder_tokenizer=rag_question_encoder_tokenizer,
generator_tokenizer=rag_decoder_tokenizer,
)
input_ids = rag_question_encoder_tokenizer(
"who sings does he love me with reba", return_tensors="pt"
).input_ids
decoder_input_ids = rag_decoder_tokenizer("Linda Davis", return_tensors="pt").input_ids
input_ids = input_ids.to(torch_device)
decoder_input_ids = decoder_input_ids.to(torch_device)
with tempfile.TemporaryDirectory() as tmp_dirname:
rag_token = RagTokenForGeneration.from_pretrained_question_encoder_generator(
"facebook/dpr-question_encoder-single-nq-base",
"facebook/bart-large-cnn",
retriever=rag_retriever,
config=rag_config,
)
# check that the from pretrained methods work
rag_token.save_pretrained(tmp_dirname)
rag_token.from_pretrained(tmp_dirname, retriever=rag_retriever)
with torch.no_grad():
output = rag_token(
input_ids,
labels=decoder_input_ids,
)
loss_pretrained = output.loss
del rag_token
question_encoder = AutoModel.from_pretrained("facebook/dpr-question_encoder-single-nq-base")
generator = AutoModelForSeq2SeqLM.from_pretrained("facebook/bart-large-cnn")
rag_token = RagTokenForGeneration(
config=rag_config, question_encoder=question_encoder, generator=generator, retriever=rag_retriever
)
with torch.no_grad():
output = rag_token(
input_ids,
labels=decoder_input_ids,
)
loss_init = output.loss
self.assertAlmostEqual(loss_pretrained.item(), loss_init.item(), places=4)
+264
View File
@@ -0,0 +1,264 @@
import json
import os
import pickle
import shutil
import tempfile
from unittest import TestCase
from unittest.mock import patch
import faiss
import numpy as np
from datasets import Dataset
from transformers.configuration_bart import BartConfig
from transformers.configuration_dpr import DPRConfig
from transformers.configuration_rag import RagConfig
from transformers.retrieval_rag import RagPyTorchDistributedRetriever, RagRetriever
from transformers.testing_utils import require_datasets, require_faiss, require_torch
from transformers.tokenization_bart import BartTokenizer
from transformers.tokenization_bert import VOCAB_FILES_NAMES as DPR_VOCAB_FILES_NAMES
from transformers.tokenization_dpr import DPRQuestionEncoderTokenizer
from transformers.tokenization_roberta import VOCAB_FILES_NAMES as BART_VOCAB_FILES_NAMES
@require_faiss
@require_datasets
@require_torch
class RagRetrieverTest(TestCase):
def setUp(self):
self.tmpdirname = tempfile.mkdtemp()
self.retrieval_vector_size = 8
# DPR tok
vocab_tokens = [
"[UNK]",
"[CLS]",
"[SEP]",
"[PAD]",
"[MASK]",
"want",
"##want",
"##ed",
"wa",
"un",
"runn",
"##ing",
",",
"low",
"lowest",
]
dpr_tokenizer_path = os.path.join(self.tmpdirname, "dpr_tokenizer")
os.makedirs(dpr_tokenizer_path, exist_ok=True)
self.vocab_file = os.path.join(dpr_tokenizer_path, DPR_VOCAB_FILES_NAMES["vocab_file"])
with open(self.vocab_file, "w", encoding="utf-8") as vocab_writer:
vocab_writer.write("".join([x + "\n" for x in vocab_tokens]))
# BART tok
vocab = [
"l",
"o",
"w",
"e",
"r",
"s",
"t",
"i",
"d",
"n",
"\u0120",
"\u0120l",
"\u0120n",
"\u0120lo",
"\u0120low",
"er",
"\u0120lowest",
"\u0120newer",
"\u0120wider",
"<unk>",
]
vocab_tokens = dict(zip(vocab, range(len(vocab))))
merges = ["#version: 0.2", "\u0120 l", "\u0120l o", "\u0120lo w", "e r", ""]
self.special_tokens_map = {"unk_token": "<unk>"}
bart_tokenizer_path = os.path.join(self.tmpdirname, "bart_tokenizer")
os.makedirs(bart_tokenizer_path, exist_ok=True)
self.vocab_file = os.path.join(bart_tokenizer_path, BART_VOCAB_FILES_NAMES["vocab_file"])
self.merges_file = os.path.join(bart_tokenizer_path, BART_VOCAB_FILES_NAMES["merges_file"])
with open(self.vocab_file, "w", encoding="utf-8") as fp:
fp.write(json.dumps(vocab_tokens) + "\n")
with open(self.merges_file, "w", encoding="utf-8") as fp:
fp.write("\n".join(merges))
def get_dpr_tokenizer(self) -> DPRQuestionEncoderTokenizer:
return DPRQuestionEncoderTokenizer.from_pretrained(os.path.join(self.tmpdirname, "dpr_tokenizer"))
def get_bart_tokenizer(self) -> BartTokenizer:
return BartTokenizer.from_pretrained(os.path.join(self.tmpdirname, "bart_tokenizer"))
def tearDown(self):
shutil.rmtree(self.tmpdirname)
def get_dummy_hf_index_retriever(self) -> RagRetriever:
dataset = Dataset.from_dict(
{
"id": ["0", "1"],
"text": ["foo", "bar"],
"title": ["Foo", "Bar"],
"embeddings": [np.ones(self.retrieval_vector_size), 2 * np.ones(self.retrieval_vector_size)],
}
)
dataset.add_faiss_index("embeddings", string_factory="Flat", metric_type=faiss.METRIC_INNER_PRODUCT)
config = RagConfig(
retrieval_vector_size=self.retrieval_vector_size,
question_encoder=DPRConfig().to_dict(),
generator=BartConfig().to_dict(),
)
with patch("transformers.retrieval_rag.load_dataset") as mock_load_dataset:
mock_load_dataset.return_value = dataset
retriever = RagRetriever(
config,
question_encoder_tokenizer=self.get_dpr_tokenizer(),
generator_tokenizer=self.get_bart_tokenizer(),
)
return retriever
def get_dummy_pytorch_distributed_retriever(self, init_retrieval, port=12345) -> RagRetriever:
dataset = Dataset.from_dict(
{
"id": ["0", "1"],
"text": ["foo", "bar"],
"title": ["Foo", "Bar"],
"embeddings": [np.ones(self.retrieval_vector_size), 2 * np.ones(self.retrieval_vector_size)],
}
)
dataset.add_faiss_index("embeddings", string_factory="Flat", metric_type=faiss.METRIC_INNER_PRODUCT)
config = RagConfig(
retrieval_vector_size=self.retrieval_vector_size,
question_encoder=DPRConfig().to_dict(),
generator=BartConfig().to_dict(),
)
with patch("transformers.retrieval_rag.load_dataset") as mock_load_dataset:
mock_load_dataset.return_value = dataset
retriever = RagPyTorchDistributedRetriever(
config,
question_encoder_tokenizer=self.get_dpr_tokenizer(),
generator_tokenizer=self.get_bart_tokenizer(),
)
if init_retrieval:
retriever.init_retrieval(port)
return retriever
def get_dummy_legacy_index_retriever(self) -> RagRetriever:
dataset = Dataset.from_dict(
{
"id": ["0", "1"],
"text": ["foo", "bar"],
"title": ["Foo", "Bar"],
"embeddings": [np.ones(self.retrieval_vector_size + 1), 2 * np.ones(self.retrieval_vector_size + 1)],
}
)
dataset.add_faiss_index("embeddings", string_factory="Flat", metric_type=faiss.METRIC_INNER_PRODUCT)
index_file_name = os.path.join(self.tmpdirname, "hf_bert_base.hnswSQ8_correct_phi_128.c_index")
dataset.save_faiss_index("embeddings", index_file_name + ".index.dpr")
pickle.dump(dataset["id"], open(index_file_name + ".index_meta.dpr", "wb"))
passages_file_name = os.path.join(self.tmpdirname, "psgs_w100.tsv.pkl")
passages = {sample["id"]: [sample["text"], sample["title"]] for sample in dataset}
pickle.dump(passages, open(passages_file_name, "wb"))
config = RagConfig(
retrieval_vector_size=self.retrieval_vector_size,
question_encoder=DPRConfig().to_dict(),
generator=BartConfig().to_dict(),
index_name="legacy",
index_path=self.tmpdirname,
passages_path=self.tmpdirname,
)
retriever = RagRetriever(
config, question_encoder_tokenizer=self.get_dpr_tokenizer(), generator_tokenizer=self.get_bart_tokenizer()
)
return retriever
def test_hf_index_retriever_retrieve(self):
n_docs = 1
retriever = self.get_dummy_hf_index_retriever()
hidden_states = np.array(
[np.ones(self.retrieval_vector_size), -np.ones(self.retrieval_vector_size)], dtype=np.float32
)
retrieved_doc_embeds, doc_ids, doc_dicts = retriever.retrieve(hidden_states, n_docs=n_docs)
self.assertEqual(retrieved_doc_embeds.shape, (2, n_docs, self.retrieval_vector_size))
self.assertEqual(len(doc_dicts), 2)
self.assertEqual(sorted(doc_dicts[0]), ["embeddings", "id", "text", "title"])
self.assertEqual(len(doc_dicts[0]["id"]), n_docs)
self.assertEqual(doc_dicts[0]["id"][0], "1") # max inner product is reached with second doc
self.assertEqual(doc_dicts[1]["id"][0], "0") # max inner product is reached with first doc
self.assertListEqual(list(doc_ids), [1, 0])
def test_pytorch_distributed_retriever_retrieve(self):
n_docs = 1
retriever = self.get_dummy_pytorch_distributed_retriever(init_retrieval=True)
hidden_states = np.array(
[np.ones(self.retrieval_vector_size), -np.ones(self.retrieval_vector_size)], dtype=np.float32
)
retrieved_doc_embeds, doc_ids, doc_dicts = retriever.retrieve(hidden_states, n_docs=n_docs)
self.assertEqual(retrieved_doc_embeds.shape, (2, n_docs, self.retrieval_vector_size))
self.assertEqual(len(doc_dicts), 2)
self.assertEqual(sorted(doc_dicts[0]), ["embeddings", "id", "text", "title"])
self.assertEqual(len(doc_dicts[0]["id"]), n_docs)
self.assertEqual(doc_dicts[0]["id"][0], "1") # max inner product is reached with second doc
self.assertEqual(doc_dicts[1]["id"][0], "0") # max inner product is reached with first doc
self.assertListEqual(list(doc_ids), [1, 0])
def test_legacy_index_retriever_retrieve(self):
n_docs = 1
retriever = self.get_dummy_legacy_index_retriever()
hidden_states = np.array(
[np.ones(self.retrieval_vector_size), -np.ones(self.retrieval_vector_size)], dtype=np.float32
)
retrieved_doc_embeds, doc_ids, doc_dicts = retriever.retrieve(hidden_states, n_docs=n_docs)
self.assertEqual(retrieved_doc_embeds.shape, (2, n_docs, self.retrieval_vector_size))
self.assertEqual(len(doc_dicts), 2)
self.assertEqual(sorted(doc_dicts[0]), ["text", "title"])
self.assertEqual(len(doc_dicts[0]["text"]), n_docs)
self.assertEqual(doc_dicts[0]["text"][0], "bar") # max inner product is reached with second doc
self.assertEqual(doc_dicts[1]["text"][0], "foo") # max inner product is reached with first doc
self.assertListEqual(list(doc_ids), [1, 0])
def test_hf_index_retriever_call(self):
import torch
n_docs = 1
retriever = self.get_dummy_hf_index_retriever()
question_input_ids = [[5, 7], [10, 11]]
hidden_states = np.array(
[np.ones(self.retrieval_vector_size), -np.ones(self.retrieval_vector_size)], dtype=np.float32
)
out = retriever(question_input_ids, hidden_states, prefix=retriever.config.generator.prefix, n_docs=n_docs)
context_input_ids, context_attention_mask, retrieved_doc_embeds = (
out["context_input_ids"],
out["context_attention_mask"],
out["retrieved_doc_embeds"],
)
self.assertEqual(retrieved_doc_embeds.shape, (2, n_docs, self.retrieval_vector_size))
self.assertIsInstance(context_input_ids, list)
self.assertIsInstance(context_attention_mask, list)
self.assertIsInstance(retrieved_doc_embeds, np.ndarray)
out = retriever(
question_input_ids,
hidden_states,
prefix=retriever.config.generator.prefix,
n_docs=n_docs,
return_tensors="pt",
)
context_input_ids, context_attention_mask, retrieved_doc_embeds, doc_ids = ( # noqa: F841
out["context_input_ids"],
out["context_attention_mask"],
out["retrieved_doc_embeds"],
out["doc_ids"],
)
self.assertEqual(retrieved_doc_embeds.shape, (2, n_docs, self.retrieval_vector_size))
self.assertIsInstance(context_input_ids, torch.Tensor)
self.assertIsInstance(context_attention_mask, torch.Tensor)
self.assertIsInstance(retrieved_doc_embeds, torch.Tensor)