LegacyIndex index download refactor
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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__":
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user