LegacyIndex index download refactor

This commit is contained in:
Your Name
2020-09-14 10:22:09 -07:00
parent 17a8621661
commit 9b561de9ca
7 changed files with 98 additions and 43 deletions
+3 -3
View File
@@ -65,9 +65,9 @@ python examples/rag/eval_rag.py \
--model_type rag_sequence \ # RAG model type (rag_token or rag_sequence)
--evaluation_set path/to/output/biencoder-nq-dev.questions \ # an input dataset for evaluation
--gold_data_path path/to/output/biencoder-nq-dev.pages \ # a dataset containing ground truth answers for samples from the evaluation_set
--predictions_filename retrieval_preds.tsv \ # name of file in which predictions will be stored
--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_filename already exists, and this option is set - we regenerate the answers, otherwise we reuse the predicsion file to calculate metrics.
--recalculate # if predictions_path already exists, and this option is set - we regenerate the answers, otherwise we reuse the predicsion file to calculate metrics.
```
@@ -78,7 +78,7 @@ python examples/rag/eval_rag.py \
--model_type rag_sequence \
--evaluation_set path/to/test.source \
--gold_data_path path/to/gold_data \
--predictions_filename e2e_preds.txt \
--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
+8 -9
View File
@@ -196,10 +196,10 @@ def get_args():
"ans - a single line of the gold file contains the expected answer string",
)
parser.add_argument(
"--predictions_filename",
"--predictions_path",
type=str,
default="predictions.txt",
help="Name of the predictions file, to be stored in the checkpoints directry",
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",
@@ -268,15 +268,14 @@ def main(args):
evaluate_batch_fn = evaluate_batch_e2e if args.eval_mode == "e2e" else evaluate_batch_retrieval
for checkpoint in checkpoints:
predictions_path = os.path.join(checkpoint, args.predictions_filename)
if os.path.exists(predictions_path) and (not args.recalculate):
logger.info("Calculating metrics based on an existing predictions file: {}".format(predictions_path))
score_fn(args, predictions_path, args.gold_data_path)
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(predictions_path))
logger.info(" Predictions will be stored under {}".format(args.predictions_path))
model = model_class.from_pretrained(checkpoint, **model_kwargs)
model.to(args.device)
@@ -291,7 +290,7 @@ def main(args):
if args.model_type != "bart":
retriever.init_retrieval(distributed_port=12345)
with open(args.evaluation_set, "r") as eval_file, open(predictions_path, "w") as preds_file:
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())
@@ -305,7 +304,7 @@ def main(args):
preds_file.write("\n".join(answers))
preds_file.flush()
score_fn(args, predictions_path, args.gold_data_path)
score_fn(args, args.predictions_path, args.gold_data_path)
if __name__ == "__main__":
+6 -5
View File
@@ -89,13 +89,14 @@ class GenerativeQAModule(BaseTransformer):
# set extra_model_params for generator configs and load_model
extra_model_params = ("encoder_layerdrop", "decoder_layerdrop", "attention_dropout", "dropout")
if self.is_rag_model:
generator_config = AutoConfig.from_pretrained(
config.pretrained_generator_name_or_path, prefix=config.prefix
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, config=config, generator_config=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)
+17 -10
View File
@@ -50,7 +50,7 @@ RAG_CONFIG_DOC = r"""
retriever_type (:obj:`str`, optional, defaults to ``hf_retriever``):
A type of index encapsulated by the ``retriever``. Possible options include:
- ``hf_retriever`` - and index build for an instance of :class:`~nlp.Datasets`
- ``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()``).
@@ -59,19 +59,28 @@ RAG_CONFIG_DOC = r"""
index_name (:obj:`str`, optional, defaults to ``train``)
The index_name of the index associated with the ``dataset``.
index_path (:obj:`str`, optional, defaults to ``None``)
The path to the serialized faiss index on disk.
passages_path: (:obj:`str`, optional, defaults to ``None``):
A path to text passages compatible with the faiss index. Required if using :class:`~transformers.retrieval_rag.LegacyIndex`
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.
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.
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`):
@@ -111,12 +120,11 @@ class RagConfig(PretrainedConfig):
dataset_split="train",
index_name="embeddings",
index_path=None,
passages_path=None,
dummy=False,
pretrained_question_encoder_tokenizer_name_or_path="facebook/dpr-question_encoder-single-nq-base",
pretrained_question_encoder_name_or_path="facebook/dpr-question_encoder-single-nq-base",
pretrained_question_encoder_name_or_path=None,
pretrained_generator_tokenizer_name_or_path="facebook/bart-large",
pretrained_generator_name_or_path="facebook/bart-large",
pretrained_generator_name_or_path=None,
**kwargs
):
super().__init__(**kwargs)
@@ -141,7 +149,6 @@ class RagConfig(PretrainedConfig):
self.retrieval_vector_size = retrieval_vector_size
self.retrieval_batch_size = retrieval_batch_size
self.passages_path = passages_path
self.index_path = index_path
self.dummy = dummy
+14 -2
View File
@@ -410,18 +410,30 @@ class PreTrainedRagModel(PreTrainedModel):
**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(
config.pretrained_question_encoder_name_or_path, config=question_encoder_config
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(config.pretrained_generator_name_or_path, **generator_kwargs)
generator = AutoModelForSeq2SeqLM.from_pretrained(pretrained_generator_name_or_path, **generator_kwargs)
return cls(config, question_encoder, generator)
+49 -13
View File
@@ -25,6 +25,7 @@ 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
@@ -91,25 +92,60 @@ class LegacyIndex(Index):
vector_size (:obj:`int`):
The dimension of indexed vectors.
index_path (:obj:`str`):
The path to the serialized faiss index on disk.
passages_path: (:obj:`str`):
A path to text passages on disk, compatible with the faiss index.
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`
"""
def __init__(self, vector_size, index_path, passages_path):
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 = []
with open(passages_path, "rb") as passages_file:
self.passages = pickle.load(passages_file)
self.index_path = index_path
self.passages = self._load_passages()
self.vector_size = vector_size
self.index = None
self._index_initialize = False
def _deserialize_from(self, index_path: str):
logger.info("Loading index from {}".format(index_path))
self.index = faiss.read_index(index_path + ".index.dpr")
with open(index_path + ".index_meta.dpr", "rb") as reader:
self.index_id_to_db_id = pickle.load(reader)
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"
@@ -122,7 +158,7 @@ class LegacyIndex(Index):
index.hnsw.efSearch = 128
index.hnsw.efConstruction = 200
self.index = index
self._deserialize_from(self.index_path)
self._deserialize_index()
self._index_initialize = True
def get_doc_dicts(self, doc_ids):
@@ -228,7 +264,7 @@ class RagRetriever(object):
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, config.passages_path)
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
+1 -1
View File
@@ -33,8 +33,8 @@ if is_torch_available() and is_datasets_available() and is_faiss_available() and
from transformers import (
BartConfig,
BartTokenizer,
BartForConditionalGeneration,
BartTokenizer,
DPRConfig,
DPRQuestionEncoder,
RagConfig,