Compare commits
52
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a2f830d5a1 | ||
|
|
8f5fd79b8f | ||
|
|
c1be41f452 | ||
|
|
135689bba3 | ||
|
|
64141bab07 | ||
|
|
3cd4a574c2 | ||
|
|
237f27f724 | ||
|
|
4274e9c223 | ||
|
|
47b137e175 | ||
|
|
82afc4b93e | ||
|
|
59ce19cde4 | ||
|
|
3abfc19ae0 | ||
|
|
5b47f0bc3b | ||
|
|
897101cfce | ||
|
|
60c8defa01 | ||
|
|
d7e169b3d9 | ||
|
|
1cf7dcbe71 | ||
|
|
82a22c20a8 | ||
|
|
d8fb4c836b | ||
|
|
6bfa18e1a7 | ||
|
|
aca8e30ddf | ||
|
|
349f85f241 | ||
|
|
3860f3144a | ||
|
|
f69a9d32fa | ||
|
|
c8c5ce0fd3 | ||
|
|
6ab7a4584b | ||
|
|
e210739bef | ||
|
|
8977533c8d | ||
|
|
cf9561a4bb | ||
|
|
f64b6c1dc8 | ||
|
|
b378005edf | ||
|
|
eaf68afffe | ||
|
|
4f546ad160 | ||
|
|
2884cd7bdb | ||
|
|
0f32ad8319 | ||
|
|
2c83e1bfd6 | ||
|
|
15af641996 | ||
|
|
95cd16275c | ||
|
|
044fa94285 | ||
|
|
f780b9f415 | ||
|
|
2f1211bbb5 | ||
|
|
00a1fc9ae4 | ||
|
|
2faaa4ad3c | ||
|
|
c1bc9fe05d | ||
|
|
6e9f30748f | ||
|
|
bc440f3e7c | ||
|
|
2094d37888 | ||
|
|
b3f9e986d9 | ||
|
|
b4a094dd97 | ||
|
|
3688823d19 | ||
|
|
e0a37450e1 | ||
|
|
720ee41342 |
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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"),
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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),
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
+706
-517
File diff suppressed because it is too large
Load Diff
+281
-198
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)),
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user