Compare commits
51
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9b561de9ca | ||
|
|
17a8621661 | ||
|
|
64796004dc | ||
|
|
df5eec9f14 | ||
|
|
1497120d89 | ||
|
|
593765088e | ||
|
|
4ad9cb1bd9 | ||
|
|
6215f3b5e0 | ||
|
|
9795dc3464 | ||
|
|
a4c7c25cd2 | ||
|
|
945d56995e | ||
|
|
e82aca09a5 | ||
|
|
845c18d9af | ||
|
|
0f3dc78c0b | ||
|
|
21e8c67bc7 | ||
|
|
972d240ae6 | ||
|
|
2b8ab2eef3 | ||
|
|
706a7c064d | ||
|
|
d37b95d39f | ||
|
|
cbf479df13 | ||
|
|
0b9b2840c3 | ||
|
|
ce3684bbcb | ||
|
|
b8a2ba8624 | ||
|
|
cc4ba034d6 | ||
|
|
84d8061dee | ||
|
|
52bbf8dfe3 | ||
|
|
930bdfaaa8 | ||
|
|
1e1a671614 | ||
|
|
edbfab74d9 | ||
|
|
3fa122af11 | ||
|
|
85929c02dd | ||
|
|
2e4dac6802 | ||
|
|
7184c2b2c2 | ||
|
|
5ed8b6460a | ||
|
|
5774e511b0 | ||
|
|
71d2aa59e0 | ||
|
|
0a6be7b00f | ||
|
|
3cb50b0c31 | ||
|
|
3bf6f0b8f1 | ||
|
|
9ea0a9e825 | ||
|
|
3f30195fd0 | ||
|
|
90bb4188e4 | ||
|
|
6c7d856b86 | ||
|
|
17007b7b66 | ||
|
|
ae3079a9c4 | ||
|
|
97c9900124 | ||
|
|
cbcfc0e948 | ||
|
|
5fd9e5d6ab | ||
|
|
20c8a9aa36 | ||
|
|
f8b1e89f8b | ||
|
|
23eaa7afcf |
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
```
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
Executable
+34
@@ -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
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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"]
|
||||
|
||||
@@ -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."
|
||||
|
||||
@@ -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
|
||||
@@ -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 "")
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user