Compare commits

...
Author SHA1 Message Date
Your Name 9b561de9ca LegacyIndex index download refactor 2020-09-14 10:22:09 -07:00
Patrick von Platen 17a8621661 add tokenizer to slow integration test and allow retriever to run on cpu 2020-09-14 12:37:43 +02:00
Your Name 64796004dc RAG integration tests 2020-09-13 16:32:23 -07:00
Your Name df5eec9f14 rename nlp->datasets 2020-09-12 14:19:43 -07:00
Your Name 1497120d89 Post-rebase refactor 2020-09-12 14:16:53 -07:00
Your Name 593765088e Merge remote-tracking branch 'upstream/master' into rag 2020-09-12 14:00:57 -07:00
Your Name 4ad9cb1bd9 Fix tests 3 2020-09-12 12:37:46 -07:00
Your Name 6215f3b5e0 Refactor test requirements 2020-09-12 12:16:34 -07:00
Your Name 9795dc3464 Fix tests 2 2020-09-12 11:37:29 -07:00
Your Name a4c7c25cd2 Fix tests 2020-09-12 11:20:30 -07:00
Your Name 945d56995e Extra dependencied for tests 2020-09-12 11:12:26 -07:00
Your Name e82aca09a5 Improve readme + post-rebase refactor 2020-09-12 11:07:18 -07:00
Your Name 845c18d9af post-merge fixes 2020-09-10 05:15:22 -07:00
Your Name 0f3dc78c0b Merge remote-tracking branch 'upstream/master' into rag 2020-09-10 04:43:01 -07:00
Your Name 21e8c67bc7 Comments fix 2020-09-10 04:37:28 -07:00
Your Name 972d240ae6 fix tests 2 2020-09-09 12:52:46 -07:00
Your Name 2b8ab2eef3 fix test 2020-09-09 12:42:56 -07:00
Your Name 706a7c064d minor cleanup plus initial tests 2020-09-09 09:46:05 -07:00
Your Name d37b95d39f Fix RAG Sequence generation 2020-09-08 04:00:29 -07:00
Your Name cbf479df13 fix quality 2020-09-07 15:50:28 -07:00
Your Name 0b9b2840c3 Refactor as per suggestions in https://github.com/huggingface/transformers/pull/6813#issuecomment-687208867 2020-09-07 15:34:19 -07:00
Your Name ce3684bbcb Fix quality errors 2020-08-30 04:25:28 -07:00
Your Name b8a2ba8624 Fix import errors 2020-08-30 04:03:57 -07:00
Your Name cc4ba034d6 Fix retrieval wit HF index 2020-08-28 14:27:42 -07:00
Your Name 84d8061dee Remove set_up_rag_env.sh file 2020-08-28 08:28:15 -07:00
Your Name 52bbf8dfe3 Add documentation and cleanup 2020-08-28 08:22:47 -07:00
Your Name 930bdfaaa8 Merge remote-tracking branch 'upstream/master' into rag 2020-08-28 00:36:14 -07:00
Your Name 1e1a671614 Finetuning refactoring and cleanup 2020-08-28 00:35:15 -07:00
Your Name edbfab74d9 Retrieval refactor 2020-08-26 15:21:01 -07:00
Your Name 3fa122af11 use_bos fix 2020-08-24 03:36:38 -07:00
Your Name 85929c02dd Various fixes + finetuning logic 2020-08-24 03:19:46 -07:00
Your Name 2e4dac6802 Merge remote-tracking branch 'upstream/master' into rag 2020-08-12 04:32:24 -07:00
Ola Piktus 7184c2b2c2 Merge pull request #1 from patrick-s-h-lewis/rag-ola
Merging latest changes
2020-08-12 13:25:09 +02:00
Aleksandra Piktus 5ed8b6460a Fix rag-token model + refactor 2020-07-30 13:05:37 -07:00
Aleksandra Piktus 5774e511b0 refactor to include modeling outputs + MPI retriever 2020-07-28 13:01:13 -07:00
Aleksandra Piktus 71d2aa59e0 Retrieval evaluation scripts 2020-07-24 06:50:27 -07:00
Aleksandra Piktus 0a6be7b00f improve comments 2020-07-16 15:27:37 -07:00
Aleksandra Piktus 3cb50b0c31 First commit 2020-07-16 13:41:44 -07:00
Aleksandra Piktus 3bf6f0b8f1 fix merge conflicts 2020-07-13 08:53:07 -07:00
Aleksandra Piktus 9ea0a9e825 Formatting / renaming prior to actual work 2020-07-13 08:50:36 -07:00
Patrick Lewis 3f30195fd0 added rag WIP 2020-07-13 08:49:32 -07:00
Aleksandra Piktus 90bb4188e4 Formatting / renaming prior to actual work 2020-07-13 08:42:33 -07:00
Patrick Lewis 6c7d856b86 path fix 2020-07-13 08:42:33 -07:00
Patrick Lewis 17007b7b66 added rag WIP 2020-07-13 08:42:33 -07:00
Aleksandra Piktus ae3079a9c4 Merge branch 'rag' of github.com:patrick-s-h-lewis/transformers into rag 2020-07-08 03:17:47 -07:00
Aleksandra Piktus 97c9900124 Formatting / renaming prior to actual work 2020-07-08 03:15:59 -07:00
Patrick Lewis cbcfc0e948 path fix 2020-07-08 03:15:59 -07:00
Patrick Lewis 5fd9e5d6ab added rag WIP 2020-07-08 03:15:59 -07:00
Aleksandra Piktus 20c8a9aa36 Formatting / renaming prior to actual work 2020-07-06 06:04:43 -07:00
Patrick Lewis f8b1e89f8b path fix 2020-07-01 11:21:11 +01:00
Patrick Lewis 23eaa7afcf added rag WIP 2020-07-01 11:13:01 +01:00
19 changed files with 3302 additions and 2 deletions
+2
View File
@@ -11,6 +11,7 @@ __pycache__/
# tests and logs
tests/fixtures
logs/
lightning_logs/
# Distribution / packaging
.Python
@@ -139,6 +140,7 @@ runs
/wandb
/examples/runs
/examples/**/*.args
/examples/rag/sweep
# data
/data
+2
View File
@@ -366,6 +366,8 @@ def generic_train(
if args.gpus > 1:
train_params["distributed_backend"] = "ddp"
train_params["accumulate_grad_batches"] = args.accumulate_grad_batches
trainer = pl.Trainer.from_argparse_args(
args,
weights_summary=None,
+85
View File
@@ -0,0 +1,85 @@
# Intro
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`.
Key files:
- `modeling_rag.py`, `tokenization_rag.py`, `configuration_rag.py` the core model implementation
- `retrieval_rag.py` - a distributed retriever built on top of the `torch.distributed` communication package. The retriever is an interface between the model and the faiss index of the encoded documents. During training, all workers initialize their own instance of the retriever, however, only the main worker loads the index into memory, which prevents OOMs on machines with multiple GPUs (we store the index in RAM). The index itself is based on the `nlp.Datasets`. We also implement a variant compatible with indices built using the original DPR implementation (https://github.com/facebookresearch/DPR)
- `eval_rag.py` - an evaluation script which allows to perform the evaluation end to end (measures the exact match and F1 on the downstream task) as well as the evaluation of the retrieval component alone (measures precision@k).
- `finetune.py` - a training script for finetuning RAG models.
# Finetuning
Our finetuning logic is based on scripts from [`examples/seq2seq`](https://github.com/huggingface/transformers/tree/master/examples/seq2seq).
Follow instructions there regarding data preprocessing. A sample finetuning command:
```
python examples/rag/finetune.py \
--data_dir $DATA_DIR \
--output_dir $OUTPUT_DIR \
--model_name_or_path $MODEL_NAME_OR_PATH \
--model_type rag_sequence \
--fp16 \
--gpus 8
```
# Evaluation
Apart for parameters specifying the model that's being evaluated and some extra parameters, the evaluation script expects paths to two files:
- `evaluation_set` - a path file specifying the input dataset for evaluation, a single datapoint per line, e.g.
```who is the owner of reading football club```
- `gold_data_path` - a path to a file contaning ground truth answers for samples from the `evaluation_set`.
We expect the following formats of the gold data file:
- for e2e evaluation, we support two formats of gold files:
- `qa` - where a single line in the following format: input [tab] output_list, e.g.:
```
who is the owner of reading football club ['Xiu Li Dai', 'Dai Yongge', 'Dai Xiuli', 'Yongge Dai']
```
- `ans` - where a single line of the gold file contains the expected output string,
```
Xiu Li Dai
```
- for retrieval evaluation, we expect a tab-separated list of Wikipedia page titles constituting positive contexts for a given query, e.g. given a question `who sings does he love me with reba`, a line with ground truth retrieval data could look as follows:
```
Does He Love You Does He Love You Red Sandy Spika dress of Reba McEntire Greatest Hits Volume Two (Reba McEntire album) Shoot for the Moon (album)
```
## Retrieval evaluation
We demonstrate how to evaluate retrieval against DPR evaluation data. You can download respective files from links listed [here](https://github.com/facebookresearch/DPR/blob/master/data/download_data.py#L39-L45).
1. Download and unzip the gold data file. We use the `biencoder-nq-dev` from https://dl.fbaipublicfiles.com/dpr/data/retriever/biencoder-nq-dev.json.gz.
2. Parse the unziped file using the `parse_dpr_relevance_data.py`
```
python examples/rag/parse_dpr_relevance_data.py --src_path path/to/unziped/biencoder-nq-dev.json --evaluation_set path/to/output/biencoder-nq-dev.questions --gold_data_path path/to/output/biencoder-nq-dev.pages
```
3. Run evaluation:
```
python examples/rag/eval_rag.py \
--model_name_or_path $MODEL_NAME_OR_PATH \ # model name or path of the model we're evaluating
--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
--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.
```
## End-to-end evaluation
```
python examples/rag/eval_rag.py \
--model_name_or_path /private/home/piktus/rag_huggingface/data/repro-rag-sequence-63/ \
--model_type rag_sequence \
--evaluation_set path/to/test.source \
--gold_data_path path/to/gold_data \
--predictions_path path/to/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
```
View File
+30
View File
@@ -0,0 +1,30 @@
import logging
import os
from pytorch_lightning.callbacks import ModelCheckpoint
logger = logging.getLogger(__name__)
def get_checkpoint_callback(output_dir, metric):
"""Saves the best model by validation ROUGE2 score."""
if metric == "rouge2":
exp = "{val_avg_rouge2:.4f}-{step_count}"
elif metric == "bleu":
exp = "{val_avg_bleu:.4f}-{step_count}"
elif metric == "em":
exp = "{val_avg_em:.4f}-{step_count}"
else:
raise NotImplementedError(
f"seq2seq callbacks only support rouge2 and bleu, got {metric}, You can make your own by adding to this function."
)
checkpoint_callback = ModelCheckpoint(
filepath=os.path.join(output_dir, exp),
monitor=f"val_{metric}",
mode="max",
save_top_k=10,
period=0, # maybe save a checkpoint every time val is run, not just end of epoch.
)
return checkpoint_callback
+312
View File
@@ -0,0 +1,312 @@
""" Evaluation script for RAG models."""
import argparse
import ast
import logging
import os
import sys
import pandas as pd
import torch
from tqdm import tqdm
from transformers import BartForConditionalGeneration, BartTokenizer, RagRetriever, RagSequence, RagToken
from transformers import logging as transformers_logging
sys.path.append(os.path.join(os.getcwd())) # noqa: E402 # isort:skip
from examples.rag.utils import exact_match_score, f1_score # noqa: E402 # isort:skip
logger = logging.getLogger(__name__)
logging.basicConfig(level=logging.INFO)
transformers_logging.set_verbosity_info()
def infer_model_type(model_name_or_path):
if "token" in model_name_or_path:
return "rag_token"
if "sequence" in model_name_or_path:
return "rag_sequence"
if "bart" in model_name_or_path:
return "bart"
return None
def metric_max_over_ground_truths(metric_fn, prediction, ground_truths):
scores_for_ground_truths = []
for ground_truth in ground_truths:
score = metric_fn(prediction, ground_truth)
scores_for_ground_truths.append(score)
return max(scores_for_ground_truths)
def get_scores(args, preds_path, gold_data_path):
hypos = [line.strip() for line in open(preds_path, "r").readlines()]
answers = []
if args.gold_data_mode == "qa":
data = pd.read_csv(gold_data_path, sep="\t", header=None)
for answer_list in data[1]:
ground_truths = ast.literal_eval(answer_list)
answers.append(ground_truths)
else:
references = [line.strip() for line in open(gold_data_path, "r").readlines()]
answers = [[reference] for reference in references]
f1 = em = total = 0
for prediction, ground_truths in zip(hypos, answers):
total += 1
em += metric_max_over_ground_truths(exact_match_score, prediction, ground_truths)
f1 += metric_max_over_ground_truths(f1_score, prediction, ground_truths)
em = 100.0 * em / total
f1 = 100.0 * f1 / total
logger.info("F1: {}".format(f1))
logger.info("EM: {}".format(em))
def get_precision_at_k(args, preds_path, gold_data_path):
k = args.k
hypos = [line.strip() for line in open(preds_path, "r").readlines()]
references = [line.strip() for line in open(gold_data_path, "r").readlines()]
em = total = 0
for hypo, reference in zip(hypos, references):
hypo_provenance = set(hypo.split("\t")[:k])
ref_provenance = set(reference.split("\t")[1 : (k + 1)])
total += 1
em += len(hypo_provenance & ref_provenance) / k
em = 100.0 * em / total
logger.info("Precision@{}: {}".format(k, em))
def evaluate_batch_retrieval(args, rag_model, tokenizer, retriever, questions):
def strip_title(title):
if title.startswith('"'):
title = title[1:]
if title.endswith('"'):
title = title[:-1]
return title
retriever_inputs = tokenizer.batch_encode_plus(
questions,
return_tensors="pt",
padding=True,
truncation=True,
)
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)
provenance_strings = []
for docs in all_docs:
provenance = [strip_title(title) for title in docs["title"]]
provenance_strings.append("\t".join(provenance))
return provenance_strings
def evaluate_batch_e2e(args, rag_model, tokenizer, retriever, questions):
with torch.no_grad():
input_ids = tokenizer.batch_encode_plus(questions, return_tensors="pt", padding=True, truncation=True)[
"input_ids"
].to(args.device)
outputs = rag_model.generate(
input_ids,
retriever=retriever,
num_beams=args.num_beams,
min_length=args.min_length,
max_length=args.max_length,
early_stopping=False,
num_return_sequences=1,
bad_words_ids=[[0, 0]], # BART likes to repeat BOS tokens, dont allow it to generate more than one
clean_up_tokenization=True,
print_docs=args.print_docs,
)
answers = tokenizer.batch_decode(outputs, skip_special_tokens=True)
if args.print_predictions:
for q, a in zip(questions, answers):
logger.info("Q: {} - A: {}".format(q, a))
return answers
def get_args():
parser = argparse.ArgumentParser()
parser.add_argument(
"--model_type",
choices=["rag_sequence", "rag_token", "bart"],
type=str,
help="RAG model type: rag_sequence, rag_token or bart, if none specified, the type is inferred from the model_name_or_path",
)
parser.add_argument(
"--retriever_type",
default=None,
choices=["hf_retriever", "legacy_retriever"],
type=str,
help="RAG model retriever type",
)
parser.add_argument(
"--index_path",
default=None,
type=str,
help="Path to the retrieval index",
)
parser.add_argument("--n_docs", default=5, type=int, help="Number of retrieved docs")
parser.add_argument(
"--model_name_or_path",
default=None,
type=str,
required=True,
help="Path to pretrained checkpoints or model identifier from huggingface.co/models",
)
parser.add_argument(
"--eval_mode",
choices=["e2e", "retrieval"],
default="e2e",
type=str,
help="Evaluation mode, e2e calculates exact match and F1 of the downstream task, retrieval calulates precision@k.",
)
parser.add_argument("--k", default=1, type=int, help="k for the precision@k calculation")
parser.add_argument(
"--evaluation_set",
default=None,
type=str,
required=True,
help="Path to a file containing evaluation samples",
)
parser.add_argument(
"--gold_data_path",
default=None,
type=str,
required=True,
help="Path to a tab-separated file with gold samples",
)
parser.add_argument(
"--gold_data_mode",
default="qa",
type=str,
choices=["qa", "ans"],
help="Format of the gold data file"
"qa - a single line in the following format: question [tab] answer_list"
"ans - a single line of the gold file contains the expected answer string",
)
parser.add_argument(
"--predictions_path",
type=str,
default="predictions.txt",
help="Path under which to store prediction files. The base dir needs to exists, the file will be generated.",
)
parser.add_argument(
"--eval_all_checkpoints",
action="store_true",
help="Evaluate all checkpoints starting with the same prefix as model_name ending and ending with step number",
)
parser.add_argument(
"--eval_batch_size",
default=8,
type=int,
help="Batch size per GPU/CPU for evaluation.",
)
parser.add_argument(
"--recalculate",
help="Recalculate predictions even if the prediction file exists",
action="store_true",
)
parser.add_argument(
"--num_beams",
default=4,
type=int,
help="Number of beams to be used when generating answers",
)
parser.add_argument("--min_length", default=1, type=int, help="Min length of the generated answers")
parser.add_argument("--max_length", default=50, type=int, help="Max length of the generated answers")
parser.add_argument(
"--print_predictions",
action="store_true",
help="If True, prints predictions while evaluating.",
)
parser.add_argument(
"--print_docs",
action="store_true",
help="If True, prints docs retried while generating.",
)
args = parser.parse_args()
args.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
return args
def main(args):
model_kwargs = {}
if args.model_type is None:
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_kwargs["n_docs"] = args.n_docs
if args.retriever_type is not None:
model_kwargs["retriever_type"] = args.retriever_type
if args.index_path is not None:
model_kwargs["index_path"] = args.index_path
else:
model_class = BartForConditionalGeneration
checkpoints = (
[f.path for f in os.scandir(args.model_name_or_path) if f.is_dir()]
if args.eval_all_checkpoints
else [args.model_name_or_path]
)
logger.info("Evaluate the following checkpoints: %s", checkpoints)
score_fn = get_scores if args.eval_mode == "e2e" else get_precision_at_k
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)
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))
model = model_class.from_pretrained(checkpoint, **model_kwargs)
model.to(args.device)
retriever = RagRetriever(model.config)
tokenizer = (
retriever.generator_tokenizer
if args.model_type != "bart" and args.eval_mode == "e2e"
else retriever.question_encoder_tokenizer
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:
questions = []
for line in tqdm(eval_file):
questions.append(line.strip())
if len(questions) == args.eval_batch_size:
answers = evaluate_batch_fn(args, model, tokenizer, retriever, questions)
preds_file.write("\n".join(answers) + "\n")
preds_file.flush()
questions = []
if len(questions) > 0:
answers = evaluate_batch_fn(args, model, tokenizer, retriever, questions)
preds_file.write("\n".join(answers))
preds_file.flush()
score_fn(args, args.predictions_path, args.gold_data_path)
if __name__ == "__main__":
args = get_args()
main(args)
+475
View File
@@ -0,0 +1,475 @@
"""Finetuning script for RAG models. Adapted from examples.seq2seq.finetune.py"""
import argparse
import glob
import logging
import os
import sys
import time
import warnings
from collections import defaultdict
from pathlib import Path
from typing import Any, Dict, List, Tuple
import numpy as np
import pytorch_lightning as pl
import torch
import torch.distributed as dist
from torch.utils.data import DataLoader
from transformers import (
AutoConfig,
AutoTokenizer,
BartForConditionalGeneration,
RagConfig,
RagRetriever,
RagSequence,
RagToken,
T5ForConditionalGeneration,
get_linear_schedule_with_warmup,
)
from transformers import logging as transformers_logging
sys.path.append(os.path.join(os.getcwd())) # noqa: E402 # noqa: E402 # isort:skip
from examples.lightning_base import BaseTransformer, add_generic_args, generic_train # noqa: E402 # isort:skip
from examples.rag.callbacks import get_checkpoint_callback # noqa: E402 # isort:skip
from examples.rag.utils import ( # noqa: E402 # isort:skip
Seq2SeqDataset,
calculate_exact_match,
is_rag_model,
set_extra_model_params,
)
from examples.seq2seq.callbacks import Seq2SeqLoggingCallback, get_early_stopping_callback # noqa: E402 # isort:skip
from examples.seq2seq.utils import ( # noqa: E402 # isort:skip
flatten_list,
get_git_info,
lmap,
pickle_save,
save_git_info,
save_json,
)
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
transformers_logging.set_verbosity_info()
class AttrDict(dict):
def __init__(self, *args, **kwargs):
super(AttrDict, self).__init__(*args, **kwargs)
self.__dict__ = self
class GenerativeQAModule(BaseTransformer):
mode = "generative_qa"
loss_names = ["loss"]
metric_names = ["em"]
val_metric = "em"
def __init__(self, hparams, **kwargs):
# when loading from a pytorch lightning checkpoint, hparams are passed as dict
if isinstance(hparams, dict):
hparams = AttrDict(hparams)
if hparams.model_type == "rag_sequence":
self.model_class = RagSequence
elif hparams.model_type == "rag_token":
self.model_class = RagToken
elif hparams.model_type == "bart":
self.model_class = BartForConditionalGeneration
else:
self.model_class = T5ForConditionalGeneration
self.is_rag_model = is_rag_model(hparams.model_type)
config_class = RagConfig if self.is_rag_model else AutoConfig
config = config_class.from_pretrained(hparams.model_name_or_path)
# 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(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)
else:
if args.prefix is not None:
setattr(config, "prefix", args.prefix)
hparams, config = set_extra_model_params(extra_model_params, hparams, config)
model = self.model_class.from_pretrained(hparams.model_name_or_path, config=config)
generator_config = config
tokenizer = (
AutoTokenizer.from_pretrained(config.pretrained_generator_tokenizer_name_or_path)
if self.is_rag_model
else AutoTokenizer.from_pretrained(hparams.model_name_or_path)
)
super().__init__(hparams, config=config, tokenizer=tokenizer, model=model)
self.retriever = RagRetriever(self.model.config) if self.is_rag_model else None
save_git_info(self.hparams.output_dir)
self.output_dir = Path(self.hparams.output_dir)
self.metrics_save_path = Path(self.output_dir) / "metrics.json"
self.hparams_save_path = Path(self.output_dir) / "hparams.pkl"
pickle_save(self.hparams, self.hparams_save_path)
self.step_count = 0
self.metrics = defaultdict(list)
self.dataset_kwargs: dict = dict(
data_dir=self.hparams.data_dir,
max_source_length=self.hparams.max_source_length,
prefix=generator_config.prefix or "",
)
n_observations_per_split = {
"train": self.hparams.n_train,
"val": self.hparams.n_val,
"test": self.hparams.n_test,
}
self.n_obs = {k: v if v >= 0 else None for k, v in n_observations_per_split.items()}
self.target_lens = {
"train": self.hparams.max_target_length,
"val": self.hparams.val_max_target_length,
"test": self.hparams.test_max_target_length,
}
assert self.target_lens["train"] <= self.target_lens["val"], f"target_lens: {self.target_lens}"
assert self.target_lens["train"] <= self.target_lens["test"], f"target_lens: {self.target_lens}"
self.hparams.git_sha = get_git_info()["repo_sha"]
self.num_workers = hparams.num_workers
self.distributed_port = self.hparams.distributed_port
def init_ddp_connection(self, global_rank: int, world_size: int, is_slurm_managing_tasks: bool = True):
logger.info("Custom init_ddp_connection.")
os.environ["MASTER_PORT"] = str(self.distributed_port)
super().init_ddp_connection(global_rank, world_size, is_slurm_managing_tasks)
if self.is_rag_model:
self.retriever.init_retrieval(self.distributed_port)
def forward(self, input_ids, **kwargs):
return self.model(input_ids, **kwargs)
def ids_to_clean_text(self, generated_ids: List[int]):
gen_text = self.tokenizer.batch_decode(
generated_ids, skip_special_tokens=True, clean_up_tokenization_spaces=True
)
return lmap(str.strip, gen_text)
def _step(self, batch: dict) -> Tuple:
source_ids, source_mask, target_ids = batch["input_ids"], batch["attention_mask"], batch["decoder_input_ids"]
if isinstance(self.model, T5ForConditionalGeneration):
decoder_input_ids = self.model._shift_right(target_ids)
lm_labels = target_ids
elif isinstance(self.model, BartForConditionalGeneration):
decoder_input_ids = target_ids[:, :-1].contiguous()
lm_labels = target_ids[:, 1:].clone()
else:
assert self.is_rag_model
generator = self.model.model.generator
if isinstance(generator, T5ForConditionalGeneration):
decoder_start_token_id = generator.config.decoder_start_token_id
decoder_input_ids = (
torch.cat(
[torch.Tensor([[decoder_start_token_id]] * target_ids.shape[0]).to(target_ids), target_ids],
dim=1,
)
if target_ids.shape[0] < self.target_lens["train"]
else generator._shift_right(target_ids)
)
elif isinstance(generator, BartForConditionalGeneration):
decoder_input_ids = target_ids
lm_labels = None
assert decoder_input_ids is not None
if lm_labels is not None:
outputs = self(
source_ids,
attention_mask=source_mask,
decoder_input_ids=decoder_input_ids,
use_cache=False,
labels=lm_labels,
return_dict=True,
)
else: # RAG models
outputs = self(
source_ids,
retriever=self.retriever,
attention_mask=source_mask,
decoder_input_ids=decoder_input_ids,
use_cache=False,
return_loss=True,
reduce=True,
label_smoothing=self.hparams.label_smoothing,
)
loss = outputs["loss"]
return (loss,)
@property
def pad(self) -> int:
return self.tokenizer.pad_token_id
def training_step(self, batch, batch_idx) -> Dict:
loss_tensors = self._step(batch)
logs = {name: loss for name, loss in zip(self.loss_names, loss_tensors)}
# tokens per batch
logs["tpb"] = batch["input_ids"].ne(self.pad).sum() + batch["decoder_input_ids"].ne(self.pad).sum()
return {"loss": loss_tensors[0], "log": logs}
def validation_step(self, batch, batch_idx) -> Dict:
return self._generative_step(batch)
def validation_epoch_end(self, outputs, prefix="val") -> Dict:
self.step_count += 1
losses = {k: torch.stack([x[k] for x in outputs]).mean() for k in self.loss_names}
loss = losses["loss"]
gen_metrics = {
k: np.array([x[k] for x in outputs]).mean() for k in self.metric_names + ["gen_time", "gen_len"]
}
metrics_tensor: torch.FloatTensor = torch.tensor(gen_metrics[self.val_metric]).type_as(loss)
gen_metrics.update({k: v.item() for k, v in losses.items()})
# fix for https://github.com/PyTorchLightning/pytorch-lightning/issues/2424
if dist.is_initialized():
dist.all_reduce(metrics_tensor, op=dist.ReduceOp.SUM)
metrics_tensor = metrics_tensor / dist.get_world_size()
gen_metrics.update({self.val_metric: metrics_tensor.item()})
losses.update(gen_metrics)
metrics = {f"{prefix}_avg_{k}": x for k, x in losses.items()}
metrics["step_count"] = self.step_count
self.save_metrics(metrics, prefix) # writes to self.metrics_save_path
preds = flatten_list([x["preds"] for x in outputs])
return {"log": metrics, "preds": preds, f"{prefix}_loss": loss, f"{prefix}_{self.val_metric}": metrics_tensor}
def save_metrics(self, latest_metrics, type_path) -> None:
self.metrics[type_path].append(latest_metrics)
save_json(self.metrics, self.metrics_save_path)
def calc_generative_metrics(self, preds, target) -> Dict:
return calculate_exact_match(preds, target)
def _generative_step(self, batch: dict) -> dict:
start_time = time.time()
generated_ids = self.model.generate(
batch["input_ids"],
retriever=self.retriever,
dedup=False, # rag specific parameter
attention_mask=batch["attention_mask"],
use_cache=True,
min_length=1,
max_length=self.target_lens["val"],
)
gen_time = (time.time() - start_time) / batch["input_ids"].shape[0]
preds: List[str] = self.ids_to_clean_text(generated_ids)
target: List[str] = self.ids_to_clean_text(batch["decoder_input_ids"])
loss_tensors = self._step(batch)
base_metrics = {name: loss for name, loss in zip(self.loss_names, loss_tensors)}
gen_metrics: Dict = self.calc_generative_metrics(preds, target)
summ_len = np.mean(lmap(len, generated_ids))
base_metrics.update(gen_time=gen_time, gen_len=summ_len, preds=preds, target=target, **gen_metrics)
return base_metrics
def test_step(self, batch, batch_idx):
return self._generative_step(batch)
def test_epoch_end(self, outputs):
return self.validation_epoch_end(outputs, prefix="test")
def get_dataset(self, type_path) -> Seq2SeqDataset:
n_obs = self.n_obs[type_path]
max_target_length = self.target_lens[type_path]
dataset = Seq2SeqDataset(
self.tokenizer,
type_path=type_path,
n_obs=n_obs,
max_target_length=max_target_length,
**self.dataset_kwargs,
)
return dataset
def get_dataloader(self, type_path: str, batch_size: int, shuffle: bool = False) -> DataLoader:
dataset = self.get_dataset(type_path)
sampler = None
if self.hparams.sortish_sampler and type_path == "train":
assert self.hparams.gpus <= 1 # TODO: assert earlier
sampler = dataset.make_sortish_sampler(batch_size)
shuffle = False
dataloader = DataLoader(
dataset,
batch_size=batch_size,
collate_fn=dataset.collate_fn,
shuffle=shuffle,
num_workers=self.num_workers,
sampler=sampler,
)
return dataloader
def train_dataloader(self) -> DataLoader:
dataloader = self.get_dataloader("train", batch_size=self.hparams.train_batch_size, shuffle=True)
t_total = (
(len(dataloader.dataset) // (self.hparams.train_batch_size * max(1, self.hparams.gpus)))
// self.hparams.accumulate_grad_batches
* float(self.hparams.max_epochs)
)
scheduler = get_linear_schedule_with_warmup(
self.opt, num_warmup_steps=self.hparams.warmup_steps, num_training_steps=t_total
)
if max(scheduler.get_last_lr()) > 0:
warnings.warn("All learning rates are 0")
self.lr_scheduler = scheduler
return dataloader
def val_dataloader(self) -> DataLoader:
return self.get_dataloader("val", batch_size=self.hparams.eval_batch_size)
def test_dataloader(self) -> DataLoader:
return self.get_dataloader("test", batch_size=self.hparams.eval_batch_size)
@pl.utilities.rank_zero_only
def on_save_checkpoint(self, checkpoint: Dict[str, Any]) -> None:
save_path = self.output_dir.joinpath("checkpoint{}".format(self.step_count))
self.model.config.save_step = self.step_count
self.model.save_pretrained(save_path)
self.tokenizer.save_pretrained(save_path)
@staticmethod
def add_model_specific_args(parser, root_dir):
BaseTransformer.add_model_specific_args(parser, root_dir)
add_generic_args(parser, root_dir)
parser.add_argument(
"--max_source_length",
default=128,
type=int,
help="The maximum total input sequence length after tokenization. Sequences longer "
"than this will be truncated, sequences shorter will be padded.",
)
parser.add_argument(
"--max_target_length",
default=25,
type=int,
help="The maximum total input sequence length after tokenization. Sequences longer "
"than this will be truncated, sequences shorter will be padded.",
)
parser.add_argument(
"--val_max_target_length",
default=25,
type=int,
help="The maximum total input sequence length after tokenization. Sequences longer "
"than this will be truncated, sequences shorter will be padded.",
)
parser.add_argument(
"--test_max_target_length",
default=25,
type=int,
help="The maximum total input sequence length after tokenization. Sequences longer "
"than this will be truncated, sequences shorter will be padded.",
)
parser.add_argument("--sortish_sampler", action="store_true", default=False)
parser.add_argument("--logger_name", type=str, choices=["default", "wandb", "wandb_shared"], default="default")
parser.add_argument("--n_train", type=int, default=-1, required=False, help="# examples. -1 means use all.")
parser.add_argument("--n_val", type=int, default=-1, required=False, help="# examples. -1 means use all.")
parser.add_argument("--n_test", type=int, default=-1, required=False, help="# examples. -1 means use all.")
parser.add_argument("--label_smoothing", type=float, default=0.0, required=False)
parser.add_argument(
"--prefix",
type=str,
default=None,
help="Prefix added at the beginning of each text, typically used with T5-based models.",
)
parser.add_argument(
"--early_stopping_patience",
type=int,
default=-1,
required=False,
help="-1 means never early stop. early_stopping_patience is measured in validation checks, not epochs. So val_check_interval will effect it.",
)
parser.add_argument(
"--distributed-port", type=int, default=-1, required=False, help="Port number for distributed training."
)
parser.add_argument(
"--model_type",
choices=["rag_sequence", "rag_token", "bart", "t5"],
type=str,
help="RAG model type: sequence or token, if none specified, the type is inferred from the model_name_or_path",
)
return parser
def main(args, model=None) -> GenerativeQAModule:
Path(args.output_dir).mkdir(exist_ok=True)
if model is None:
model: GenerativeQAModule = GenerativeQAModule(args)
dataset = Path(args.data_dir).name
if (
args.logger_name == "default"
or args.fast_dev_run
or str(args.output_dir).startswith("/tmp")
or str(args.output_dir).startswith("/var")
):
logger = True # don't pollute wandb logs unnecessarily
elif args.logger_name == "wandb":
from pytorch_lightning.loggers import WandbLogger
project = os.environ.get("WANDB_PROJECT", dataset)
logger = WandbLogger(name=model.output_dir.name, project=project)
elif args.logger_name == "wandb_shared":
from pytorch_lightning.loggers import WandbLogger
logger = WandbLogger(name=model.output_dir.name, project=f"hf_{dataset}")
es_callback = (
get_early_stopping_callback(model.val_metric, args.early_stopping_patience)
if args.early_stopping_patience >= 0
else False
)
trainer: pl.Trainer = generic_train(
model,
args,
logging_callback=Seq2SeqLoggingCallback(),
checkpoint_callback=get_checkpoint_callback(args.output_dir, model.val_metric),
early_stopping_callback=es_callback,
logger=logger,
)
pickle_save(model.hparams, model.output_dir / "hparams.pkl")
if not args.do_predict:
return model
model.hparams.test_checkpoint = ""
checkpoints = list(sorted(glob.glob(os.path.join(args.output_dir, "*.ckpt"), recursive=True)))
if checkpoints:
model.hparams.test_checkpoint = checkpoints[-1]
trainer.resume_from_checkpoint = checkpoints[-1] # best checkpoint
trainer.logger.log_hyperparams(model.hparams)
# test() without a model tests using the best checkpoint automatically
trainer.test()
return model
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser = pl.Trainer.add_argparse_args(parser)
parser = GenerativeQAModule.add_model_specific_args(parser, os.getcwd())
args = parser.parse_args()
main(args)
+34
View File
@@ -0,0 +1,34 @@
# Add parent directory to python path to access lightning_base.py
export PYTHONPATH="../":"${PYTHONPATH}"
# A sample finetuning run, you need to specify data_dir, output_dir and model_name_or_path
# run ./examples/rag/finetune.sh --help to see all the possible options
python examples/rag/finetune.py \
--data_dir $DATA_DIR \
--output_dir $OUTPUT_DIR \
--model_name_or_path $MODLE_NAME_OR_PATH \
--model_type rag_sequence \
--fp16 \
--gpus 8 \
--do_train \
--do_predict \
--n_val -1 \
--val_check_interval 0.25 \
--train_batch_size 8 \
--eval_batch_size 1 \
--max_source_length 128 \
--max_target_length 25 \
--val_max_target_length 25 \
--test_max_target_length 25 \
--label_smoothing 0.1 \
--dropout 0.1 \
--attention_dropout 0.1 \
--weight_decay 0.001 \
--adam_epsilon 1e-08 \
--max_grad_norm 0.1 \
--lr_scheduler polynomial \
--learning_rate 3e-05 \
--num_train_epochs 100 \
--warmup_steps 500 \
--gradient_accumulation_steps 1
+47
View File
@@ -0,0 +1,47 @@
"""
This script reads DPR retriever training data and parses each datapoint. We save a line per datapoint.
Each line consists of the query followed by a tab-separated list of Wikipedia page titles constituting
positive contexts for a given query.
"""
import argparse
import json
from tqdm import tqdm
def main():
parser = argparse.ArgumentParser()
# Required parameters
parser.add_argument(
"--src_path",
type=str,
default="biencoder-nq-dev.json",
help="Path to raw DPR training data",
)
parser.add_argument(
"--evaluation_set",
type=str,
help="where to store parsed evaluation_set file",
)
parser.add_argument(
"--gold_data_path",
type=str,
help="where to store parsed gold_data_path file",
)
args = parser.parse_args()
with open(args.src_path, "r") as src_file, open(args.evaluation_set, "w") as eval_file, open(
args.gold_data_path, "w"
) as gold_file:
dpr_records = json.load(src_file)
for dpr_record in tqdm(dpr_records):
question = dpr_record["question"]
contexts = [context["title"] for context in dpr_record["positive_ctxs"]]
eval_file.write(question + "\n")
gold_file.write("\t".join(contexts) + "\n")
if __name__ == "__main__":
main()
+174
View File
@@ -0,0 +1,174 @@
import linecache
import re
import string
from collections import Counter
from logging import getLogger
from pathlib import Path
from typing import Dict, List
import torch
from torch.utils.data import Dataset
from examples.seq2seq.utils import SortishSampler, trim_batch
from transformers import BartTokenizer, T5Tokenizer
def encode_line(tokenizer, line, max_length, padding_side, pad_to_max_length=True, return_tensors="pt"):
extra_kw = {"add_prefix_space": True} if isinstance(tokenizer, BartTokenizer) else {}
tokenizer.padding_side = padding_side
return tokenizer(
[line],
max_length=max_length,
padding="max_length" if pad_to_max_length else None,
truncation=True,
return_tensors=return_tensors,
add_special_tokens=True,
**extra_kw,
)
class Seq2SeqDataset(Dataset):
def __init__(
self,
tokenizer,
data_dir,
max_source_length,
max_target_length,
type_path="train",
n_obs=None,
src_lang=None,
tgt_lang=None,
prefix="",
):
super().__init__()
self.src_file = Path(data_dir).joinpath(type_path + ".source")
self.tgt_file = Path(data_dir).joinpath(type_path + ".target")
self.src_lens = self.get_char_lens(self.src_file)
self.max_source_length = max_source_length
self.max_target_length = max_target_length
assert min(self.src_lens) > 0, f"found empty line in {self.src_file}"
self.tokenizer = tokenizer
self.prefix = prefix
if n_obs is not None:
self.src_lens = self.src_lens[:n_obs]
self.pad_token_id = self.tokenizer.pad_token_id
self.src_lang = src_lang
self.tgt_lang = tgt_lang
def __len__(self):
return len(self.src_lens)
def __getitem__(self, index) -> Dict[str, torch.Tensor]:
index = index + 1 # linecache starts at 1
source_line = self.prefix + linecache.getline(str(self.src_file), index).rstrip("\n")
tgt_line = linecache.getline(str(self.tgt_file), index).rstrip("\n")
assert source_line, f"empty source line for index {index}"
assert tgt_line, f"empty tgt line for index {index}"
# Need to add eos token manually for T5
if isinstance(self.tokenizer, T5Tokenizer):
source_line += self.tokenizer.eos_token
tgt_line += self.tokenizer.eos_token
# Pad source to the left and target to the right
source_inputs = encode_line(self.tokenizer, source_line, self.max_source_length, "right") # "left")
target_inputs = encode_line(self.tokenizer, tgt_line, self.max_target_length, "right")
source_ids = source_inputs["input_ids"].squeeze()
target_ids = target_inputs["input_ids"].squeeze()
src_mask = source_inputs["attention_mask"].squeeze()
return {
"input_ids": source_ids,
"attention_mask": src_mask,
"decoder_input_ids": target_ids,
}
@staticmethod
def get_char_lens(data_file):
return [len(x) for x in Path(data_file).open().readlines()]
def collate_fn(self, batch) -> Dict[str, torch.Tensor]:
input_ids = torch.stack([x["input_ids"] for x in batch])
masks = torch.stack([x["attention_mask"] for x in batch])
target_ids = torch.stack([x["decoder_input_ids"] for x in batch])
pad_token_id = self.pad_token_id
y = trim_batch(target_ids, pad_token_id)
source_ids, source_mask = trim_batch(input_ids, pad_token_id, attention_mask=masks)
batch = {
"input_ids": source_ids,
"attention_mask": source_mask,
"decoder_input_ids": y,
}
return batch
def make_sortish_sampler(self, batch_size):
return SortishSampler(self.src_lens, batch_size)
logger = getLogger(__name__)
def normalize_answer(s):
"""Lower text and remove punctuation, articles and extra whitespace."""
def remove_articles(text):
return re.sub(r"\b(a|an|the)\b", " ", text)
def white_space_fix(text):
return " ".join(text.split())
def remove_punc(text):
exclude = set(string.punctuation)
return "".join(ch for ch in text if ch not in exclude)
def lower(text):
return text.lower()
return white_space_fix(remove_articles(remove_punc(lower(s))))
def f1_score(prediction, ground_truth):
prediction_tokens = normalize_answer(prediction).split()
ground_truth_tokens = normalize_answer(ground_truth).split()
common = Counter(prediction_tokens) & Counter(ground_truth_tokens)
num_same = sum(common.values())
if num_same == 0:
return 0
precision = 1.0 * num_same / len(prediction_tokens)
recall = 1.0 * num_same / len(ground_truth_tokens)
f1 = (2 * precision * recall) / (precision + recall)
return f1
def exact_match_score(prediction, ground_truth):
return normalize_answer(prediction) == normalize_answer(ground_truth)
def calculate_exact_match(output_lns: List[str], reference_lns: List[str]) -> Dict:
assert len(output_lns) == len(reference_lns)
em = 0
for hypo, pred in zip(output_lns, reference_lns):
em += exact_match_score(hypo, pred)
if len(output_lns) > 0:
em /= len(output_lns)
return {"em": em}
def is_rag_model(model_prefix):
return model_prefix.startswith("rag")
def set_extra_model_params(extra_params, hparams, config):
equivalent_param = {p: p for p in extra_params}
# T5 models don't have `dropout` param, they have `dropout_rate` instead
equivalent_param["dropout"] = "dropout_rate"
for p in extra_params:
if getattr(hparams, p, None):
if not hasattr(config, p) and not hasattr(config, equivalent_param[p]):
logger.info("config doesn't have a `{}` attribute".format(p))
delattr(hparams, p)
continue
set_p = p if hasattr(config, p) else equivalent_param[p]
setattr(config, set_p, getattr(hparams, p))
delattr(hparams, p)
return hparams, config
+1 -1
View File
@@ -89,7 +89,7 @@ extras["onnxruntime"] = ["onnxruntime>=1.4.0", "onnxruntime-tools>=1.4.2"]
extras["serving"] = ["pydantic", "uvicorn", "fastapi", "starlette"]
extras["all"] = extras["serving"] + ["tensorflow", "torch"]
extras["testing"] = ["pytest", "pytest-xdist", "timeout-decorator", "psutil", "parameterized"]
extras["testing"] = ["pytest", "pytest-xdist", "timeout-decorator", "psutil", "parameterized", "faiss", "datasets"]
# sphinx-rtd-theme==0.5.0 introduced big changes in the style.
extras["docs"] = ["recommonmark", "sphinx", "sphinx-markdown-tables", "sphinx-rtd-theme==0.4.3", "sphinx-copybutton"]
extras["quality"] = ["black >= 20.8b1", "isort >= 5", "flake8 >= 3.8.3"]
+8
View File
@@ -84,6 +84,7 @@ from .file_utils import (
cached_path,
is_apex_available,
is_datasets_available,
is_faiss_available,
is_psutil_available,
is_py3nvml_available,
is_tf_available,
@@ -712,6 +713,13 @@ if is_tf_available():
from .trainer_tf import TFTrainer
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 .retrieval_rag import RagRetriever
from .tokenization_rag import RagDefaultTokenizer
if not is_tf_available() and not is_torch_available():
logger.warning(
"Neither PyTorch nor TensorFlow >= 2.0 have been found."
+158
View File
@@ -0,0 +1,158 @@
# coding=utf-8
# Copyright 2020, The RAG Authors and The HuggingFace Inc. team.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
""" RAG model configuration """
from .configuration_utils import PretrainedConfig
from .file_utils import add_start_docstrings_to_callable
RAG_PRETRAINED_CONFIG_ARCHIVE_MAP = {
"facebook/rag-sequence-nq": "TBA",
"facebook/rag-token-nq": "TBA",
}
RAG_CONFIG_DOC = r"""
:class:`~transformers.RagConfig` is the configuration class to store the configuration of a `RagModel`.
Args:
vocab_size (:obj:`int`, optional, defaults to ``None``):
Vocabulary size of the underlying generator model.
is_encoder_decoder (:obj:`bool`, `optional`, defaults to :obj:`False`):
Whether the model is used as an encoder/decoder or not.
title_sep (:obj:`str`, optional, defaults to ``" / "``):
Separator inserted between the title and the text of the retrieved document when running
`:func:`~transformers.RagModel.contextualize``.
doc_sep (:obj:`str`, optional, defaults to ``" // "``):
Separator inserted between the the text of the retrieved document and the original input when running
`:func:`~transformers.RagModel.contextualize``.
n_docs (:obj:`int`, optional, defaults to ``5``):
Number of retrieved docs.
max_combined_length (:int:`bool`, optional, defaults to ``300``):
Max length of contextualized input returned by `:func:`~transformers.RagModel.contextualize``.
retrieval_vector_size (:obj:`int`, optional, defaults to ``768``):
Dimensionality of the document embeddings indexed by the ``retriever``.
retrieval_batch_size (:obj:`int`, optional, defaults to ``8``):
Retrieval batch size - the number of queries issues concurrently to the faiss index excapsulated
by the ``retriever``.
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`
- ``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()``).
dataset_split (:obj:`str`, optional, defaults to ``train``)
Which split of the ``dataset`` to load.
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`
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``.
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``.
Args linked to the tokenizer - they have to be compatible with equivalent parameters of the ``generator``:
prefix (:obj:`str`, `optional`):
A specific prompt that should be added at the beginning of each text before calling the model.
bos_token_id (:obj:`int`, `optional`):
The id of the `beginning-of-stream` token.
pad_token_id (:obj:`int`, `optional`):
The id of the `padding` token.
eos_token_id (:obj:`int`, `optional`)"
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.
"""
@add_start_docstrings_to_callable(RAG_CONFIG_DOC)
class RagConfig(PretrainedConfig):
model_type = "rag"
def __init__(
self,
vocab_size=None,
is_encoder_decoder=True,
prefix=None,
bos_token_id=None,
pad_token_id=None,
eos_token_id=None,
decoder_start_token_id=None,
title_sep=" / ",
doc_sep=" // ",
n_docs=5,
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_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,
**kwargs
):
super().__init__(**kwargs)
self.vocab_size = vocab_size
self.is_encoder_decoder = is_encoder_decoder
self.prefix = prefix
self.bos_token_id = bos_token_id
self.pad_token_id = pad_token_id
self.eos_token_id = eos_token_id
self.decoder_start_token_id = decoder_start_token_id
self.title_sep = title_sep
self.doc_sep = doc_sep
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.index_path = index_path
self.dummy = dummy
self.pretrained_question_encoder_tokenizer_name_or_path = pretrained_question_encoder_tokenizer_name_or_path
self.pretrained_question_encoder_name_or_path = pretrained_question_encoder_name_or_path
self.pretrained_generator_tokenizer_name_or_path = pretrained_generator_tokenizer_name_or_path
self.pretrained_generator_name_or_path = pretrained_generator_name_or_path
+13
View File
@@ -119,6 +119,15 @@ try:
except ImportError:
_has_apex = False
try:
import faiss # noqa: F401
_faiss_available = True
except ImportError:
_faiss_available = False
default_cache_path = os.path.join(torch_cache_home, "transformers")
@@ -171,6 +180,10 @@ def is_apex_available():
return _has_apex
def is_faiss_available():
return _faiss_available
def add_start_docstrings(*docstr):
def docstring_decorator(fn):
fn.__doc__ = "".join(docstr) + (fn.__doc__ if fn.__doc__ is not None else "")
+9 -1
View File
@@ -399,7 +399,15 @@ class GenerationMixin:
# get encoder and store encoder outputs
encoder = self.get_encoder()
encoder_outputs: ModelOutput = encoder(input_ids, attention_mask=attention_mask, return_dict=True)
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)
# Expand input ids if num_beams > 1 or num_return_sequences > 1
if num_return_sequences > 1 or num_beams > 1:
+912
View File
@@ -0,0 +1,912 @@
# coding=utf-8
# Copyright 2020, The RAG Authors and The HuggingFace Inc. team.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""RAG model implementation."""
import copy
import os
from dataclasses import dataclass
from typing import List, Optional, Tuple
import torch
from .configuration_auto import AutoConfig
from .configuration_dpr import DPRConfig
from .configuration_rag import RagConfig
from .configuration_utils import PretrainedConfig
from .file_utils import add_start_docstrings_to_callable, replace_return_docstrings
from .modeling_auto import AutoModelForSeq2SeqLM
from .modeling_dpr import DPRQuestionEncoder
from .modeling_outputs import ModelOutput
from .modeling_t5 import T5ForConditionalGeneration
from .modeling_utils import PreTrainedModel
from .retrieval_rag import RagRetriever
from .utils import logging
logger = logging.get_logger(__name__)
_CONFIG_FOR_DOC = "RagConfig"
@dataclass
class BaseModelOutputWithDocs(ModelOutput):
"""
Base class for model's outputs, with potential hidden states and attentions.
Args:
last_hidden_state (:obj:`torch.FloatTensor` of shape :obj:`(batch_size, sequence_length, hidden_size)`):
Sequence of hidden-states at the output of the last layer of the model.
hidden_states (:obj:`tuple(torch.FloatTensor)`, `optional`, returned when ``output_hidden_states=True`` is passed or when ``config.output_hidden_states=True``):
Tuple of :obj:`torch.FloatTensor` (one for the output of the embeddings + one for the output of each layer)
of shape :obj:`(batch_size, sequence_length, hidden_size)`.
Hidden-states of the model at the output of each layer plus the initial embedding outputs.
attentions (:obj:`tuple(torch.FloatTensor)`, `optional`, returned when ``output_attentions=True`` is passed or when ``config.output_attentions=True``):
Tuple of :obj:`torch.FloatTensor` (one for each layer) of shape
:obj:`(batch_size, num_heads, sequence_length, sequence_length)`.
Attentions weights after the attention softmax, used to compute the weighted average in the self-attention
heads.
doc_scores (:obj:`torch.FloatTensor` of shape :obj:`(batch_size, config.n_docs)`):
Scores of retrieved documents.
"""
last_hidden_state: torch.FloatTensor
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
attentions: Optional[Tuple[torch.FloatTensor]] = None
doc_scores: Optional[torch.FloatTensor] = None
@dataclass
class Seq2SeqLMOutputWithDocs(ModelOutput):
"""
Outputs for sequence-to-sequence language models with retrieval in the loop.
Args:
loss (:obj:`torch.FloatTensor` of shape :obj:`(1,)`, `optional`, returned when :obj:`labels` is provided):
Languaged modeling loss.
logits (:obj:`torch.FloatTensor` of shape :obj:`(batch_size, sequence_length, config.vocab_size)` if the ``logits are marginalized or :obj:`(batch_size * config.n_docs, sequence_length, config.vocab_size)` if they aren't):
Prediction scores of the language modeling head (scores for each vocabulary token before SoftMax).
past_key_values (:obj:`List[torch.FloatTensor]`, `optional`, returned when ``use_cache=True`` is passed or when ``config.use_cache=True``):
List of :obj:`torch.FloatTensor` of length :obj:`config.n_layers`, with each tensor of shape
:obj:`(2, batch_size, num_heads, sequence_length, embed_size_per_head)`).
Contains pre-computed hidden-states (key and values in the attention blocks) of the decoder that can be
used (see ``past_key_values`` input) to speed up sequential decoding.
decoder_hidden_states (:obj:`tuple(torch.FloatTensor)`, `optional`, returned when ``output_hidden_states=True`` is passed or when ``config.output_hidden_states=True``):
Tuple of :obj:`torch.FloatTensor` (one for the output of the embeddings + one for the output of each layer)
of shape :obj:`(batch_size, sequence_length, hidden_size)`.
Hidden-states of the decoder at the output of each layer plus the initial embedding outputs.
decoder_attentions (:obj:`tuple(torch.FloatTensor)`, `optional`, returned when ``output_attentions=True`` is passed or when ``config.output_attentions=True``):
Tuple of :obj:`torch.FloatTensor` (one for each layer) of shape
:obj:`(batch_size, num_heads, sequence_length, sequence_length)`.
Attentions weights of the decoder, after the attention softmax, used to compute the weighted average in the
self-attention heads.
encoder_last_hidden_state (:obj:`torch.FloatTensor` of shape :obj:`(batch_size, sequence_length, hidden_size)`, `optional`):
Sequence of hidden-states at the output of the last layer of the encoder of the model.
encoder_hidden_states (:obj:`tuple(torch.FloatTensor)`, `optional`, returned when ``output_hidden_states=True`` is passed or when ``config.output_hidden_states=True``):
Tuple of :obj:`torch.FloatTensor` (one for the output of the embeddings + one for the output of each layer)
of shape :obj:`(batch_size, sequence_length, hidden_size)`.
Hidden-states of the encoder at the output of each layer plus the initial embedding outputs.
encoder_attentions (:obj:`tuple(torch.FloatTensor)`, `optional`, returned when ``output_attentions=True`` is passed or when ``config.output_attentions=True``):
Tuple of :obj:`torch.FloatTensor` (one for each layer) of shape
:obj:`(batch_size, num_heads, sequence_length, sequence_length)`.
Attentions weights of the encoder, after the attention softmax, used to compute the weighted average in the
self-attention heads.
doc_scores (:obj:`torch.FloatTensor` of shape :obj:`(batch_size, config.n_docs)`):
Scores of retrieved documents.
"""
loss: Optional[torch.FloatTensor] = None
logits: torch.FloatTensor = None
past_key_values: Optional[List[torch.FloatTensor]] = None
decoder_hidden_states: Optional[Tuple[torch.FloatTensor]] = None
decoder_attentions: Optional[Tuple[torch.FloatTensor]] = None
encoder_last_hidden_state: Optional[torch.FloatTensor] = None
encoder_hidden_states: Optional[Tuple[torch.FloatTensor]] = None
encoder_attentions: Optional[Tuple[torch.FloatTensor]] = None
doc_scores: Optional[torch.FloatTensor] = None
# Reshape from [batch_size, n_docs, dims] to [batch_size * n_docs, dims]
def _stack_ctxt(tensor):
return tensor.view(-1, *tensor.shape[2:])
# Reshape from [batch_size * n_docs, dims] to [batch_size, n_docs, dims]
def _unstack_ctxt(tensor, n_docs):
return tensor.view(-1, n_docs, *tensor.shape[1:])
RAG_START_DOCSTRING = r"""
RAG 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.
The model is compatible with :class:`~transformers.DPRQuestionEncoder` as the ``question_encoder``. As for the ``generator``,
two compatible architectures have been tested: :class:`~transformers.BartForConditionalGeneration`
and :class:`~transformers.T5ForConditionalGeneration`.
This model is a PyTorch `torch.nn.Module <https://pytorch.org/docs/stable/nn.html#torch.nn.Module>`_ sub-class.
Use it as a regular PyTorch Module and refer to the PyTorch documentation for all matter related to general
usage and behavior.
Args:
config (:class:`~transformers.RagConfig`): Model configuration class with all the parameters of the model.
Initializing with a config file does not load the weights associated with the model, only the configuration.
Check out the :meth:`~transformers.PreTrainedModel.from_pretrained` method to load the model weights.
"""
RAG_CORE_DOCSTRING = r"""
A base RAG model calculating raw sequence logits and document retrieval scores.
The model takes a question encoder and a generator as inputs to the constructor, so it can be a base
for various RAG architectures encapsualting different retrievers and generators.
Args:
config (:class:`~transformers.RagConfig`): Model configuration class with all the parameters of the model.
Initializing with a config file does not load the weights associated with the model, only the configuration.
Check out the :meth:`~transformers.PreTrainedModel.from_pretrained` method to load the model weights.
question_encoder (:class:`transformers.PreTrainedModel`):
An encoder model compatible with the faiss index encapsulated by the ``retriever``.
generator (:class:`transformers.PreTrainedModel`):
A seq2seq model used as the generator in the RAG architecture.
"""
RAG_FORWARD_INPUTS_DOCSTRING = r"""
Args:
input_ids (:obj:`torch.LongTensor` of shape :obj:`(batch_size, sequence_length)`):
Indices of input sequence tokens in the vocabulary.
:class:`~transformers.RagConfig`, used to initialize the model, specifies which generator to use, it also specifies a compatible
generator tokenizer. Use that tokenizer class to obtain the indices.
retriever (:class:`~transformers.RagRetriever`):
A retriever class encapsulating a faiss index queried to obtain context documents for current inputs.
attention_mask (:obj:`torch.Tensor` of shape :obj:`(batch_size, sequence_length)`, `optional`, defaults to :obj:`None`):
Mask to avoid performing attention on padding token indices in input_ids.
Mask values selected in ``[0, 1]``:
``1`` for tokens that are NOT MASKED, ``0`` for MASKED tokens.
encoder_outputs (:obj:`tuple(tuple(torch.FloatTensor)`, `optional`, defaults to :obj:`None`):
Tuple consists of (`last_hidden_state`, `optional`: `hidden_states`, `optional`: `attentions`, `doc_scores`)
`last_hidden_state` of shape :obj:`(batch_size, n_docs * sequence_length, hidden_size)` is a sequence of hidden-states at the output of the last layer of the encoder.
`doc_scores` of shape :obj:`(batch_size, n_docs)` store retrieval scores of documents retrieved for each input in the batch.
Used by the (:class:`~transformers.RagToken`) model during decoding.
decoder_input_ids (:obj:`torch.LongTensor` of shape :obj:`(batch_size, target_sequence_length)`, `optional`, defaults to :obj:`None`):
Provide for generation tasks. `None` by default, constuct as per instructions for the generator model you're using with your RAG instance.
past_key_values (:obj:`tuple(tuple(torch.FloatTensor))`):
Tuple consists of two elements: ``encoder_outputs`` of the RAG model (see ``encoder_outputs``) and ``past_key_values`` of the underlying generator.
Can be used to speed up decoding. ``past_key_values`` are used in the (:class:`~transformers.RagToken`)
model during decoding.
use_cache (:obj:`bool`, `optional`, defaults to :obj:`True`):
If `use_cache` is True, ``past_key_values`` are returned and can be used to speed up decoding (see
``past_key_values``).
generator_kwargs (remaining dictionary of keyword arguments, `optional`):
Additional keyword arguments will be passed to the generator forward pass.
"""
RAG_LOSS_INPUTS_DOCSTRING = r"""
return_loss (:obj:`bool`, `optional`, defaults to :obj:`False`):
If :obj:`True`, computes the loss which is returned as part of the :class:`~transformers.file_utils.Seq2SeqLMOutputWithDocs`.
Otherwise, loss defaults to :obj:`None`.
reduce (:obj:`bool`, `optional`, defaults to :obj:`False`):
Only relevant if ``return_loss`` is set to :obj:`True`. If :obj:`True`, the NLL loss is reduced using the ``torch.Tensor.sum`` operation.
label_smoothing (:obj:`float`, `optional`, defaults to ``0.0``):
Only relevant if ``return_loss`` is set to :obj:`True`. Controls the ``epsilon`` parameter value for label smoothing in the loss calculation.
If set to ``0.0``, no label smoothing is performed.
"""
@add_start_docstrings_to_callable(RAG_CORE_DOCSTRING)
class RagModel(torch.nn.Module):
def __init__(
self,
config,
question_encoder,
generator,
):
super().__init__()
self.config = config
self.question_encoder = question_encoder
self.generator = generator
self.n_docs = self.config.n_docs
def contextualize(self, input_ids, retriever, print_docs=False):
"""
Adds context to every input in the batch by querying the retriever.
Args:
input_ids (:obj:`torch.LongTensor` of shape :obj:`(batch_size, sequence_length)`, `optional`):
The sequence used as a prompt for the generation. If :obj:`None` the method initializes
it as an empty :obj:`torch.LongTensor` of shape :obj:`(1,)`.
retriever (:class:`~transformers.RagRetriever`):
A retriever class encapsulating a faiss index queried to obtain context documents for current inputs.
print_docs (:obj:`bool`, `optional`, defaults to :obj:`True`):
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 three elements: contextualized ``input_ids``,
compatible ``attention_mask`` and scores of the retrieved documents.
"""
question_encoder_input_ids, input_strings = retriever.preprocess_query(input_ids, self.generator.config.prefix)
query_vectors = self.question_encoder(question_encoder_input_ids)[0]
doc_vectors, docs = retriever.retrieve(query_vectors.cpu().detach().to(torch.float32), n_docs=self.n_docs)
doc_vectors = doc_vectors.to(query_vectors)
doc_scores = torch.bmm(query_vectors.unsqueeze(1), doc_vectors.transpose(1, 2)).squeeze(1)
# T5 tokenizer doesn't add eos token by default even with add_special_tokens set to True
add_eos = (input_ids == self.config.eos_token_id).any() and isinstance(
self.generator, T5ForConditionalGeneration
)
input_ids, attention_mask = retriever.postprocess_docs(
doc_scores, docs, input_strings, add_eos, self.generator.config.prefix, print_docs
)
return input_ids, attention_mask, doc_scores
@add_start_docstrings_to_callable(RAG_FORWARD_INPUTS_DOCSTRING)
@replace_return_docstrings(output_type=Seq2SeqLMOutputWithDocs, config_class=_CONFIG_FOR_DOC)
def forward(
self,
input_ids,
retriever: RagRetriever,
attention_mask=None,
encoder_outputs=None,
decoder_input_ids=None,
past_key_values=None,
use_cache=None,
print_docs=False,
**generator_kwargs
):
r"""
print_docs (:obj:`bool`, `optional`, defaults to :obj:`True`):
If :obj:`True`, documents retrieved during the forward pass will be logged. Intended for debugging purposes.
Returns:
"""
# encoder_outputs are pre-computed during RAG-token generation
if encoder_outputs is not None:
doc_scores = encoder_outputs.doc_scores
else:
# Add context documents to input
input_ids, attention_mask, doc_scores = self.contextualize(input_ids, retriever, print_docs)
# Decoder input without context documents
if decoder_input_ids is not None:
decoder_input_ids = decoder_input_ids.repeat_interleave(self.n_docs, dim=0)
outputs = self.generator(
input_ids,
attention_mask=attention_mask,
encoder_outputs=encoder_outputs,
decoder_input_ids=decoder_input_ids,
past_key_values=past_key_values,
use_cache=use_cache,
**generator_kwargs,
)
return Seq2SeqLMOutputWithDocs(
loss=None,
logits=outputs.logits,
past_key_values=outputs.past_key_values,
decoder_hidden_states=outputs.decoder_hidden_states,
decoder_attentions=outputs.decoder_attentions,
encoder_last_hidden_state=outputs.encoder_last_hidden_state,
encoder_hidden_states=outputs.encoder_hidden_states,
encoder_attentions=outputs.encoder_attentions,
doc_scores=doc_scores,
)
class RAGEncoder(torch.nn.Module):
r"""
RAG is an encoder-decoder model, however, we don't exaplicitly implement an encoder and a decoder layes,
like it's done e.g. in BART and T5 implementations - for RAG these are encapsulated inside the generaotr instance.
This is a dummy model simulating RAG encoder output, need for compatibility with transformers generation code
"""
def __init__(self, rag_model: RagModel):
super().__init__()
self.rag_model = rag_model
def forward(self, input_ids=None, retriever=None, attention_mask=None, return_dict=True):
ctxt_input_ids, ctxt_attention_mask, doc_scores = self.rag_model.contextualize(input_ids, retriever)
encoder = self.rag_model.generator.get_encoder()
encoder_outputs = encoder(
input_ids=ctxt_input_ids, attention_mask=ctxt_attention_mask, return_dict=return_dict
)
# needed to satisfy assertion that encoder_outputs.last_hidden_state.shape[0] == batch_size in generation_utils
unstacked_x = _unstack_ctxt(encoder_outputs.last_hidden_state, self.rag_model.n_docs)
return BaseModelOutputWithDocs(
last_hidden_state=unstacked_x,
hidden_states=encoder_outputs.hidden_states,
attentions=ctxt_attention_mask,
doc_scores=doc_scores,
)
class PreTrainedRagModel(PreTrainedModel):
r"""
RAG models encapsulate two trainable components - a question encoder and a generator, but as such they don't have any trainable parameters.
We specialize `:func:`~transformers.PreTrainedModel.from_pretrained`` and `:func:`~transformers.PreTrainedModel.save_pretrained`` to reflect this.
"""
config_class = RagConfig
def __init__(
self,
config: RagConfig,
question_encoder: PreTrainedModel = None,
generator: PreTrainedModel = None,
):
super().__init__(config)
self.config = config
if question_encoder is None:
# TODO(piktus): To be replaced with AutoConfig / AutoModel once it supports DPRQuestionEncoder
question_encoder_config = DPRConfig.from_pretrained(self.config.pretrained_question_encoder_name_or_path)
question_encoder = DPRQuestionEncoder(question_encoder_config)
if generator is None:
generaotr_config = AutoConfig.from_pretrained(self.config.pretrained_generator_name_or_path)
generator = AutoModelForSeq2SeqLM.from_config(generaotr_config)
self.n_docs = self.config.n_docs
self._validate_configs_match(self.config, generator.config)
self.model = RagModel(config, question_encoder, generator)
@classmethod
def from_pretrained(cls, pretrained_model_name_or_path=None, **kwargs):
r"""
Instantiates a pretrained RAG model from a pre-trained model configuration. Since RAG doesn't have any trainable parameters
other than those encapsulated by the ``question _encoder`` and the ``generator``, we call `:func:`~transformers.PreTrainedModel.from_pretrained``
for the ``question_encoder`` and the ``generator`` respectively.
Parameters:
pretrained_model_name_or_path (:obj:`str`, `optional`):
A string specifying the model to be loaded. See :func:`~transformers.PreTrainedModel.from_pretrained`` for details.
config (:obj:`Union[PretrainedConfig, str]`, `optional`):
Can be either:
- an instance of a class derived from :class:`~transformers.PretrainedConfig`,
- a string valid as input to :func:`~transformers.PretrainedConfig.from_pretrained`.
See :func:`~transformers.PreTrainedModel.from_pretrained`` for more details.
generator_config (:obj:`str`, `optional`):
A string valid as input to :func:`~transformers.PretrainedConfig.from_pretrained`. Will be passed
to the :func:`~transformers.PreTrainedModel.from_pretrained`` function initializing the ``generator`` model.
question_encoder_config (:obj:`str`, `optional`):
A string valid as input to :func:`~transformers.PretrainedConfig.from_pretrained`. Will be passed
to the :func:`~transformers.PreTrainedModel.from_pretrained`` function initializing the ``question_encoder`` model.
kwargs (remaining dictionary of keyword arguments, `optional`):
`kwargs`` will be passed to the configuration class initialization function (:func:`~transformers.PretrainedConfig.from_pretrained`).
Each key of ``kwargs`` that corresponds to a configuration attribute will be used to override said attribute
with the supplied ``kwargs`` value. Remaining keys that do not correspond to any configuration
attribute will be passed to the underlying model's ``__init__`` function.
"""
config = kwargs.pop("config", None)
generator_config = kwargs.pop("generator_config", None)
question_encoder_config = kwargs.pop("question_encoder_config", None)
assert pretrained_model_name_or_path is not None or config is not None
if not isinstance(config, PretrainedConfig):
config = cls.config_class.from_pretrained(
config if config is not None else pretrained_model_name_or_path,
**kwargs,
)
assert pretrained_model_name_or_path is not None or config.pretrained_question_encoder_name_or_path is not None
pretrained_question_encoder_name_or_path = (
config.pretrained_question_encoder_name_or_path
if config.pretrained_question_encoder_name_or_path is not None
else os.path.join(pretrained_model_name_or_path, "question_encoder")
)
# TODO(piktus): To be replaced with AutoModel once it supports DPRQuestionEncoder
question_encoder = DPRQuestionEncoder.from_pretrained(
pretrained_question_encoder_name_or_path, config=question_encoder_config
)
assert pretrained_model_name_or_path is not None or config.pretrained_generator_name_or_path is not None
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(pretrained_model_name_or_path, "generator")
)
generator_kwargs = {}
if generator_config is not None:
setattr(generator_config, "return_dict", True)
generator_kwargs["config"] = generator_config
else:
generator_kwargs["return_dict"] = True
generator = AutoModelForSeq2SeqLM.from_pretrained(pretrained_generator_name_or_path, **generator_kwargs)
return cls(config, question_encoder, generator)
def save_pretrained(self, save_directory):
r"""
Save a model and its configuration file to a directory, so that it can be re-loaded using the
`:func:`~transformers.PreTrainedRagModel.from_pretrained`` class method.
Arguments:
save_directory (:obj:`str`):
Base directory to which to save. Will be created if it doesn't exist. The generator model
will be saved to save_directory/generator directory. The question encoder model will be saved
to save_directory/genquestion_encoder 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)
generator_output_dir = os.path.join(save_directory, "generator")
self.model.generator.save_pretrained(generator_output_dir)
qe_output_dir = os.path.join(save_directory, "question_encoder")
self.model.question_encoder.save_pretrained(qe_output_dir)
config = copy.deepcopy(self.config)
config.pretrained_generator_name_or_path = generator_output_dir
config.pretrained_question_encoder_name_or_path = qe_output_dir
config.vocab_size = self.model.generator.config.vocab_size
config.save_pretrained(save_directory)
def shift_tokens_left(self, input_ids, pad_token_id=None):
"""Shift input ids one token to the left, and add a pad to right"""
if pad_token_id is None:
pad_token_id = self.config.pad_token_id
return torch.cat([input_ids[:, 1:], input_ids.new(input_ids.shape[0], 1).fill_(pad_token_id)], 1)
def shift_tokens_right(self, input_ids, start_token_id=None):
"""Shift input ids one token to the right, and pad with start_token_id"""
if start_token_id is None:
start_token_id = self.config.decoder_start_token_id
shifted_input_ids = input_ids.new_zeros(input_ids.shape)
shifted_input_ids[:, 1:] = input_ids[:, :-1].clone()
shifted_input_ids[:, 0] = start_token_id
return shifted_input_ids
def _validate_configs_match(self, rag_config, gen_config):
assert rag_config.pad_token_id == gen_config.pad_token_id, "pad_token_id mismatch: {} vs. {}".format(
rag_config.pad_token_id, gen_config.pad_token_id
)
assert rag_config.bos_token_id == gen_config.bos_token_id, "bos_token_id mismatch: {} vs. {}".format(
rag_config.bos_token_id, gen_config.bos_token_id
)
assert rag_config.eos_token_id == gen_config.eos_token_id, "eos_token_id mismatch: {} vs. {}".format(
rag_config.eos_token_id, gen_config.eos_token_id
)
assert (
rag_config.decoder_start_token_id == gen_config.decoder_start_token_id
), "decoder_start_token_id mismatch: {} vs. {}".format(
rag_config.decoder_start_token_id, gen_config.decoder_start_token_id
)
assert (
rag_config.is_encoder_decoder == gen_config.is_encoder_decoder
), "pad_token_id mismatch: {} vs. {}".format(rag_config.is_encoder_decoder, gen_config.is_encoder_decoder)
assert rag_config.vocab_size == gen_config.vocab_size, "vocab_size mismatch: {} vs. {}".format(
rag_config.vocab_size, gen_config.vocab_size
)
@add_start_docstrings_to_callable(
"""A RAG-sequence model impementation. It performs RAG-sequence specific marginalization in the forward pass
and specializes some of the functions of :class:`~transformers.PreTrainedModel` to enable RAG-sequence generation.
""",
RAG_START_DOCSTRING,
)
class RagSequence(PreTrainedRagModel):
base_model_prefix = "rag_sequence"
@add_start_docstrings_to_callable(RAG_FORWARD_INPUTS_DOCSTRING, RAG_LOSS_INPUTS_DOCSTRING)
@replace_return_docstrings(output_type=Seq2SeqLMOutputWithDocs, config_class=_CONFIG_FOR_DOC)
def forward(
self,
input_ids,
retriever,
attention_mask=None,
encoder_outputs=None,
decoder_input_ids=None,
past_key_values=None,
use_cache=None,
return_loss=False,
reduce=False,
label_smoothing=0.0,
score=False,
**generator_kwargs
):
r"""
score (:obj:`bool`, `optional`, defaults to :obj:`False`):
A flag passed as an argument to `:func:`~transformers.RagSequence.get_nll``. If :obj:`True`,
we exclude the BOS token's score while scoring the sequence.
Returns:
"""
if return_loss:
use_cache = False
outputs = self.model(
input_ids,
retriever,
attention_mask,
encoder_outputs,
decoder_input_ids,
past_key_values,
use_cache,
**generator_kwargs,
)
if return_loss:
assert decoder_input_ids is not None
loss = self.get_nll(
outputs.logits,
outputs.doc_scores,
decoder_input_ids,
reduce=reduce,
epsilon=label_smoothing,
score=score,
)
return Seq2SeqLMOutputWithDocs(
loss=loss,
logits=outputs.logits,
past_key_values=outputs.past_key_values,
decoder_hidden_states=outputs.decoder_hidden_states,
decoder_attentions=outputs.decoder_attentions,
encoder_last_hidden_state=outputs.encoder_last_hidden_state,
encoder_hidden_states=outputs.encoder_hidden_states,
encoder_attentions=outputs.encoder_attentions,
doc_scores=outputs.doc_scores,
)
return Seq2SeqLMOutputWithDocs(
loss=None,
logits=outputs.logits,
past_key_values=outputs.past_key_values,
decoder_hidden_states=outputs.decoder_hidden_states,
decoder_attentions=outputs.decoder_attentions,
encoder_last_hidden_state=outputs.encoder_last_hidden_state,
encoder_hidden_states=outputs.encoder_hidden_states,
encoder_attentions=outputs.encoder_attentions,
doc_scores=outputs.doc_scores,
)
def generate(
self,
input_ids,
retriever,
dedup=True,
print_docs=False,
num_return_sequences=1,
num_beams=1,
attention_mask=None,
**kwargs
):
"""
Implements RAG sequence "thorough" decoding.
Args:
input_ids (:obj:`torch.LongTensor` of shape :obj:`(batch_size, sequence_length)`, `optional`):
The sequence used as a prompt for the generation. If :obj:`None` the method initializes
it as an empty :obj:`torch.LongTensor` of shape :obj:`(1,)`.
retriever (:class:`~transformers.RagRetriever`):
A retriever class encapsulating a faiss index queried to obtain context documents for current inputs.
dedup (:obj:`bool`, `optional`, defaults to :obj:`True`):
Controls whether we want to deduplicate the generations from different context documents for a given input.
Has to be set to :obj:`False` if used while training with distributed backend.
print_docs (:obj:`bool`, `optional`, defaults to :obj:`True`):
If :obj:`True`, documents retrieved during the forward pass will be printed out. Intended for debugging purposes.
num_return_sequences(:obj:`int`, `optional`, defaults to 1):
The number of independently computed returned sequences for each element in the batch. Note that this is not the value
we pass to the ``generator``'s `:func:`~transformers.PreTrainedModel.generate`` function, where we set ``num_return_sequences``
to `num_beams`.
num_beams (:obj:`int`, `optional`, defaults to ``1``):
Number of beams for beam search. ``1`` means no beam search.
attention_mask (:obj:`torch.LongTensor` of shape :obj:`(batch_size, sequence_length)`, `optional`, defaults to :obj:`None`):
Mask to avoid performing attention on padding token indices. Mask values are in ``[0, 1]``, 1 for
tokens that are not masked, and 0 for masked tokens.
kwargs:
Additional kwargs will be passed to the the ``generator``'s `:func:`~transformers.PreTrainedModel.generate`` function call.
Return:
:obj:`torch.LongTensor` of shape :obj:`(batch_size * num_return_sequences, sequence_length)`:
The generated sequences. The second dimension (sequence_length) is either equal to :obj:`max_length` or
shorter if all batches finished early due to the :obj:`eos_token_id`.
"""
def _get_unique_rows(_input_ids):
return torch.stack(list({str(k.tolist()): k for k in _input_ids}.values()))
ctxt_input_ids, _, _ = self.model.contextualize(input_ids, retriever, print_docs=print_docs)
rag_num_return_sequences = num_return_sequences
hypos = []
for index in range(len(input_ids)):
# first, generate beams from documents:
generator_input_ids = ctxt_input_ids[index * self.n_docs : (index + 1) * self.n_docs] # (n_docs, max_len)
output_sequences = self.model.generator.generate(
generator_input_ids, num_return_sequences=num_beams, num_beams=num_beams, attention_mask=None, **kwargs
) # n_docs * n_beam, tgt_len
if dedup:
output_sequences = _get_unique_rows(output_sequences) # dedup, max_output_len
# then, run model forwards to get nll scores:
new_input_ids = input_ids[index : index + 1].repeat(len(output_sequences), 1)
outputs = self.forward(
new_input_ids, retriever=retriever, decoder_input_ids=output_sequences, return_loss=True, score=True
)
top_cand_inds = (-outputs["loss"]).topk(rag_num_return_sequences)[1]
if logging.get_verbosity() == logging.DEBUG:
output_strings = self.model.generator_tokenizer.batch_decode(output_sequences)
logger.debug("Hypos with scores:")
for score, hypo in zip(outputs.loss, output_strings):
logger.debug("\t{} {}".format(score, hypo))
hypos.append(output_sequences[top_cand_inds])
return self._cat_and_pad(hypos, pad_token_id=self.config.pad_token_id)
def get_nll(self, seq_logits, doc_scores, target, reduce=False, epsilon=0.0, score=False):
target = self.shift_tokens_left(target)
# bos_token_id is None for T5
use_bos = self.config.bos_token_id is not None and target[:, 0].eq(self.config.bos_token_id).all()
def _mask_pads(ll, smooth_obj):
pad_mask = target.eq(self.config.pad_token_id)
if pad_mask.any():
ll.masked_fill_(pad_mask, 0.0)
smooth_obj.masked_fill_(pad_mask, 0.0)
return ll.squeeze(-1), smooth_obj.squeeze(-1)
seq_logprobs = torch.nn.functional.log_softmax(seq_logits, dim=-1).view(
seq_logits.shape[0] // self.n_docs, self.n_docs, -1, seq_logits.size(-1)
) # batch_size x n_docs x tgt_len x dim
doc_logprobs = torch.nn.functional.log_softmax(doc_scores, dim=1).unsqueeze(-1).unsqueeze(-1)
# RAG-sequence marginaliation
first_token_scores = seq_logprobs[:, :, :1, :]
second_token_scores = seq_logprobs[:, :, 1:2, :]
remainder = seq_logprobs[:, :, 2:, :]
rag_logprobs = torch.cat([first_token_scores, second_token_scores + doc_logprobs, remainder], dim=2)
# calcualate loss
target = target.unsqueeze(1).unsqueeze(-1).repeat(1, self.n_docs, 1, 1)
assert target.dim() == rag_logprobs.dim()
ll = rag_logprobs.gather(dim=-1, index=target)
smooth_obj = rag_logprobs.sum(dim=-1, keepdim=True) # total sum of all (normalised) logits
ll, smooth_obj = _mask_pads(ll, smooth_obj)
# sum over tokens, exclude bos while scoring
ll = ll[:, :, 1:].sum(2) if score and use_bos else ll.sum(2)
smooth_obj = smooth_obj.sum(2)
ll = ll.logsumexp(1) # logsumexp over docs
smooth_obj = smooth_obj.logsumexp(1)
nll_loss = -ll
smooth_loss = -smooth_obj
if reduce:
nll_loss = nll_loss.sum()
smooth_loss = smooth_loss.sum()
eps_i = epsilon / rag_logprobs.size(-1)
loss = (1.0 - epsilon) * nll_loss + eps_i * smooth_loss
return loss
@staticmethod
def _cat_and_pad(tensors, pad_token_id):
output = (
tensors[0].new(sum([t.shape[0] for t in tensors]), max([t.shape[1] for t in tensors])).fill_(pad_token_id)
)
ind = 0
for t in tensors:
output[ind : ind + t.shape[0], : t.shape[1]] = t
ind += t.shape[0]
return output
@add_start_docstrings_to_callable(
"""A RAG-token model impementation. It performs RAG-token specific marginalization in the forward pass
and specializes some of the functions of :class:`~transformers.PreTrainedModel` to enable RAG-token generation.
""",
RAG_START_DOCSTRING,
)
class RagToken(PreTrainedRagModel):
base_model_prefix = "rag_token"
def adjust_logits_during_generation(self, logits, cur_len, max_length):
return self.model.generator.adjust_logits_during_generation(logits, cur_len, max_length)
def prepare_inputs_for_generation(
self, decoder_input_ids, past, attention_mask, use_cache, encoder_outputs, **kwargs
):
last_hidden_state = encoder_outputs["last_hidden_state"]
doc_scores = encoder_outputs["doc_scores"]
attention_mask = encoder_outputs["attentions"]
beam_size = decoder_input_ids.shape[0] // doc_scores.shape[0]
doc_scores = doc_scores.repeat_interleave(beam_size, dim=0) # batch_size -> batch_size * beam_size
attention_mask = attention_mask.repeat_interleave(beam_size, dim=0) # batch_size -> batch_size * beam_size
encoder_outputs = BaseModelOutputWithDocs(
last_hidden_state=_stack_ctxt(last_hidden_state),
hidden_states=encoder_outputs.hidden_states,
attentions=attention_mask,
doc_scores=doc_scores,
)
print_docs = getattr(kwargs, "print_docs", False)
return {
"input_ids": None,
"retriever": kwargs["retriever"],
"encoder_outputs": encoder_outputs,
"attention_mask": attention_mask,
"decoder_input_ids": decoder_input_ids,
"past_key_values": past,
"use_cache": use_cache,
"marginalize": True,
"print_docs": print_docs,
}
@staticmethod
def _reorder_cache(past, beam_idx):
"""Reorders cache for generation. BART-inspired but we need to take care of the extra dimension for docs"""
def _reorder_stacked(t):
n_docs = t.shape[0] // beam_idx.shape[0]
t = _unstack_ctxt(t, n_docs).index_select(0, beam_idx)
return _stack_ctxt(t)
def _reorder_buffer(attn_cache):
for k, input_buffer_k in attn_cache.items():
if input_buffer_k is not None:
attn_cache[k] = _reorder_stacked(input_buffer_k)
return attn_cache
reordered_past = []
for layer_past in past:
# get the correct batch idx from decoder layer's batch dim for cross and self-attn
layer_past_new = {attn_key: _reorder_buffer(attn_cache) for attn_key, attn_cache in layer_past.items()}
reordered_past.append(layer_past_new)
return reordered_past
def marginalize(self, seq_logits, doc_scores):
# RAG-token marginalization
seq_logprobs = torch.nn.functional.log_softmax(seq_logits, dim=-1).view(
seq_logits.shape[0] // self.n_docs, self.n_docs, -1, seq_logits.size(-1)
)
doc_logprobs = torch.log_softmax(doc_scores, dim=1)
log_prob_sum = seq_logprobs + doc_logprobs.unsqueeze(-1).unsqueeze(-1)
return torch.logsumexp(log_prob_sum, dim=1)
@add_start_docstrings_to_callable(RAG_FORWARD_INPUTS_DOCSTRING, RAG_LOSS_INPUTS_DOCSTRING)
@replace_return_docstrings(output_type=Seq2SeqLMOutputWithDocs, config_class=_CONFIG_FOR_DOC)
def forward(
self,
input_ids,
retriever,
attention_mask=None,
encoder_outputs=None,
decoder_input_ids=None,
past_key_values=None,
use_cache=None,
return_loss=False,
reduce=False,
label_smoothing=0.0,
marginalize=False,
**generator_kwargs
):
r"""
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`.
Returns:
"""
if return_loss:
use_cache = False
outputs = self.model(
input_ids,
retriever,
attention_mask,
encoder_outputs,
decoder_input_ids,
past_key_values,
use_cache,
**generator_kwargs,
)
if return_loss:
assert decoder_input_ids is not None
loss = self.get_nll(
outputs.logits, outputs.doc_scores, decoder_input_ids, reduce=reduce, epsilon=label_smoothing
)
return Seq2SeqLMOutputWithDocs(
loss=loss,
logits=outputs.logits,
past_key_values=outputs.past_key_values,
decoder_hidden_states=outputs.decoder_hidden_states,
decoder_attentions=outputs.decoder_attentions,
encoder_last_hidden_state=outputs.encoder_last_hidden_state,
encoder_hidden_states=outputs.encoder_hidden_states,
encoder_attentions=outputs.encoder_attentions,
doc_scores=outputs.doc_scores,
)
logits = self.marginalize(outputs.logits, outputs.doc_scores) if marginalize else outputs.logits
return Seq2SeqLMOutputWithDocs(
loss=None,
logits=logits,
past_key_values=outputs.past_key_values,
decoder_hidden_states=outputs.decoder_hidden_states,
decoder_attentions=outputs.decoder_attentions,
encoder_last_hidden_state=outputs.encoder_last_hidden_state,
encoder_hidden_states=outputs.encoder_hidden_states,
encoder_attentions=outputs.encoder_attentions,
doc_scores=outputs.doc_scores,
)
def get_input_embeddings(self):
return self.model.generator.get_input_embeddings()
def get_output_embeddings(self):
return self.model.generator.get_output_embeddings()
def get_encoder(self):
return RAGEncoder(self.model)
def get_nll(self, seq_logits, doc_scores, target, reduce=False, epsilon=0.0):
target = self.shift_tokens_left(target)
def _mask_pads(ll, smooth_obj):
pad_mask = target.eq(self.config.pad_token_id)
if pad_mask.any():
ll.masked_fill_(pad_mask, 0.0)
smooth_obj.masked_fill_(pad_mask, 0.0)
return ll.squeeze(-1), smooth_obj.squeeze(-1)
rag_logprobs = self.marginalize(seq_logits, doc_scores)
target = target.unsqueeze(-1)
assert target.dim() == rag_logprobs.dim()
ll = rag_logprobs.gather(dim=-1, index=target)
smooth_obj = rag_logprobs.sum(dim=-1, keepdim=True) # total sum of all (normalised) logits
ll, smooth_obj = _mask_pads(ll, smooth_obj)
ll = ll.sum(1) # sum over tokens
smooth_obj = smooth_obj.sum(1)
nll_loss = -ll
smooth_loss = -smooth_obj
if reduce:
nll_loss = nll_loss.sum()
smooth_loss = smooth_loss.sum()
eps_i = epsilon / rag_logprobs.size(-1)
loss = (1.0 - epsilon) * nll_loss + eps_i * smooth_loss
return loss
+502
View File
@@ -0,0 +1,502 @@
# coding=utf-8
# Copyright 2020, The RAG Authors and The HuggingFace Inc. team.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""RAG Retriever model implementation."""
import os
import pickle
import time
import faiss
import numpy as np
import psutil
import torch
import torch.distributed as dist
from datasets import load_dataset
from .file_utils import cached_path, is_remote_url
from .tokenization_auto import AutoTokenizer
from .tokenization_dpr import DPRQuestionEncoderTokenizer
from .tokenization_t5 import T5Tokenizer
from .utils import logging
logger = logging.get_logger(__name__)
class Index(object):
"""
A base class for the Indices encapsulated by the :class:`~transformers.RagRetriever`.
"""
def __init__(self, *args, **kwargs):
pass
def get_doc_dicts(self, doc_ids):
"""
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)`):
A tensor of document indices.
"""
pass
def get_top_docs(self, query_vectors, n_docs):
"""
For each query in the batch, retrieves ``n_docs`` documents.
Args:
query_vectors (:obj:`np.array` 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.
"""
raise NotImplementedError
def is_initialized(self):
"""
Returns :obj:`True` if index is already initialized.
"""
raise NotImplementedError
def init_index(self):
"""
A function responsible for loading the index into memory. Should be called only once per training run of a RAG model.
E.g. if the model is trained on multiple GPUs in a distributed setup, only one of the workers will load the index.
"""
raise NotImplementedError
class LegacyIndex(Index):
"""
An index which can be deserialized from the files built using https://github.com/facebookresearch/DPR.
We use default faiss index parameters as specified in that repository.
Args:
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`
"""
INDEX_FILENAME = "hf_bert_base.hnswSQ8_correct_phi_128.c_index"
PASSAGE_FILENAME = "psgs_w100.tsv.pkl"
def __init__(self, vector_size, index_path):
self.index_id_to_db_id = []
self.index_path = index_path
self.passages = self._load_passages()
self.vector_size = vector_size
self.index = None
self._index_initialize = False
def _resolve_path(self, index_path, filename):
assert os.path.isdir(index_path) or is_remote_url(index_path), "Please specify a valid ``index_path``."
archive_file = os.path.join(index_path, filename)
try:
# Load from URL or cache if already cached
resolved_archive_file = cached_path(archive_file)
if resolved_archive_file is None:
raise EnvironmentError
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"- or '{index_path}' is the correct path to a directory containing a file named {filename}.\n\n"
)
raise EnvironmentError(msg)
if resolved_archive_file == archive_file:
logger.info("loading file {}".format(archive_file))
else:
logger.info("loading file {} from cache at {}".format(archive_file, resolved_archive_file))
return resolved_archive_file
def _load_passages(self):
passages_path = self._resolve_path(self.index_path, self.PASSAGE_FILENAME)
with open(passages_path, "rb") as passages_file:
passages = pickle.load(passages_file)
return passages
def _deserialize_index(self):
logger.info("Loading index from {}".format(self.index_path))
resolved_index_path = self._resolve_path(self.index_path, self.INDEX_FILENAME + ".index.dpr")
self.index = faiss.read_index(resolved_index_path)
resolved_meta_path = self._resolve_path(self.index_path, self.INDEX_FILENAME + ".index_meta.dpr")
with open(resolved_meta_path, "rb") as metadata_file:
self.index_id_to_db_id = pickle.load(metadata_file)
assert (
len(self.index_id_to_db_id) == self.index.ntotal
), "Deserialized index_id_to_db_id should match faiss index size"
def is_initialized(self):
return self._index_initialize
def init_index(self):
index = faiss.IndexHNSWFlat(self.vector_size + 1, 512)
index.hnsw.efSearch = 128
index.hnsw.efConstruction = 200
self.index = index
self._deserialize_index()
self._index_initialize = True
def get_doc_dicts(self, doc_ids):
doc_list = []
for doc_ids_i in doc_ids:
ids = [str(int(doc_id)) for doc_id in doc_ids_i]
docs = [self.passages[doc_id] for doc_id in ids]
doc_list.append(docs)
doc_dicts = []
for docs in doc_list:
doc_dict = {}
doc_dict["title"] = [doc[1] for doc in docs]
doc_dict["text"] = [doc[0] for doc in docs]
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))
_, 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)
class HFIndex(Index):
"""
A wrapper around an instance of :class:`~datasets.Datasets`. If ``index_path`` is set to ``None``,
we load the pre-computed index available with the :class:`~datasets.arrow_dataset.Dataset`, otherwise, we load the index from the indicated path on disk.
Args:
dataset (:obj:`str`, optional, defaults to ``wiki_dpr``):
A datatset identifier of the indexed dataset on HuggingFace AWS bucket (list all available datasets and ids with ``datasets.list_datasets()``).
dataset_split (:obj:`str`, optional, defaults to ``train``)
Which split of the ``dataset`` to load.
index_name (:obj:`str`, optional, defaults to ``train``)
The index_name of the index associated with the ``dataset``. The index loaded from ``index_path`` will be saved under this name.
index_path (:obj:`str`, optional, defaults to ``None``)
The path to the serialized faiss index on disk.
"""
def __init__(
self,
dataset,
dataset_split,
index_name,
index_path,
dummy,
):
super().__init__()
self.dataset = dataset
self.dataset_split = dataset_split
self.index_name = index_name
self.index_path = index_path
self.dummy = dummy
self.index = load_dataset(self.dataset, with_index=False, split=self.dataset_split, dummy=self.dummy)
self._index_initialize = False
def is_initialized(self):
return self._index_initialize
def init_index(self):
if self.index_path is not None:
self.index.load_faiss_index(index_name=self.index_name, file=self.index_path)
else:
self.index = load_dataset(
self.dataset,
with_embeddings=True,
with_index=True,
split=self.dataset_split,
index_name=self.index_name,
dummy=self.dummy,
)
self._index_initialize = True
def get_doc_dicts(self, doc_ids):
return [self.index[doc_ids[i].tolist()] for i in range(doc_ids.shape[0])]
def get_top_docs(self, query_vectors, n_docs=5):
_, docs = self.index.get_nearest_examples_batch("embeddings", query_vectors, 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)
class RagRetriever(object):
"""
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.
"""
def __init__(self, config):
super().__init__()
assert (
config.retriever_type == "hf_retriever" or config.retriever_type == "legacy_retriever"
), "invalid retirever type"
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)
)
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.n_docs = config.n_docs
self.batch_size = config.retrieval_batch_size
if torch.cuda.is_available():
self.batch_size *= torch.cuda.device_count()
self.config = config
def init_retrieval(self, distributed_port):
"""
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.
Args:
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``.
"""
logger.info("initializing retrieval")
# initializing a separate process group for retrievel as the default
# nccl backend doesn't support gather/scatter operations while gloo
# is too slow to replace nccl for the core gpu communication
if dist.is_initialized():
logger.info("dist initialized")
# needs to be set manually
os.environ["GLOO_SOCKET_IFNAME"] = self._infer_socket_ifname()
# avoid clash with the NCCL port
os.environ["MASTER_PORT"] = str(distributed_port + 1)
self.process_group = dist.new_group(ranks=None, backend="gloo")
# 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()
# 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)
return target_tensor
def _infer_socket_ifname(self):
addrs = psutil.net_if_addrs()
# a hacky way to deal with varying network interface names
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):
"""
Retrieves documents for specified ``query_vectors``. 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)`:
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.
"""
# 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)
# distributed training
world_size = dist.get_world_size(group=self.process_group)
# 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)
# scatter logic
n_queries = query_vectors.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))
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]])
return doc_vectors, self.retriever.get_doc_dicts(doc_ids)
+66
View File
@@ -0,0 +1,66 @@
# coding=utf-8
# Copyright 2020, The RAG Authors and The HuggingFace Inc. team.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Tokenization classes for RAG."""
from .tokenization_bart import BartTokenizer, BartTokenizerFast
VOCAB_FILES_NAMES = {
"vocab_file": "vocab.json",
"merges_file": "merges.txt",
}
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",
},
}
class RagDefaultTokenizer(BartTokenizer):
r"""
Constructs a RagDefaultTokenizer.
:class:`~transformers.RagDefaultTokenizer` 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
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
+472
View File
@@ -0,0 +1,472 @@
# coding=utf-8
# Copyright 2020, The RAG Authors and The HuggingFace Inc. team.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import copy
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
from .test_configuration_common import ConfigTester
from .test_modeling_common import ids_tensor
TOLERANCE = 1e-4
if is_torch_available() and is_datasets_available() and is_faiss_available() and is_psutil_available():
import torch
from transformers import (
BartConfig,
BartForConditionalGeneration,
BartTokenizer,
DPRConfig,
DPRQuestionEncoder,
RagConfig,
RagRetriever,
RagSequence,
RagToken,
)
def _assert_tensors_equal(a, b, atol=1e-12, prefix=""):
"""If tensors not close, or a and b arent both tensors, raise a nice Assertion error."""
if a is None and b is None:
return True
try:
if torch.allclose(a, b, atol=atol):
return True
raise
except Exception:
msg = "{} != {}".format(a, b)
if prefix:
msg = prefix + ": " + msg
raise AssertionError(msg)
def require_retrieval(test_case):
"""
Decorator marking a test that requires a set of dependencies necessary for pefrorm retrieval with
:class:`~transformers.RagRetriever`.
These tests are skipped when respective libraries are not installed.
"""
if not (is_torch_available() and is_datasets_available() and is_faiss_available() and is_psutil_available()):
test_case = unittest.skip("test requires PyTorch")(test_case)
return test_case
class RagModelTester:
def __init__(
self,
parent,
):
# Global params
self.parent = parent
self.batch_size = 13
self.seq_length = 7
# RAG params
self.n_docs = 3
self.vocab_size = 50265
self.bos_token_id = 0
self.pad_token_id = 1
self.eos_token_id = 2
self.decoder_start_token_id = 2
self.max_combined_length = 123
self.retrieval_vector_size = 768
self.retrieval_batch_size = 8
self.rag_config = RagConfig(
n_docs=self.n_docs,
vocab_size=self.vocab_size,
bos_token_id=self.bos_token_id,
pad_token_id=self.pad_token_id,
eos_token_id=self.eos_token_id,
decoder_start_token_id=self.decoder_start_token_id,
max_combined_length=self.max_combined_length,
retrieval_vector_size=self.retrieval_vector_size,
retrieval_batch_size=self.retrieval_batch_size,
)
# BART params
self.hidden_size = 16
self.num_hidden_layers = 2
self.num_attention_heads = 4
self.intermediate_size = 4
self.hidden_dropout_prob = 0.1
self.attention_probs_dropout_prob = 0.1
self.max_position_embeddings = 20
self.bart_config = BartConfig(
vocab_size=self.vocab_size,
d_model=self.hidden_size,
encoder_layers=self.num_hidden_layers,
decoder_layers=self.num_hidden_layers,
encoder_attention_heads=self.num_attention_heads,
decoder_attention_heads=self.num_attention_heads,
encoder_ffn_dim=self.intermediate_size,
decoder_ffn_dim=self.intermediate_size,
dropout=self.hidden_dropout_prob,
attention_dropout=self.attention_probs_dropout_prob,
max_position_embeddings=self.max_position_embeddings,
eos_token_id=self.eos_token_id,
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
self.dpr_vocab_size = 51
self.hidden_size = 20
self.num_hidden_layers = 3
self.num_attention_heads = 5
self.intermediate_size = 5
self.hidden_act = "gelu"
self.hidden_dropout_prob = 0.2
self.attention_probs_dropout_prob = 0.2
self.max_position_embeddings = 19
self.type_vocab_size = 17
self.initializer_range = 0.02
self.projection_dim = 0
self.dpr_config = DPRConfig(
projection_dim=self.projection_dim,
vocab_size=self.dpr_vocab_size,
hidden_size=self.hidden_size,
num_hidden_layers=self.num_hidden_layers,
num_attention_heads=self.num_attention_heads,
intermediate_size=self.intermediate_size,
hidden_act=self.hidden_act,
hidden_dropout_prob=self.hidden_dropout_prob,
attention_probs_dropout_prob=self.attention_probs_dropout_prob,
max_position_embeddings=self.max_position_embeddings,
type_vocab_size=self.type_vocab_size,
is_decoder=False,
initializer_range=self.initializer_range,
return_dict=True,
)
def prepare_inputs(self):
input_ids = ids_tensor([self.batch_size, self.seq_length], self.vocab_size).clamp(
3,
)
input_ids[:, -1] = self.eos_token_id
attention_mask = input_ids.ne(self.pad_token_id)
return input_ids, attention_mask
@require_torch
@require_retrieval
class RagModelTest(unittest.TestCase):
all_model_classes = (
(RagSequence, RagToken)
if is_torch_available() and is_datasets_available() and is_faiss_available() and is_psutil_available()
else ()
)
def setUp(self):
self.model_tester = RagModelTester(self)
self.config_tester = ConfigTester(self, config_class=RagConfig, hidden_size=37)
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()
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)
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),
)
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_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.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)
# no cache
result = model(
input_ids,
retriever=None,
decoder_input_ids=decoder_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,
)
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(
input_ids,
retriever=None,
decoder_input_ids=decoder_input_ids,
attention_mask=attention_mask,
return_loss=True,
reduce=True,
)
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, torch.Size([]))
@require_torch
@require_retrieval
class RagModelIntegrationTests(unittest.TestCase):
def get_rag_config(self):
return RagConfig(
bos_token_id=0,
decoder_start_token_id=2,
eos_token_id=2,
is_encoder_decoder=True,
pad_token_id=1,
vocab_size=50264,
title_sep=" / ",
doc_sep=" // ",
n_docs=5,
max_combined_length=300,
retriever_type="hf_retriever",
dataset="wiki_dpr",
dataset_split="train",
index_name="exact",
index_path=None,
dummy=True,
retrieval_vector_size=768,
retrieval_batch_size=8,
pretrained_question_encoder_name_or_path="facebook/dpr-question_encoder-single-nq-base",
pretrained_generator_tokenizer_name_or_path="facebook/bart-large-cnn",
pretrained_generator_name_or_path="facebook/bart-large-cnn",
)
@slow
def test_rag_sequence_inference(self):
rag_config = self.get_rag_config()
rag_retriever = RagRetriever(rag_config)
rag_tokenizer = BartTokenizer.from_pretrained("facebook/bart-large-cnn")
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
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,
)
expected_shape = torch.Size([5, 5, 50264])
self.assertEqual(output.logits.shape, expected_shape)
expected_loss = torch.tensor([38.7446])
_assert_tensors_equal(expected_loss, output.loss, atol=TOLERANCE)
expected_doc_scores = torch.tensor([[75.0286, 74.4998, 74.0804, 74.0306, 73.9504]])
_assert_tensors_equal(expected_doc_scores, output.doc_scores, atol=TOLERANCE)
@slow
def test_rag_token_inference(self):
rag_config = self.get_rag_config()
rag_retriever = RagRetriever(rag_config)
rag_tokenizer = BartTokenizer.from_pretrained("facebook/bart-large-cnn")
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
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,
)
expected_shape = torch.Size([5, 5, 50264])
self.assertEqual(output.logits.shape, expected_shape)
expected_loss = torch.tensor([38.7045])
_assert_tensors_equal(expected_loss, output.loss, atol=TOLERANCE)
expected_doc_scores = torch.tensor([[75.0286, 74.4998, 74.0804, 74.0306, 73.9504]])
_assert_tensors_equal(expected_doc_scores, output.doc_scores, atol=TOLERANCE)