Compare commits
86
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9b561de9ca | ||
|
|
17a8621661 | ||
|
|
64796004dc | ||
|
|
df5eec9f14 | ||
|
|
1497120d89 | ||
|
|
593765088e | ||
|
|
4ad9cb1bd9 | ||
|
|
6215f3b5e0 | ||
|
|
9795dc3464 | ||
|
|
a4c7c25cd2 | ||
|
|
945d56995e | ||
|
|
e82aca09a5 | ||
|
|
b76cb1c3df | ||
|
|
563ffb3dc3 | ||
|
|
1ad49cde3a | ||
|
|
4753816e39 | ||
|
|
0a8c17d53c | ||
|
|
4cbd50e611 | ||
|
|
ae736163d0 | ||
|
|
e841b75dec | ||
|
|
0054a48cdd | ||
|
|
221d4c63a3 | ||
|
|
8fcbe486e1 | ||
|
|
77950c485a | ||
|
|
514486739c | ||
|
|
e9a2f772bc | ||
|
|
df4594a9da | ||
|
|
d6c08b07a0 | ||
|
|
db38f7ce29 | ||
|
|
3bd95b0faf | ||
|
|
eb2feb5d90 | ||
|
|
66a5a6fda8 | ||
|
|
9ccdb1d517 | ||
|
|
60698936fc | ||
|
|
e0c3bc8ee0 | ||
|
|
c356b9878d | ||
|
|
5afd3f6196 | ||
|
|
15a189049e | ||
|
|
7fd1febf38 | ||
|
|
d1691d90e5 | ||
|
|
63e539459d | ||
|
|
054db06b1b | ||
|
|
b482ad474a | ||
|
|
845c18d9af | ||
|
|
762cba3bda | ||
|
|
0f3dc78c0b | ||
|
|
21e8c67bc7 | ||
|
|
49e9be0639 | ||
|
|
4ee1053dcf | ||
|
|
972d240ae6 | ||
|
|
2b8ab2eef3 | ||
|
|
706a7c064d | ||
|
|
76818cc4c6 | ||
|
|
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 |
No files matched your search
@@ -77,7 +77,7 @@ jobs:
|
||||
- v0.3-torch_and_tf-{{ checksum "setup.py" }}
|
||||
- v0.3-{{ checksum "setup.py" }}
|
||||
- run: pip install --upgrade pip
|
||||
- run: pip install git+https://github.com/huggingface/nlp
|
||||
- run: pip install git+https://github.com/huggingface/datasets
|
||||
- run: pip install .[sklearn,tf-cpu,torch,testing]
|
||||
- run: pip install codecov pytest-cov
|
||||
- save_cache:
|
||||
@@ -104,7 +104,7 @@ jobs:
|
||||
- v0.3-torch-{{ checksum "setup.py" }}
|
||||
- v0.3-{{ checksum "setup.py" }}
|
||||
- run: pip install --upgrade pip
|
||||
- run: pip install git+https://github.com/huggingface/nlp
|
||||
- run: pip install git+https://github.com/huggingface/datasets
|
||||
- run: pip install .[sklearn,torch,testing]
|
||||
- save_cache:
|
||||
key: v0.3-torch-{{ checksum "setup.py" }}
|
||||
@@ -129,7 +129,7 @@ jobs:
|
||||
- v0.3-tf-{{ checksum "setup.py" }}
|
||||
- v0.3-{{ checksum "setup.py" }}
|
||||
- run: pip install --upgrade pip
|
||||
- run: pip install git+https://github.com/huggingface/nlp
|
||||
- run: pip install git+https://github.com/huggingface/datasets
|
||||
- run: pip install .[sklearn,tf-cpu,testing]
|
||||
- save_cache:
|
||||
key: v0.3-tf-{{ checksum "setup.py" }}
|
||||
|
||||
@@ -46,7 +46,7 @@ jobs:
|
||||
pip install --upgrade pip
|
||||
pip install torch!=1.6.0
|
||||
pip install .[sklearn,testing,onnxruntime]
|
||||
pip install git+https://github.com/huggingface/nlp
|
||||
pip install git+https://github.com/huggingface/datasets
|
||||
|
||||
- name: Are GPUs recognized by our DL frameworks
|
||||
run: |
|
||||
|
||||
@@ -43,7 +43,7 @@ jobs:
|
||||
pip install --upgrade pip
|
||||
pip install torch!=1.6.0
|
||||
pip install .[sklearn,testing,onnxruntime]
|
||||
pip install git+https://github.com/huggingface/nlp
|
||||
pip install git+https://github.com/huggingface/datasets
|
||||
|
||||
- name: Are GPUs recognized by our DL frameworks
|
||||
run: |
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -134,6 +134,18 @@ Follow these steps to start contributing:
|
||||
it with `pip uninstall transformers` before reinstalling it in editable
|
||||
mode with the `-e` flag.)
|
||||
|
||||
To run the full test suite, you might need the additional dependency on `datasets` which requires a separate source
|
||||
install:
|
||||
|
||||
```bash
|
||||
$ git clone https://github.com/huggingface/datasets
|
||||
$ cd datasets
|
||||
$ pip install -e .
|
||||
```
|
||||
|
||||
If you have already cloned that repo, you might need to `git pull` to get the most recent changes in the `datasets`
|
||||
library.
|
||||
|
||||
5. Develop the features on your branch.
|
||||
|
||||
As you work on the features, you should make sure that the test suite
|
||||
|
||||
@@ -134,7 +134,10 @@ conversion utilities for the following models:
|
||||
26. `Funnel Transformer <https://github.com/laiguokun/Funnel-Transformer>`_ (from CMU/Google Brain) released with the paper
|
||||
`Funnel-Transformer: Filtering out Sequential Redundancy for Efficient Language Processing
|
||||
<https://arxiv.org/abs/2006.03236>`_ by Zihang Dai, Guokun Lai, Yiming Yang, Quoc V. Le.
|
||||
27. `Other community models <https://huggingface.co/models>`_, contributed by the `community
|
||||
27. `Bert For Sequence Generation <https://tfhub.dev/s?module-type=text-generation&subtype=module,placeholder>`_ (from Google) released with the paper
|
||||
`Leveraging Pre-trained Checkpoints for Sequence Generation Tasks
|
||||
<https://arxiv.org/abs/1907.12461>`_ by Sascha Rothe, Shashi Narayan, Aliaksei Severyn.
|
||||
28. `Other community models <https://huggingface.co/models>`_, contributed by the `community
|
||||
<https://huggingface.co/users>`_.
|
||||
|
||||
.. toctree::
|
||||
@@ -221,6 +224,7 @@ conversion utilities for the following models:
|
||||
model_doc/mbart
|
||||
model_doc/funnel
|
||||
model_doc/lxmert
|
||||
model_doc/bertgeneration
|
||||
internal/modeling_utils
|
||||
internal/tokenization_utils
|
||||
internal/pipelines_utils
|
||||
@@ -1,12 +1,13 @@
|
||||
Configuration
|
||||
----------------------------------------------------
|
||||
|
||||
The base class ``PretrainedConfig`` implements the common methods for loading/saving a configuration either from a
|
||||
local file or directory, or from a pretrained model configuration provided by the library (downloaded from
|
||||
HuggingFace's AWS S3 repository).
|
||||
The base class :class:`~transformers.PretrainedConfig` implements the common methods for loading/saving a configuration
|
||||
either from a local file or directory, or from a pretrained model configuration provided by the library (downloaded
|
||||
from HuggingFace's AWS S3 repository).
|
||||
|
||||
``PretrainedConfig``
|
||||
~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
PretrainedConfig
|
||||
~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.PretrainedConfig
|
||||
:members:
|
||||
@@ -21,12 +21,25 @@ previous features. To inject custom behavior you can subclass them and override
|
||||
- **setup_wandb** -- Setups wandb (see `here <https://docs.wandb.com/huggingface>`__ for more information).
|
||||
- **create_optimizer_and_scheduler** -- Setups the optimizer and learning rate scheduler if they were not passed at
|
||||
init.
|
||||
- **compute_loss** - Computes the loss on a batch of training inputs.
|
||||
- **training_step** -- Performs a training step.
|
||||
- **prediction_step** -- Performs an evaluation/test step.
|
||||
- **run_model** (TensorFlow only) -- Basic pass through the model.
|
||||
- **evaluate** -- Runs an evaluation loop and returns metrics.
|
||||
- **predict** -- Returns predictions (with metrics if labels are available) on a test set.
|
||||
|
||||
Here is an example of how to customize :class:`~transformers.Trainer` using a custom loss function:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
from transformers import Trainer
|
||||
class MyTrainer(Trainer):
|
||||
def compute_loss(self, model, inputs):
|
||||
labels = inputs.pop("labels")
|
||||
outputs = models(**inputs)
|
||||
logits = outputs[0]
|
||||
return my_custom_loss(logits, labels)
|
||||
|
||||
|
||||
``Trainer``
|
||||
~~~~~~~~~~~
|
||||
|
||||
@@ -0,0 +1,82 @@
|
||||
BertGeneration
|
||||
----------------------------------------------------
|
||||
|
||||
Overview
|
||||
~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
The BertGeneration model is a BERT model that can be leveraged for sequence-to-sequence tasks using :class:`~transformers.EncoderDecoderModel` as proposed in `Leveraging Pre-trained Checkpoints for Sequence Generation Tasks <https://arxiv.org/abs/1907.12461>`__ by Sascha Rothe, Shashi Narayan, Aliaksei Severyn.
|
||||
|
||||
The abstract from the paper is the following:
|
||||
|
||||
*Unsupervised pre-training of large neural models has recently revolutionized Natural Language Processing. By warm-starting from the publicly released checkpoints, NLP practitioners have pushed the state-of-the-art on multiple benchmarks while saving significant amounts of compute time. So far the focus has been mainly on the Natural Language Understanding tasks. In this paper, we demonstrate the efficacy of pre-trained checkpoints for Sequence Generation. We developed a Transformer-based sequence-to-sequence model that is compatible with publicly available pre-trained BERT, GPT-2 and RoBERTa checkpoints and conducted an extensive empirical study on the utility of initializing our model, both encoder and decoder, with these checkpoints. Our models result in new state-of-the-art results on Machine Translation, Text Summarization, Sentence Splitting, and Sentence Fusion.*
|
||||
|
||||
Usage:
|
||||
|
||||
- The model can be used in combination with the :class:`~transformers.EncoderDecoderModel` to leverage two bert pretrained bert checkpoints for subsequent fine-tuning.
|
||||
|
||||
::
|
||||
|
||||
# leverage checkpoints for Bert2Bert model...
|
||||
encoder = BertGenerationEncoder.from_pretrained("bert-large-uncased", bos_token_id=101, eos_token_id=102) # use BERT's cls token as BOS token and sep token as EOS token
|
||||
decoder = BertGenerationDecoder.from_pretrained("bert-large-uncased", add_cross_attention=True, is_decoder=True, bos_token_id=101, eos_token_id=102) # add cross attention layers and use BERT's cls token as BOS token and sep token as EOS token
|
||||
bert2bert = EncoderDecoderModel(encoder=encoder, decoder=decoder)
|
||||
|
||||
# create tokenizer...
|
||||
tokenizer = BertTokenizer.from_pretrained("bert-large-uncased")
|
||||
|
||||
input_ids = tokenizer('This is a long article to summarize', add_special_tokens=False, return_tensors="pt").input_ids
|
||||
labels = tokenizer('This is a short summary', return_tensors="pt").input_ids
|
||||
|
||||
# train...
|
||||
loss = bert2bert(input_ids=input_ids, decoder_input_ids=labels, labels=labels, return_dict=True).loss
|
||||
loss.backward()
|
||||
|
||||
|
||||
- Pretrained :class:`~transformers.EncoderDecoderModel` are also directly available in the model hub, *e.g.*:
|
||||
|
||||
|
||||
::
|
||||
|
||||
# instantiate sentence fusion model
|
||||
sentence_fuser = EncoderDecoderModel.from_pretrained("google/roberta2roberta_L-24_discofuse")
|
||||
tokenizer = AutoTokenizer.from_pretrained("google/roberta2roberta_L-24_discofuse")
|
||||
|
||||
input_ids = tokenizer('This is the first sentence. This is the second sentence.', add_special_tokens=False, return_tensors="pt").input_ids
|
||||
|
||||
outputs = sentence_fuser.generate(input_ids)
|
||||
|
||||
print(tokenizer.decode(outputs[0]))
|
||||
|
||||
|
||||
Tips:
|
||||
|
||||
- :class:`~transformers.BertGenerationEncoder` and :class:`~transformers.BertGenerationDecoder` should be used in combination with :class:`~transformers.EncoderDecoder`.
|
||||
- For summarization, sentence splitting, sentence fusion and translation, no special tokens are required for the input. Therefore, no EOS token should be added to the end of the input.
|
||||
|
||||
The original code can be found `here <https://tfhub.dev/s?module-type=text-generation&subtype=module,placeholder>`__.
|
||||
|
||||
BertGenerationConfig
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.BertGenerationConfig
|
||||
:members:
|
||||
|
||||
|
||||
BertGenerationTokenizer
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.BertGenerationTokenizer
|
||||
:members:
|
||||
|
||||
BertGenerationEncoder
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.BertGenerationEncoder
|
||||
:members:
|
||||
|
||||
|
||||
BertGenerationDecoder
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.BertGenerationDecoder
|
||||
:members:
|
||||
@@ -69,6 +69,9 @@ Funnel specific outputs
|
||||
.. autoclass:: transformers.modeling_funnel.FunnelForPreTrainingOutput
|
||||
:members:
|
||||
|
||||
.. autoclass:: transformers.modeling_tf_funnel.TFFunnelForPreTrainingOutput
|
||||
:members:
|
||||
|
||||
|
||||
FunnelBaseModel
|
||||
~~~~~~~~~~~~~~~
|
||||
@@ -124,3 +127,59 @@ FunnelForQuestionAnswering
|
||||
|
||||
.. autoclass:: transformers.FunnelForQuestionAnswering
|
||||
:members:
|
||||
|
||||
|
||||
TFFunnelBaseModel
|
||||
~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.TFFunnelBaseModel
|
||||
:members:
|
||||
|
||||
|
||||
TFFunnelModel
|
||||
~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.TFFunnelModel
|
||||
:members:
|
||||
|
||||
|
||||
TFFunnelModelForPreTraining
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.TFFunnelForPreTraining
|
||||
:members:
|
||||
|
||||
|
||||
TFFunnelForMaskedLM
|
||||
~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.TFFunnelForMaskedLM
|
||||
:members:
|
||||
|
||||
|
||||
TFFunnelForSequenceClassification
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.TFFunnelForSequenceClassification
|
||||
:members:
|
||||
|
||||
|
||||
TFFunnelForMultipleChoice
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.TFFunnelForMultipleChoice
|
||||
:members:
|
||||
|
||||
|
||||
TFFunnelForTokenClassification
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.TFFunnelForTokenClassification
|
||||
:members:
|
||||
|
||||
|
||||
TFFunnelForQuestionAnswering
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.TFFunnelForQuestionAnswering
|
||||
:members:
|
||||
@@ -136,6 +136,13 @@ Then log in using the same credentials as on huggingface.co. To upload your mode
|
||||
|
||||
This will upload the folder containing the weights, tokenizer and configuration we prepared in the previous section.
|
||||
|
||||
By default you will be prompted to confirm that you want these files to be uploaded. If you are uploading multiple models and need to script that process, you can add `-y` to bypass the prompt. For example:
|
||||
|
||||
::
|
||||
|
||||
transformers-cli upload -y path/to/awesome-name-you-picked/
|
||||
|
||||
|
||||
If you want to upload a single file (a new version of your model, or the other framework checkpoint you want to add),
|
||||
just type:
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
# Long Form Question Answering
|
||||
|
||||
This folder contains the code for the Long Form Question answering [demo](http://35.226.96.115:8080/) as well as methods to train and use a fully end-to-end Long Form Question Answering system using the [🤗transformers](https://github.com/huggingface/transformers) and [🤗nlp](https://github.com/huggingface/nlp) libraries.
|
||||
This folder contains the code for the Long Form Question answering [demo](http://35.226.96.115:8080/) as well as methods to train and use a fully end-to-end Long Form Question Answering system using the [🤗transformers](https://github.com/huggingface/transformers) and [🤗datasets](https://github.com/huggingface/datasets) libraries.
|
||||
|
||||
You can use these methods to train your own system by following along the associate [notebook](https://github.com/huggingface/notebooks/blob/master/longform-qa/Long_Form_Question_Answering_with_ELI5_and_Wikipedia.ipynb) or [blog post](https://yjernite.github.io/lfqa.html).
|
||||
@@ -1,5 +1,5 @@
|
||||
import datasets
|
||||
import faiss
|
||||
import nlp
|
||||
import numpy as np
|
||||
import streamlit as st
|
||||
import torch
|
||||
@@ -45,7 +45,7 @@ def load_models():
|
||||
def load_indexes():
|
||||
if LOAD_DENSE_INDEX:
|
||||
faiss_res = faiss.StandardGpuResources()
|
||||
wiki40b_passages = nlp.load_dataset(path="wiki_snippets", name="wiki40b_en_100_0")["train"]
|
||||
wiki40b_passages = datasets.load_dataset(path="wiki_snippets", name="wiki40b_en_100_0")["train"]
|
||||
wiki40b_passage_reps = np.memmap(
|
||||
"wiki40b_passages_reps_32_l-8_h-768_b-512-512.dat",
|
||||
dtype="float32",
|
||||
@@ -63,7 +63,7 @@ def load_indexes():
|
||||
|
||||
@st.cache(allow_output_mutation=True)
|
||||
def load_train_data():
|
||||
eli5 = nlp.load_dataset("eli5", name="LFQA_reddit")
|
||||
eli5 = datasets.load_dataset("eli5", name="LFQA_reddit")
|
||||
eli5_train = eli5["train_eli5"]
|
||||
eli5_train_q_reps = np.memmap(
|
||||
"eli5_questions_reps.dat", dtype="float32", mode="r", shape=(eli5_train.num_rows, 128)
|
||||
|
||||
@@ -4,8 +4,8 @@ import os # noqa: F401
|
||||
from random import choice, randint
|
||||
from time import time
|
||||
|
||||
import datasets # noqa: F401
|
||||
import faiss # noqa: F401
|
||||
import nlp # noqa: F401
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import torch
|
||||
|
||||
@@ -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
|
||||
```
|
||||
Whitespace-only changes.
@@ -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
|
||||
@@ -12,7 +12,7 @@ faiss
|
||||
streamlit
|
||||
elasticsearch
|
||||
pandas
|
||||
nlp
|
||||
datasets
|
||||
fire
|
||||
pytest
|
||||
conllu
|
||||
@@ -16,5 +16,5 @@ python distillation.py \
|
||||
--train_batch_size=$BS --eval_batch_size=$BS \
|
||||
--tokenizer_name Helsinki-NLP/opus-mt-en-ro \
|
||||
--warmup_steps 500 --logger_name wandb \
|
||||
--fp16_opt_level O1 --task translation --normalize_hidden \
|
||||
--fp16_opt_level O1 --task translation --normalize_hidden --num_sanity_val_steps=0 \
|
||||
"$@"
|
||||
@@ -13,5 +13,5 @@ python distillation.py \
|
||||
--train_batch_size=$BS --eval_batch_size=$BS \
|
||||
--tokenizer_name $m --model_name_or_path $m \
|
||||
--warmup_steps 500 --sortish_sampler --logger_name wandb \
|
||||
--gpus 1 --fp16_opt_level=O1 --task translation \
|
||||
--gpus 1 --fp16_opt_level=O1 --task translation --num_sanity_val_steps=0 \
|
||||
"$@"
|
||||
@@ -5,25 +5,25 @@ from tqdm import tqdm
|
||||
|
||||
|
||||
def download_wmt_dataset(src_lang="ro", tgt_lang="en", dataset="wmt16", save_dir=None) -> None:
|
||||
"""Download a dataset using the nlp package and save it to the format expected by finetune.py
|
||||
"""Download a dataset using the datasets package and save it to the format expected by finetune.py
|
||||
Format of save_dir: train.source, train.target, val.source, val.target, test.source, test.target.
|
||||
|
||||
Args:
|
||||
src_lang: <str> source language
|
||||
tgt_lang: <str> target language
|
||||
dataset: <str> wmt16, wmt17, etc. wmt16 is a good start as it's small. To get the full list run `import nlp; print([d.id for d in nlp.list_datasets() if "wmt" in d.id])`
|
||||
dataset: <str> wmt16, wmt17, etc. wmt16 is a good start as it's small. To get the full list run `import datasets; print([d.id for d in datasets.list_datasets() if "wmt" in d.id])`
|
||||
save_dir: <str>, where to save the datasets, defaults to f'{dataset}-{src_lang}-{tgt_lang}'
|
||||
|
||||
Usage:
|
||||
>>> download_wmt_dataset('ro', 'en', dataset='wmt16') # saves to wmt16-ro-en
|
||||
"""
|
||||
try:
|
||||
import nlp
|
||||
import datasets
|
||||
except (ModuleNotFoundError, ImportError):
|
||||
raise ImportError("run pip install nlp")
|
||||
raise ImportError("run pip install datasets")
|
||||
pair = f"{src_lang}-{tgt_lang}"
|
||||
print(f"Converting {dataset}-{pair}")
|
||||
ds = nlp.load_dataset(dataset, pair)
|
||||
ds = datasets.load_dataset(dataset, pair)
|
||||
if save_dir is None:
|
||||
save_dir = f"{dataset}-{pair}"
|
||||
save_dir = Path(save_dir)
|
||||
|
||||
@@ -3,7 +3,6 @@ import glob
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
import warnings
|
||||
from collections import defaultdict
|
||||
from pathlib import Path
|
||||
from typing import Dict, List, Tuple
|
||||
@@ -11,7 +10,6 @@ from typing import Dict, List, Tuple
|
||||
import numpy as np
|
||||
import pytorch_lightning as pl
|
||||
import torch
|
||||
from packaging import version
|
||||
from torch.utils.data import DataLoader
|
||||
|
||||
from lightning_base import BaseTransformer, add_generic_args, generic_train
|
||||
@@ -68,6 +66,8 @@ class SummarizationModule(BaseTransformer):
|
||||
default_val_metric = "rouge2"
|
||||
|
||||
def __init__(self, hparams, **kwargs):
|
||||
if hparams.sortish_sampler and hparams.gpus > 1:
|
||||
hparams.replace_sampler_ddp = False
|
||||
super().__init__(hparams, num_labels=None, mode=self.mode, **kwargs)
|
||||
use_task_specific_params(self.model, "summarization")
|
||||
save_git_info(self.hparams.output_dir)
|
||||
@@ -114,6 +114,10 @@ class SummarizationModule(BaseTransformer):
|
||||
)
|
||||
self.eval_beams = self.model.config.num_beams if self.hparams.eval_beams is None else self.hparams.eval_beams
|
||||
assert self.eval_beams >= 1, f"got self.eval_beams={self.eval_beams}. Need an integer > 1"
|
||||
if self.hparams.eval_max_gen_length is not None:
|
||||
self.eval_max_length = self.hparams.eval_max_gen_length
|
||||
else:
|
||||
self.eval_max_length = self.model.config.max_length
|
||||
self.val_metric = self.default_val_metric if self.hparams.val_metric is None else self.hparams.val_metric
|
||||
|
||||
def freeze_embeds(self):
|
||||
@@ -209,12 +213,15 @@ class SummarizationModule(BaseTransformer):
|
||||
|
||||
def _generative_step(self, batch: dict) -> dict:
|
||||
t0 = time.time()
|
||||
|
||||
# parser.add_argument('--eval_max_gen_length', type=int, default=None, help='never generate more than n tokens')
|
||||
generated_ids = self.model.generate(
|
||||
batch["input_ids"],
|
||||
attention_mask=batch["attention_mask"],
|
||||
use_cache=True,
|
||||
decoder_start_token_id=self.decoder_start_token_id,
|
||||
num_beams=self.eval_beams,
|
||||
max_length=self.eval_max_length,
|
||||
)
|
||||
gen_time = (time.time() - t0) / batch["input_ids"].shape[0]
|
||||
preds: List[str] = self.ids_to_clean_text(generated_ids)
|
||||
@@ -248,8 +255,7 @@ class SummarizationModule(BaseTransformer):
|
||||
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)
|
||||
sampler = dataset.make_sortish_sampler(batch_size, distributed=self.hparams.gpus > 1)
|
||||
shuffle = False
|
||||
|
||||
dataloader = DataLoader(
|
||||
@@ -321,6 +327,7 @@ class SummarizationModule(BaseTransformer):
|
||||
parser.add_argument(
|
||||
"--val_metric", type=str, default=None, required=False, choices=["bleu", "rouge2", "loss", None]
|
||||
)
|
||||
parser.add_argument("--eval_max_gen_length", type=int, default=None, help="never generate more than n tokens")
|
||||
parser.add_argument("--save_top_k", type=int, default=1, required=False, help="How many checkpoints to save")
|
||||
parser.add_argument(
|
||||
"--early_stopping_patience",
|
||||
@@ -356,8 +363,6 @@ def main(args, model=None) -> SummarizationModule:
|
||||
model: SummarizationModule = SummarizationModule(args)
|
||||
else:
|
||||
model: SummarizationModule = TranslationModule(args)
|
||||
if version.parse(torch.__version__) == version.parse("1.6") and args.fp16:
|
||||
warnings.warn("FP16 only seems to work with torch 1.5+apex")
|
||||
dataset = Path(args.data_dir).name
|
||||
if (
|
||||
args.logger_name == "default"
|
||||
|
||||
@@ -36,6 +36,7 @@ def generate_summaries_or_translations(
|
||||
device: str = DEFAULT_DEVICE,
|
||||
fp16=False,
|
||||
task="summarization",
|
||||
prefix=None,
|
||||
**generate_kwargs,
|
||||
) -> Dict:
|
||||
"""Save model.generate results to <out_file>, and return how long it took."""
|
||||
@@ -51,9 +52,10 @@ def generate_summaries_or_translations(
|
||||
start_time = time.time()
|
||||
# update config with task specific params
|
||||
use_task_specific_params(model, task)
|
||||
if prefix is None:
|
||||
prefix = prefix or getattr(model.config, "prefix", "") or ""
|
||||
for examples_chunk in tqdm(list(chunks(examples, batch_size))):
|
||||
if "t5" in model_name:
|
||||
examples_chunk = [model.config.prefix + text for text in examples_chunk]
|
||||
examples_chunk = [prefix + text for text in examples_chunk]
|
||||
batch = tokenizer(examples_chunk, return_tensors="pt", truncation=True, padding="longest").to(device)
|
||||
summaries = model.generate(
|
||||
input_ids=batch.input_ids,
|
||||
@@ -78,6 +80,9 @@ def run_generate():
|
||||
parser.add_argument("--reference_path", type=str, required=False, help="like cnn_dm/test.target")
|
||||
parser.add_argument("--score_path", type=str, required=False, default="metrics.json", help="where to save metrics")
|
||||
parser.add_argument("--device", type=str, required=False, default=DEFAULT_DEVICE, help="cuda, cuda:1, cpu etc.")
|
||||
parser.add_argument(
|
||||
"--prefix", type=str, required=False, default=None, help="will be added to the begininng of src examples"
|
||||
)
|
||||
parser.add_argument("--task", type=str, default="summarization", help="used for task_specific_params + metrics")
|
||||
parser.add_argument("--bs", type=int, default=8, required=False, help="batch size")
|
||||
parser.add_argument(
|
||||
@@ -103,6 +108,7 @@ def run_generate():
|
||||
device=args.device,
|
||||
fp16=args.fp16,
|
||||
task=args.task,
|
||||
prefix=args.prefix,
|
||||
**parsed,
|
||||
)
|
||||
if args.reference_path is None:
|
||||
|
||||
@@ -160,7 +160,7 @@ def test_opus_mt_distill_script():
|
||||
metrics = load_json(model.metrics_save_path)
|
||||
first_step_stats = metrics["val"][0]
|
||||
last_step_stats = metrics["val"][-1]
|
||||
assert len(metrics["val"]) == (args.max_epochs / args.val_check_interval) + 1 # +1 accounts for val_sanity_check
|
||||
assert len(metrics["val"]) >= (args.max_epochs / args.val_check_interval) # +1 accounts for val_sanity_check
|
||||
|
||||
assert last_step_stats["val_avg_gen_time"] >= 0.01
|
||||
|
||||
|
||||
@@ -34,6 +34,7 @@ CHEAP_ARGS = {
|
||||
"supervise_forward": True,
|
||||
"normalize_hidden": True,
|
||||
"label_smoothing": 0.2,
|
||||
"eval_max_gen_length": None,
|
||||
"eval_beams": 1,
|
||||
"val_metric": "loss",
|
||||
"save_top_k": 1,
|
||||
@@ -148,9 +149,9 @@ class TestSummarizationDistiller(unittest.TestCase):
|
||||
no_teacher=True,
|
||||
freeze_encoder=True,
|
||||
gpus=2,
|
||||
sortish_sampler=False,
|
||||
sortish_sampler=True,
|
||||
)
|
||||
self._test_distiller_cli(updates)
|
||||
self._test_distiller_cli(updates, check_contents=False)
|
||||
|
||||
def test_distill_no_teacher(self):
|
||||
updates = dict(student_encoder_layers=2, student_decoder_layers=1, no_teacher=True)
|
||||
|
||||
+74
-16
@@ -1,6 +1,7 @@
|
||||
import itertools
|
||||
import json
|
||||
import linecache
|
||||
import math
|
||||
import os
|
||||
import pickle
|
||||
from logging import getLogger
|
||||
@@ -10,6 +11,7 @@ from typing import Callable, Dict, Iterable, List, Union
|
||||
import git
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from rouge_score import rouge_scorer, scoring
|
||||
from sacrebleu import corpus_bleu
|
||||
from torch import nn
|
||||
@@ -111,8 +113,11 @@ class AbstractSeq2SeqDataset(Dataset):
|
||||
def get_char_lens(data_file):
|
||||
return [len(x) for x in Path(data_file).open().readlines()]
|
||||
|
||||
def make_sortish_sampler(self, batch_size):
|
||||
return SortishSampler(self.src_lens, batch_size)
|
||||
def make_sortish_sampler(self, batch_size, distributed=False):
|
||||
if distributed:
|
||||
return DistributedSortishSampler(self, batch_size)
|
||||
else:
|
||||
return SortishSampler(self.src_lens, batch_size)
|
||||
|
||||
def __getitem__(self, item):
|
||||
raise NotImplementedError("You must implement this")
|
||||
@@ -191,24 +196,77 @@ class SortishSampler(Sampler):
|
||||
def __init__(self, data, batch_size):
|
||||
self.data, self.bs = data, batch_size
|
||||
|
||||
def key(self, i):
|
||||
return self.data[i]
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self.data)
|
||||
|
||||
def __iter__(self):
|
||||
idxs = np.random.permutation(len(self.data))
|
||||
sz = self.bs * 50
|
||||
ck_idx = [idxs[i : i + sz] for i in range(0, len(idxs), sz)]
|
||||
sort_idx = np.concatenate([sorted(s, key=self.key, reverse=True) for s in ck_idx])
|
||||
sz = self.bs
|
||||
ck_idx = [sort_idx[i : i + sz] for i in range(0, len(sort_idx), sz)]
|
||||
max_ck = np.argmax([self.key(ck[0]) for ck in ck_idx]) # find the chunk with the largest key,
|
||||
ck_idx[0], ck_idx[max_ck] = ck_idx[max_ck], ck_idx[0] # then make sure it goes first.
|
||||
sort_idx = np.concatenate(np.random.permutation(ck_idx[1:])) if len(ck_idx) > 1 else np.array([], dtype=np.int)
|
||||
sort_idx = np.concatenate((ck_idx[0], sort_idx))
|
||||
return iter(sort_idx)
|
||||
return iter(sortish_sampler_indices(self.data, self.bs))
|
||||
|
||||
|
||||
def sortish_sampler_indices(data: List, bs: int) -> np.array:
|
||||
"Go through the text data by order of src length with a bit of randomness. From fastai repo."
|
||||
|
||||
def key_fn(i):
|
||||
return data[i]
|
||||
|
||||
idxs = np.random.permutation(len(data))
|
||||
sz = bs * 50
|
||||
ck_idx = [idxs[i : i + sz] for i in range(0, len(idxs), sz)]
|
||||
sort_idx = np.concatenate([sorted(s, key=key_fn, reverse=True) for s in ck_idx])
|
||||
sz = bs
|
||||
ck_idx = [sort_idx[i : i + sz] for i in range(0, len(sort_idx), sz)]
|
||||
max_ck = np.argmax([key_fn(ck[0]) for ck in ck_idx]) # find the chunk with the largest key,
|
||||
ck_idx[0], ck_idx[max_ck] = ck_idx[max_ck], ck_idx[0] # then make sure it goes first.
|
||||
sort_idx = np.concatenate(np.random.permutation(ck_idx[1:])) if len(ck_idx) > 1 else np.array([], dtype=np.int)
|
||||
sort_idx = np.concatenate((ck_idx[0], sort_idx))
|
||||
return sort_idx
|
||||
|
||||
|
||||
class DistributedSortishSampler(Sampler):
|
||||
"""Copied from torch DistributedSampler"""
|
||||
|
||||
def __init__(self, dataset, batch_size, num_replicas=None, rank=None):
|
||||
if num_replicas is None:
|
||||
if not dist.is_available():
|
||||
raise RuntimeError("Requires distributed package to be available")
|
||||
num_replicas = dist.get_world_size()
|
||||
if rank is None:
|
||||
if not dist.is_available():
|
||||
raise RuntimeError("Requires distributed package to be available")
|
||||
rank = dist.get_rank()
|
||||
self.dataset = dataset
|
||||
self.num_replicas = num_replicas
|
||||
self.rank = rank
|
||||
self.epoch = 0
|
||||
self.num_samples = int(math.ceil(len(self.dataset) * 1.0 / self.num_replicas))
|
||||
self.total_size = self.num_samples * self.num_replicas
|
||||
self.batch_size = batch_size
|
||||
|
||||
def __iter__(self) -> Iterable:
|
||||
g = torch.Generator()
|
||||
g.manual_seed(self.epoch)
|
||||
available_indices = self.get_indices_for_rank() # indices[self.rank: self.total_size: self.num_replicas]
|
||||
|
||||
sortish_data = [self.dataset.src_lens[i] for i in available_indices]
|
||||
sortish_indices = sortish_sampler_indices(sortish_data, self.batch_size)
|
||||
indices = [available_indices[i] for i in sortish_indices]
|
||||
assert len(indices) == self.num_samples
|
||||
return iter(indices)
|
||||
|
||||
def get_indices_for_rank(self) -> np.array:
|
||||
indices = list(range(len(self.dataset)))
|
||||
# add extra samples to make it evenly divisible
|
||||
indices += indices[: (self.total_size - len(indices))]
|
||||
assert len(indices) == self.total_size
|
||||
# subsample
|
||||
available_indices = indices[self.rank : self.total_size : self.num_replicas]
|
||||
return available_indices
|
||||
|
||||
def __len__(self):
|
||||
return self.num_samples
|
||||
|
||||
def set_epoch(self, epoch):
|
||||
self.epoch = epoch
|
||||
|
||||
|
||||
logger = getLogger(__name__)
|
||||
|
||||
@@ -0,0 +1,37 @@
|
||||
---
|
||||
language:
|
||||
- en
|
||||
- de
|
||||
license: apache-2.0
|
||||
datasets:
|
||||
- wmt14
|
||||
---
|
||||
|
||||
# bert2bert_L-24_wmt_de_en EncoderDecoder model
|
||||
|
||||
The model was introduced in
|
||||
[this paper](https://arxiv.org/abs/1907.12461) by Sascha Rothe, Shashi Narayan, Aliaksei Severyn and first released in [this repository](https://tfhub.dev/google/bertseq2seq/bert24_de_en/1).
|
||||
|
||||
The model is an encoder-decoder model that was initialized on the `bert-large` checkpoints for both the encoder
|
||||
and decoder and fine-tuned on German to English translation on the WMT dataset, which is linked above.
|
||||
|
||||
Disclaimer: The model card has been written by the Hugging Face team.
|
||||
|
||||
## How to use
|
||||
|
||||
You can use this model for translation, *e.g.*
|
||||
|
||||
```python
|
||||
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained("google/bert2bert_L-24_wmt_de_en", pad_token="<pad>", eos_token="</s>", bos_token="<s>")
|
||||
model = AutoModelForSeq2SeqLM.from_pretrained("google/bert2bert_L-24_wmt_de_en")
|
||||
|
||||
sentence = "Willst du einen Kaffee trinken gehen mit mir?"
|
||||
|
||||
input_ids = tokenizer(sentence, return_tensors="pt", add_special_tokens=False).input_ids
|
||||
output_ids = model.generate(input_ids)[0]
|
||||
print(tokenizer.decode(output_ids, skip_special_tokens=True))
|
||||
# should output
|
||||
# Want to drink a kaffee go with me? .
|
||||
```
|
||||
@@ -0,0 +1,36 @@
|
||||
---
|
||||
language:
|
||||
- en
|
||||
- de
|
||||
license: apache-2.0
|
||||
datasets:
|
||||
- wmt14
|
||||
---
|
||||
|
||||
# bert2bert_L-24_wmt_en_de EncoderDecoder model
|
||||
|
||||
The model was introduced in
|
||||
[this paper](https://arxiv.org/abs/1907.12461) by Sascha Rothe, Shashi Narayan, Aliaksei Severyn and first released in [this repository](https://tfhub.dev/google/bertseq2seq/bert24_en_de/1).
|
||||
|
||||
The model is an encoder-decoder model that was initialized on the `bert-large` checkpoints for both the encoder
|
||||
and decoder and fine-tuned on English to German translation on the WMT dataset, which is linked above.
|
||||
|
||||
Disclaimer: The model card has been written by the Hugging Face team.
|
||||
|
||||
## How to use
|
||||
|
||||
You can use this model for translation, *e.g.*
|
||||
|
||||
```python
|
||||
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained("google/bert2bert_L-24_wmt_en_de", pad_token="<pad>", eos_token="</s>", bos_token="<s>")
|
||||
model = AutoModelForSeq2SeqLM.from_pretrained("google/bert2bert_L-24_wmt_en_de")
|
||||
|
||||
sentence = "Would you like to grab a coffee with me this week?"
|
||||
|
||||
input_ids = tokenizer(sentence, return_tensors="pt", add_special_tokens=False).input_ids
|
||||
output_ids = model.generate(input_ids)[0]
|
||||
print(tokenizer.decode(output_ids, skip_special_tokens=True))
|
||||
# should output
|
||||
# Möchten Sie diese Woche einen Kaffee mit mir schnappen?
|
||||
@@ -0,0 +1,56 @@
|
||||
---
|
||||
language: en
|
||||
license: apache-2.0
|
||||
datasets:
|
||||
- xsum
|
||||
---
|
||||
|
||||
# Roberta2Roberta_L-24_bbc EncoderDecoder model
|
||||
|
||||
The model was introduced in
|
||||
[this paper](https://arxiv.org/abs/1907.12461) by Sascha Rothe, Shashi Narayan, Aliaksei Severyn and first released in [this repository](https://tfhub.dev/google/bertseq2seq/roberta24_bbc/1).
|
||||
|
||||
The model is an encoder-decoder model that was initialized on the `roberta-large` checkpoints for both the encoder
|
||||
and decoder and fine-tuned on extreme summarization on the BBC XSum dataset, which is linked above.
|
||||
|
||||
Disclaimer: The model card has been written by the Hugging Face team.
|
||||
|
||||
## How to use
|
||||
|
||||
You can use this model for extreme summarization, *e.g.*
|
||||
|
||||
```python
|
||||
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained("google/roberta2roberta_L-24_bbc")
|
||||
model = AutoModelForSeq2SeqLM.from_pretrained("google/roberta2roberta_L-24_bbc")
|
||||
|
||||
article = """The problem is affecting people using the older
|
||||
versions of the PlayStation 3, called the "Fat"
|
||||
model.The problem isn't affecting the newer PS3
|
||||
Slim systems that have been on sale since
|
||||
September last year.Sony have also said they are
|
||||
aiming to have the problem fixed shortly but is
|
||||
advising some users to avoid using their console
|
||||
for the time being."We hope to resolve this
|
||||
problem within the next 24 hours," a statement
|
||||
reads. "In the meantime, if you have a model other
|
||||
than the new slim PS3, we advise that you do not
|
||||
use your PS3 system, as doing so may result in
|
||||
errors in some functionality, such as recording
|
||||
obtained trophies, and not being able to restore
|
||||
certain data."We believe we have identified that
|
||||
this problem is being caused by a bug in the clock
|
||||
functionality incorporated in the system."The
|
||||
PlayStation Network is used by millions of people
|
||||
around the world.It allows users to play their
|
||||
friends at games like Fifa over the internet and
|
||||
also do things like download software or visit
|
||||
online stores."""
|
||||
|
||||
input_ids = tokenizer(article, return_tensors="pt").input_ids
|
||||
output_ids = model.generate(input_ids)[0]
|
||||
print(tokenizer.decode(output_ids, skip_special_tokens=True))
|
||||
# should output
|
||||
# Some Sony PlayStation gamers are being advised to stay away from the network because of a problem with the PlayStation 3 network.
|
||||
```
|
||||
@@ -0,0 +1,73 @@
|
||||
---
|
||||
language: en
|
||||
license: apache-2.0
|
||||
datasets:
|
||||
- cnn_dailymail
|
||||
---
|
||||
|
||||
# Roberta2Roberta_L-24_cnn_daily_mail EncoderDecoder model
|
||||
|
||||
The model was introduced in
|
||||
[this paper](https://arxiv.org/abs/1907.12461) by Sascha Rothe, Shashi Narayan, Aliaksei Severyn and first released in [this repository](https://tfhub.dev/google/bertseq2seq/roberta24_cnndm/1).
|
||||
|
||||
The model is an encoder-decoder model that was initialized on the `roberta-large` checkpoints for both the encoder
|
||||
and decoder and fine-tuned on summarization on the CNN / Dailymail dataset, which is linked above.
|
||||
|
||||
Disclaimer: The model card has been written by the Hugging Face team.
|
||||
|
||||
## How to use
|
||||
|
||||
You can use this model for summarization, *e.g.*
|
||||
|
||||
```python
|
||||
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained("google/roberta2roberta_L-24_cnn_daily_mail")
|
||||
model = AutoModelForSeq2SeqLM.from_pretrained("google/roberta2roberta_L-24_cnn_daily_mail")
|
||||
|
||||
article = """ (The Hollywood Reporter)"The Rocky Horror Picture
|
||||
Show" is the latest musical getting the small-
|
||||
screen treatment. Fox is developing a two-hour
|
||||
remake of the 1975 cult classic to be directed,
|
||||
executive-produced and choreographed by Kenneth
|
||||
Ortega ("High School Musical"). The project,
|
||||
tentatively titled "The Rocky Horror Picture Show
|
||||
Event," is casting-contingent. The special will be
|
||||
filmed in advance and not air live, but few
|
||||
details beyond that are known. In addition to
|
||||
Ortega, Gail Berman and Lou Adler, who produced
|
||||
the original film, are also attached as executive
|
||||
producers. The special will be produced by Fox 21
|
||||
Television Studios, and Berman's The Jackal Group.
|
||||
The special is timed to celebrate the 40th
|
||||
anniversary of the film, which has grossed more
|
||||
than $112 million and still plays in theaters
|
||||
across the country. TV premiere dates: The
|
||||
complete guide . This isn't the first stab at
|
||||
adapting "The Rocky Horror Picture Show." In 2002,
|
||||
Fox unveiled plans for an adaptation timed to the
|
||||
30th anniversary that never came to fruition. The
|
||||
faces of pilot season 2015 . Fox's "Glee" covered
|
||||
several of the show's most popular songs for a
|
||||
Season 2 episode and even released a special "The
|
||||
Rocky Horror Glee Show" EP. There is no plan yet
|
||||
for when the adaptation will air. Fox also has a
|
||||
live musical production of "Grease", starring
|
||||
Julianne Hough and Vanessa Hudgens, scheduled to
|
||||
air on Jan. 31, 2016. Broadcast TV scorecard .
|
||||
Following in the footsteps of "The Sound of Music"
|
||||
and "Peter Pan," NBC recently announced plans to
|
||||
air a live version of The Wiz later this year.
|
||||
Ortega's credits include "Gilmore Girls," "This Is
|
||||
It" and "Hocus Pocus." He is repped by Paradigm
|
||||
and Hanson, Jacobson. ©2015 The Hollywood
|
||||
Reporter. All rights reserved."""
|
||||
|
||||
input_ids = tokenizer(article, return_tensors="pt").input_ids
|
||||
output_ids = model.generate(input_ids)[0]
|
||||
print(tokenizer.decode(output_ids, skip_special_tokens=True))
|
||||
# should output
|
||||
# Fox is developing a two-hour remake of the 1975 cult classic. The special will be directed, executive-produced and choreographed by Kenneth Ortega.
|
||||
# The special is timed to celebrate the 40th anniversary of the film, which has grossed more than $112 million.
|
||||
|
||||
```
|
||||
@@ -0,0 +1,35 @@
|
||||
---
|
||||
language: en
|
||||
license: apache-2.0
|
||||
datasets:
|
||||
- discofuse
|
||||
---
|
||||
|
||||
# Roberta2Roberta_L-24_discofuse EncoderDecoder model
|
||||
|
||||
The model was introduced in
|
||||
[this paper](https://arxiv.org/abs/1907.12461) by Sascha Rothe, Shashi Narayan, Aliaksei Severyn and first released in [this repository](https://tfhub.dev/google/bertseq2seq/roberta24_discofuse/1).
|
||||
|
||||
The model is an encoder-decoder model that was initialized on the `roberta-large` checkpoints for both the encoder
|
||||
and decoder and fine-tuned on sentencefusion on the discofuse dataset, which is linked above.
|
||||
|
||||
Disclaimer: The model card has been written by the Hugging Face team.
|
||||
|
||||
## How to use
|
||||
|
||||
You can use this model for sentence fusion, *e.g.*
|
||||
|
||||
```python
|
||||
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained("google/roberta2roberta_L-24_discofuse")
|
||||
model = AutoModelForSeq2SeqLM.from_pretrained("google/roberta2roberta_L-24_discofuse")
|
||||
|
||||
discofuse = """As a run-blocker, Zeitler moves relatively well. Zeitler often struggles at the point of contact in space."""
|
||||
|
||||
input_ids = tokenizer(discofuse, return_tensors="pt").input_ids
|
||||
output_ids = model.generate(input_ids)[0]
|
||||
print(tokenizer.decode(output_ids, skip_special_tokens=True))
|
||||
# should output
|
||||
# As a run-blocker, Zeitler moves relatively well. However, Zeitler often struggles at the point of contact in space.
|
||||
```
|
||||
@@ -0,0 +1,37 @@
|
||||
---
|
||||
language: en
|
||||
license: apache-2.0
|
||||
datasets:
|
||||
- gigaword
|
||||
---
|
||||
|
||||
# Roberta2Roberta_L-24_gigaword EncoderDecoder model
|
||||
|
||||
The model was introduced in
|
||||
[this paper](https://arxiv.org/abs/1907.12461) by Sascha Rothe, Shashi Narayan, Aliaksei Severyn and first released in [this repository](https://tfhub.dev/google/bertseq2seq/roberta24_gigaword/1).
|
||||
|
||||
The model is an encoder-decoder model that was initialized on the `roberta-large` checkpoints for both the encoder
|
||||
and decoder and fine-tuned on headline generation using the Gigaword dataset, which is linked above.
|
||||
|
||||
Disclaimer: The model card has been written by the Hugging Face team.
|
||||
|
||||
## How to use
|
||||
|
||||
You can use this model for extreme summarization, *e.g.*
|
||||
|
||||
```python
|
||||
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained("google/roberta2roberta_L-24_gigaword")
|
||||
model = AutoModelForSeq2SeqLM.from_pretrained("google/roberta2roberta_L-24_gigaword")
|
||||
|
||||
article = """australian shares closed down #.# percent monday
|
||||
following a weak lead from the united states and
|
||||
lower commodity prices , dealers said ."""
|
||||
|
||||
input_ids = tokenizer(article, return_tensors="pt").input_ids
|
||||
output_ids = model.generate(input_ids)[0]
|
||||
print(tokenizer.decode(output_ids, skip_special_tokens=True))
|
||||
# should output
|
||||
# australian shares close down #.# percent.
|
||||
```
|
||||
@@ -0,0 +1,34 @@
|
||||
---
|
||||
language: en
|
||||
license: apache-2.0
|
||||
---
|
||||
|
||||
# Roberta2Roberta_L-24_wikisplit EncoderDecoder model
|
||||
|
||||
The model was introduced in
|
||||
[this paper](https://arxiv.org/abs/1907.12461) by Sascha Rothe, Shashi Narayan, Aliaksei Severyn and first released in [this repository](https://tfhub.dev/google/bertseq2seq/roberta24_cnndm/1).
|
||||
|
||||
The model is an encoder-decoder model that was initialized on the `roberta-large` checkpoints for both the encoder
|
||||
and decoder and fine-tuned on sentence splitting on the [WikiSplit](https://github.com/google-research-datasets/wiki-split) dataset.
|
||||
|
||||
Disclaimer: The model card has been written by the Hugging Face team.
|
||||
|
||||
## How to use
|
||||
|
||||
You can use this model for sentence splitting, *e.g.*
|
||||
|
||||
```python
|
||||
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained("google/roberta2roberta_L-24_wikisplit")
|
||||
model = AutoModelForSeq2SeqLM.from_pretrained("google/roberta2roberta_L-24_wikisplit")
|
||||
|
||||
long_sentence = """Due to the hurricane, Lobsterfest has been canceled, making Bob very happy about it and he decides to open Bob 's Burgers for customers who were planning on going to Lobsterfest."""
|
||||
|
||||
input_ids = tokenizer(long_sentence, return_tensors="pt").input_ids
|
||||
output_ids = model.generate(input_ids)[0]
|
||||
print(tokenizer.decode(output_ids, skip_special_tokens=True))
|
||||
# should output
|
||||
# Due Due hurricane, Lobsterfest has been canceled, making Bob very happy about it. He decides to open B
|
||||
# ob's Burgers for customers who were planning on going to Lobsterfest.com.
|
||||
```
|
||||
@@ -0,0 +1,30 @@
|
||||
# LayoutLM
|
||||
|
||||
## Model description
|
||||
|
||||
LayoutLM is a simple but effective pre-training method of text and layout for document image understanding and information extraction tasks, such as form understanding and receipt understanding. LayoutLM archives the SOTA results on multiple datasets. For more details, please refer to our paper:
|
||||
|
||||
[LayoutLM: Pre-training of Text and Layout for Document Image Understanding](https://arxiv.org/abs/1912.13318)
|
||||
Yiheng Xu, Minghao Li, Lei Cui, Shaohan Huang, Furu Wei, Ming Zhou, [KDD 2020](https://www.kdd.org/kdd2020/accepted-papers)
|
||||
|
||||
## Training data
|
||||
|
||||
We pre-train LayoutLM on IIT-CDIP Test Collection 1.0\* dataset with two settings.
|
||||
|
||||
* LayoutLM-Base, Uncased (11M documents, 2 epochs): 12-layer, 768-hidden, 12-heads, 113M parameters **(This Model)**
|
||||
* LayoutLM-Large, Uncased (11M documents, 2 epochs): 24-layer, 1024-hidden, 16-heads, 343M parameters
|
||||
|
||||
## Citation
|
||||
|
||||
If you find LayoutLM useful in your research, please cite the following paper:
|
||||
|
||||
``` latex
|
||||
@misc{xu2019layoutlm,
|
||||
title={LayoutLM: Pre-training of Text and Layout for Document Image Understanding},
|
||||
author={Yiheng Xu and Minghao Li and Lei Cui and Shaohan Huang and Furu Wei and Ming Zhou},
|
||||
year={2019},
|
||||
eprint={1912.13318},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.CL}
|
||||
}
|
||||
```
|
||||
@@ -0,0 +1,30 @@
|
||||
# LayoutLM
|
||||
|
||||
## Model description
|
||||
|
||||
LayoutLM is a simple but effective pre-training method of text and layout for document image understanding and information extraction tasks, such as form understanding and receipt understanding. LayoutLM archives the SOTA results on multiple datasets. For more details, please refer to our paper:
|
||||
|
||||
[LayoutLM: Pre-training of Text and Layout for Document Image Understanding](https://arxiv.org/abs/1912.13318)
|
||||
Yiheng Xu, Minghao Li, Lei Cui, Shaohan Huang, Furu Wei, Ming Zhou, [KDD 2020](https://www.kdd.org/kdd2020/accepted-papers)
|
||||
|
||||
## Training data
|
||||
|
||||
We pre-train LayoutLM on IIT-CDIP Test Collection 1.0\* dataset with two settings.
|
||||
|
||||
* LayoutLM-Base, Uncased (11M documents, 2 epochs): 12-layer, 768-hidden, 12-heads, 113M parameters
|
||||
* LayoutLM-Large, Uncased (11M documents, 2 epochs): 24-layer, 1024-hidden, 16-heads, 343M parameters **(This Model)**
|
||||
|
||||
## Citation
|
||||
|
||||
If you find LayoutLM useful in your research, please cite the following paper:
|
||||
|
||||
``` latex
|
||||
@misc{xu2019layoutlm,
|
||||
title={LayoutLM: Pre-training of Text and Layout for Document Image Understanding},
|
||||
author={Yiheng Xu and Minghao Li and Lei Cui and Shaohan Huang and Furu Wei and Ming Zhou},
|
||||
year={2019},
|
||||
eprint={1912.13318},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.CL}
|
||||
}
|
||||
```
|
||||
@@ -0,0 +1,101 @@
|
||||
---
|
||||
language: bn
|
||||
tags:
|
||||
- bert
|
||||
- bengali
|
||||
- bengali-lm
|
||||
- bangla
|
||||
license: MIT
|
||||
datasets:
|
||||
- common_crawl
|
||||
- wikipedia
|
||||
- oscar
|
||||
---
|
||||
|
||||
|
||||
# Bangla BERT Base
|
||||
A long way passed. Here is our **Bangla-Bert**! It is now available in huggingface model hub.
|
||||
|
||||
[Bangla-Bert-Base](https://github.com/sagorbrur/bangla-bert) is a pretrained language model of Bengali language using mask language modeling described in [BERT](https://arxiv.org/abs/1810.04805) and it's github [repository](https://github.com/google-research/bert)
|
||||
|
||||
|
||||
|
||||
## Pretrain Corpus Details
|
||||
Corpus was downloaded from two main sources:
|
||||
|
||||
* Bengali commoncrawl copurs downloaded from [OSCAR](https://oscar-corpus.com/)
|
||||
* [Bengali Wikipedia Dump Dataset](https://dumps.wikimedia.org/bnwiki/latest/)
|
||||
|
||||
After downloading these corpus, we preprocessed it as a Bert format. which is one sentence per line and an extra newline for new documents.
|
||||
|
||||
```
|
||||
sentence 1
|
||||
sentence 2
|
||||
|
||||
sentence 1
|
||||
sentence 2
|
||||
|
||||
```
|
||||
|
||||
## Building Vocab
|
||||
We used [BNLP](https://github.com/sagorbrur/bnlp) package for training bengali sentencepiece model with vocab size 102025. We preprocess the output vocab file as Bert format.
|
||||
Our final vocab file availabe at [https://github.com/sagorbrur/bangla-bert](https://github.com/sagorbrur/bangla-bert) and also at [huggingface](https://huggingface.co/sagorsarker/bangla-bert-base) model hub.
|
||||
|
||||
## Training Details
|
||||
* Bangla-Bert was trained with code provided in Google BERT's github repository (https://github.com/google-research/bert)
|
||||
* Currently released model follows bert-base-uncased model architecture (12-layer, 768-hidden, 12-heads, 110M parameters)
|
||||
* Total Training Steps: 1 Million
|
||||
* The model was trained on a single Google Cloud TPU
|
||||
|
||||
## Evaluation Results
|
||||
|
||||
After training 1 millions steps here is the evaluation resutls.
|
||||
|
||||
```
|
||||
global_step = 1000000
|
||||
loss = 2.2406516
|
||||
masked_lm_accuracy = 0.60641736
|
||||
masked_lm_loss = 2.201459
|
||||
next_sentence_accuracy = 0.98625
|
||||
next_sentence_loss = 0.040997364
|
||||
perplexity = numpy.exp(2.2406516) = 9.393331287442784
|
||||
Loss for final step: 2.426227
|
||||
|
||||
|
||||
```
|
||||
|
||||
**NB: If you use this model for any nlp task please share evaluation results with us. We will add it here.**
|
||||
|
||||
|
||||
## How to Use
|
||||
You can use this model directly with a pipeline for masked language modeling:
|
||||
|
||||
```py
|
||||
from transformers import BertForMaskedLM, BertTokenizer, pipeline
|
||||
|
||||
model = BertForMaskedLM.from_pretrained("bangla-bert-base")
|
||||
tokenizer = BertTokenizer.from_pretrained("bangla-bert-base")
|
||||
nlp = pipeline('fill-mask', model=model, tokenizer=tokenizer)
|
||||
for pred in nlp(f"আমি বাংলায় {nlp.tokenizer.mask_token} গাই।"):
|
||||
print(pred)
|
||||
|
||||
# {'sequence': '[CLS] আমি বাংলায গান গাই । [SEP]', 'score': 0.13404667377471924, 'token': 2552, 'token_str': 'গান'}
|
||||
|
||||
```
|
||||
|
||||
|
||||
## Author
|
||||
[Sagor Sarker](https://github.com/sagorbrur)
|
||||
|
||||
## Acknowledgements
|
||||
|
||||
* Thanks to Google [TensorFlow Research Cloud (TFRC)](https://www.tensorflow.org/tfrc) for providing the free TPU credits - thank you!
|
||||
* Thank to all the people around, who always helping us to build something for Bengali.
|
||||
|
||||
## Reference
|
||||
* https://github.com/google-research/bert
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -3,7 +3,7 @@ language:
|
||||
- hi
|
||||
- en
|
||||
datasets:
|
||||
- LinCE
|
||||
- lince
|
||||
license: "MIT"
|
||||
tags:
|
||||
- codeswitching
|
||||
|
||||
@@ -3,7 +3,7 @@ language:
|
||||
- hi
|
||||
- en
|
||||
datasets:
|
||||
- LinCE
|
||||
- lince
|
||||
license: "MIT"
|
||||
tags:
|
||||
- codeswitching
|
||||
|
||||
@@ -3,7 +3,7 @@ language:
|
||||
- hi
|
||||
- en
|
||||
datasets:
|
||||
- LinCE
|
||||
- lince
|
||||
license: "MIT"
|
||||
tags:
|
||||
- codeswitching
|
||||
|
||||
@@ -3,7 +3,7 @@ language:
|
||||
- ne
|
||||
- en
|
||||
datasets:
|
||||
- LinCE
|
||||
- lince
|
||||
license: "MIT"
|
||||
tags:
|
||||
- codeswitching
|
||||
|
||||
@@ -3,7 +3,7 @@ language:
|
||||
- es
|
||||
- en
|
||||
datasets:
|
||||
- LinCE
|
||||
- lince
|
||||
license: "MIT"
|
||||
tags:
|
||||
- codeswitching
|
||||
|
||||
@@ -3,7 +3,7 @@ language:
|
||||
- es
|
||||
- en
|
||||
datasets:
|
||||
- LinCE
|
||||
- lince
|
||||
license: "MIT"
|
||||
tags:
|
||||
- codeswitching
|
||||
|
||||
@@ -3,7 +3,7 @@ language:
|
||||
- es
|
||||
- en
|
||||
datasets:
|
||||
- LinCE
|
||||
- lince
|
||||
license: "MIT"
|
||||
tags:
|
||||
- codeswitching
|
||||
|
||||
@@ -3,7 +3,7 @@ language:
|
||||
- es
|
||||
- en
|
||||
datasets:
|
||||
- LinCE
|
||||
- lince
|
||||
license: "MIT"
|
||||
tags:
|
||||
- codeswitching
|
||||
|
||||
@@ -7,6 +7,7 @@ known_first_party = transformers
|
||||
known_third_party =
|
||||
absl
|
||||
conllu
|
||||
datasets
|
||||
elasticsearch
|
||||
fairseq
|
||||
faiss
|
||||
@@ -16,7 +17,6 @@ known_third_party =
|
||||
git
|
||||
h5py
|
||||
matplotlib
|
||||
nlp
|
||||
nltk
|
||||
numpy
|
||||
packaging
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -22,6 +22,7 @@ from .configuration_albert import ALBERT_PRETRAINED_CONFIG_ARCHIVE_MAP, AlbertCo
|
||||
from .configuration_auto import ALL_PRETRAINED_CONFIG_ARCHIVE_MAP, CONFIG_MAPPING, AutoConfig
|
||||
from .configuration_bart import BartConfig
|
||||
from .configuration_bert import BERT_PRETRAINED_CONFIG_ARCHIVE_MAP, BertConfig
|
||||
from .configuration_bert_generation import BertGenerationConfig
|
||||
from .configuration_camembert import CAMEMBERT_PRETRAINED_CONFIG_ARCHIVE_MAP, CamembertConfig
|
||||
from .configuration_ctrl import CTRL_PRETRAINED_CONFIG_ARCHIVE_MAP, CTRLConfig
|
||||
from .configuration_distilbert import DISTILBERT_PRETRAINED_CONFIG_ARCHIVE_MAP, DistilBertConfig
|
||||
@@ -82,7 +83,8 @@ from .file_utils import (
|
||||
add_start_docstrings,
|
||||
cached_path,
|
||||
is_apex_available,
|
||||
is_nlp_available,
|
||||
is_datasets_available,
|
||||
is_faiss_available,
|
||||
is_psutil_available,
|
||||
is_py3nvml_available,
|
||||
is_tf_available,
|
||||
@@ -142,6 +144,7 @@ from .tokenization_albert import AlbertTokenizer
|
||||
from .tokenization_auto import TOKENIZER_MAPPING, AutoTokenizer
|
||||
from .tokenization_bart import BartTokenizer, BartTokenizerFast
|
||||
from .tokenization_bert import BasicTokenizer, BertTokenizer, BertTokenizerFast, WordpieceTokenizer
|
||||
from .tokenization_bert_generation import BertGenerationTokenizer
|
||||
from .tokenization_bert_japanese import BertJapaneseTokenizer, CharacterTokenizer, MecabTokenizer
|
||||
from .tokenization_camembert import CamembertTokenizer
|
||||
from .tokenization_ctrl import CTRLTokenizer
|
||||
@@ -207,6 +210,7 @@ if is_torch_available():
|
||||
DataCollatorForLanguageModeling,
|
||||
DataCollatorForNextSentencePrediction,
|
||||
DataCollatorForPermutationLanguageModeling,
|
||||
DataCollatorForSOP,
|
||||
DataCollatorWithPadding,
|
||||
default_data_collator,
|
||||
)
|
||||
@@ -214,6 +218,7 @@ if is_torch_available():
|
||||
GlueDataset,
|
||||
GlueDataTrainingArguments,
|
||||
LineByLineTextDataset,
|
||||
LineByLineWithSOPTextDataset,
|
||||
SquadDataset,
|
||||
SquadDataTrainingArguments,
|
||||
TextDataset,
|
||||
@@ -277,6 +282,11 @@ if is_torch_available():
|
||||
BertPreTrainedModel,
|
||||
load_tf_weights_in_bert,
|
||||
)
|
||||
from .modeling_bert_generation import (
|
||||
BertGenerationDecoder,
|
||||
BertGenerationEncoder,
|
||||
load_tf_weights_in_bert_generation,
|
||||
)
|
||||
from .modeling_camembert import (
|
||||
CAMEMBERT_PRETRAINED_MODEL_ARCHIVE_LIST,
|
||||
CamembertForCausalLM,
|
||||
@@ -583,6 +593,17 @@ if is_tf_available():
|
||||
TFFlaubertModel,
|
||||
TFFlaubertWithLMHeadModel,
|
||||
)
|
||||
from .modeling_tf_funnel import (
|
||||
TF_FUNNEL_PRETRAINED_MODEL_ARCHIVE_LIST,
|
||||
TFFunnelBaseModel,
|
||||
TFFunnelForMaskedLM,
|
||||
TFFunnelForMultipleChoice,
|
||||
TFFunnelForPreTraining,
|
||||
TFFunnelForQuestionAnswering,
|
||||
TFFunnelForSequenceClassification,
|
||||
TFFunnelForTokenClassification,
|
||||
TFFunnelModel,
|
||||
)
|
||||
from .modeling_tf_gpt2 import (
|
||||
TF_GPT2_PRETRAINED_MODEL_ARCHIVE_LIST,
|
||||
TFGPT2DoubleHeadsModel,
|
||||
@@ -692,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."
|
||||
|
||||
@@ -40,6 +40,7 @@ class UserCommands(BaseTransformersCLICommand):
|
||||
upload_parser.add_argument(
|
||||
"--filename", type=str, default=None, help="Optional: override individual object filename on S3."
|
||||
)
|
||||
upload_parser.add_argument("-y", "--yes", action="store_true", help="Optional: answer Yes to the prompt")
|
||||
upload_parser.set_defaults(func=lambda args: UploadCommand(args))
|
||||
|
||||
|
||||
@@ -221,10 +222,11 @@ class UploadCommand(BaseUserCommand):
|
||||
)
|
||||
)
|
||||
|
||||
choice = input("Proceed? [Y/n] ").lower()
|
||||
if not (choice == "" or choice == "y" or choice == "yes"):
|
||||
print("Abort")
|
||||
exit()
|
||||
if not self.args.yes:
|
||||
choice = input("Proceed? [Y/n] ").lower()
|
||||
if not (choice == "" or choice == "y" or choice == "yes"):
|
||||
print("Abort")
|
||||
exit()
|
||||
print(ANSI.bold("Uploading... This might take a while if files are large"))
|
||||
for filepath, filename in files:
|
||||
try:
|
||||
|
||||
@@ -14,12 +14,13 @@
|
||||
# limitations under the License.
|
||||
""" Auto Config class. """
|
||||
|
||||
|
||||
import re
|
||||
from collections import OrderedDict
|
||||
|
||||
from .configuration_albert import ALBERT_PRETRAINED_CONFIG_ARCHIVE_MAP, AlbertConfig
|
||||
from .configuration_bart import BART_PRETRAINED_CONFIG_ARCHIVE_MAP, BartConfig
|
||||
from .configuration_bert import BERT_PRETRAINED_CONFIG_ARCHIVE_MAP, BertConfig
|
||||
from .configuration_bert_generation import BertGenerationConfig
|
||||
from .configuration_camembert import CAMEMBERT_PRETRAINED_CONFIG_ARCHIVE_MAP, CamembertConfig
|
||||
from .configuration_ctrl import CTRL_PRETRAINED_CONFIG_ARCHIVE_MAP, CTRLConfig
|
||||
from .configuration_distilbert import DISTILBERT_PRETRAINED_CONFIG_ARCHIVE_MAP, DistilBertConfig
|
||||
@@ -77,118 +78,126 @@ ALL_PRETRAINED_CONFIG_ARCHIVE_MAP = dict(
|
||||
|
||||
CONFIG_MAPPING = OrderedDict(
|
||||
[
|
||||
(
|
||||
"retribert",
|
||||
RetriBertConfig,
|
||||
),
|
||||
(
|
||||
"t5",
|
||||
T5Config,
|
||||
),
|
||||
(
|
||||
"mobilebert",
|
||||
MobileBertConfig,
|
||||
),
|
||||
(
|
||||
"distilbert",
|
||||
DistilBertConfig,
|
||||
),
|
||||
(
|
||||
"albert",
|
||||
AlbertConfig,
|
||||
),
|
||||
(
|
||||
"camembert",
|
||||
CamembertConfig,
|
||||
),
|
||||
(
|
||||
"xlm-roberta",
|
||||
XLMRobertaConfig,
|
||||
),
|
||||
("retribert", RetriBertConfig),
|
||||
("t5", T5Config),
|
||||
("mobilebert", MobileBertConfig),
|
||||
("distilbert", DistilBertConfig),
|
||||
("albert", AlbertConfig),
|
||||
("bert-generation", BertGenerationConfig),
|
||||
("camembert", CamembertConfig),
|
||||
("xlm-roberta", XLMRobertaConfig),
|
||||
("pegasus", PegasusConfig),
|
||||
(
|
||||
"marian",
|
||||
MarianConfig,
|
||||
),
|
||||
(
|
||||
"mbart",
|
||||
MBartConfig,
|
||||
),
|
||||
(
|
||||
"bart",
|
||||
BartConfig,
|
||||
),
|
||||
(
|
||||
"reformer",
|
||||
ReformerConfig,
|
||||
),
|
||||
(
|
||||
"longformer",
|
||||
LongformerConfig,
|
||||
),
|
||||
(
|
||||
"roberta",
|
||||
RobertaConfig,
|
||||
),
|
||||
(
|
||||
"flaubert",
|
||||
FlaubertConfig,
|
||||
),
|
||||
(
|
||||
"bert",
|
||||
BertConfig,
|
||||
),
|
||||
(
|
||||
"openai-gpt",
|
||||
OpenAIGPTConfig,
|
||||
),
|
||||
(
|
||||
"gpt2",
|
||||
GPT2Config,
|
||||
),
|
||||
(
|
||||
"transfo-xl",
|
||||
TransfoXLConfig,
|
||||
),
|
||||
(
|
||||
"xlnet",
|
||||
XLNetConfig,
|
||||
),
|
||||
(
|
||||
"xlm",
|
||||
XLMConfig,
|
||||
),
|
||||
(
|
||||
"ctrl",
|
||||
CTRLConfig,
|
||||
),
|
||||
(
|
||||
"electra",
|
||||
ElectraConfig,
|
||||
),
|
||||
(
|
||||
"encoder-decoder",
|
||||
EncoderDecoderConfig,
|
||||
),
|
||||
(
|
||||
"funnel",
|
||||
FunnelConfig,
|
||||
),
|
||||
(
|
||||
"lxmert",
|
||||
LxmertConfig,
|
||||
),
|
||||
("marian", MarianConfig),
|
||||
("mbart", MBartConfig),
|
||||
("bart", BartConfig),
|
||||
("reformer", ReformerConfig),
|
||||
("longformer", LongformerConfig),
|
||||
("roberta", RobertaConfig),
|
||||
("flaubert", FlaubertConfig),
|
||||
("bert", BertConfig),
|
||||
("openai-gpt", OpenAIGPTConfig),
|
||||
("gpt2", GPT2Config),
|
||||
("transfo-xl", TransfoXLConfig),
|
||||
("xlnet", XLNetConfig),
|
||||
("xlm", XLMConfig),
|
||||
("ctrl", CTRLConfig),
|
||||
("electra", ElectraConfig),
|
||||
("encoder-decoder", EncoderDecoderConfig),
|
||||
("funnel", FunnelConfig),
|
||||
("lxmert", LxmertConfig),
|
||||
]
|
||||
)
|
||||
|
||||
MODEL_NAMES_MAPPING = OrderedDict(
|
||||
[
|
||||
("retribert", "RetriBERT"),
|
||||
("t5", "T5"),
|
||||
("mobilebert", "MobileBERT"),
|
||||
("distilbert", "DistilBERT"),
|
||||
("albert", "ALBERT"),
|
||||
("bert-generation", "Bert Generation"),
|
||||
("camembert", "CamemBERT"),
|
||||
("xlm-roberta", "XLM-RoBERTa"),
|
||||
("pegasus", "Pegasus"),
|
||||
("marian", "Marian"),
|
||||
("mbart", "mBART"),
|
||||
("bart", "BART"),
|
||||
("reformer", "Reformer"),
|
||||
("longformer", "Longformer"),
|
||||
("roberta", "RoBERTa"),
|
||||
("flaubert", "FlauBERT"),
|
||||
("bert", "BERT"),
|
||||
("openai-gpt", "OpenAI GPT"),
|
||||
("gpt2", "OpenAI GPT-2"),
|
||||
("transfo-xl", "Transformer-XL"),
|
||||
("xlnet", "XLNet"),
|
||||
("xlm", "XLM"),
|
||||
("ctrl", "CTRL"),
|
||||
("electra", "ELECTRA"),
|
||||
("encoder-decoder", "Encoder decoder"),
|
||||
("funnel", "Funnel Transformer"),
|
||||
("lxmert", "LXMERT"),
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def _list_model_options(indent, config_to_class=None, use_model_types=True):
|
||||
if config_to_class is None and not use_model_types:
|
||||
raise ValueError("Using `use_model_types=False` requires a `config_to_class` dictionary.")
|
||||
if use_model_types:
|
||||
if config_to_class is None:
|
||||
model_type_to_name = {model_type: config.__name__ for model_type, config in CONFIG_MAPPING.items()}
|
||||
else:
|
||||
model_type_to_name = {
|
||||
model_type: config_to_class[config].__name__
|
||||
for model_type, config in CONFIG_MAPPING.items()
|
||||
if config in config_to_class
|
||||
}
|
||||
lines = [
|
||||
f"{indent}- **{model_type}** -- :class:`~transformers.{cls_name}` ({MODEL_NAMES_MAPPING[model_type]} model)"
|
||||
for model_type, cls_name in model_type_to_name.items()
|
||||
]
|
||||
else:
|
||||
config_to_name = {config.__name__: clas.__name__ for config, clas in config_to_class.items()}
|
||||
config_to_model_name = {
|
||||
config.__name__: MODEL_NAMES_MAPPING[model_type] for model_type, config in CONFIG_MAPPING.items()
|
||||
}
|
||||
lines = [
|
||||
f"{indent}- :class:`~transformers.{config_name}` configuration class: :class:`~transformers.{cls_name}` ({config_to_model_name[config_name]} model)"
|
||||
for config_name, cls_name in config_to_name.items()
|
||||
]
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def replace_list_option_in_docstrings(config_to_class=None, use_model_types=True):
|
||||
def docstring_decorator(fn):
|
||||
docstrings = fn.__doc__
|
||||
lines = docstrings.split("\n")
|
||||
i = 0
|
||||
while i < len(lines) and re.search(r"^(\s*)List options\s*$", lines[i]) is None:
|
||||
i += 1
|
||||
if i < len(lines):
|
||||
indent = re.search(r"^(\s*)List options\s*$", lines[i]).groups()[0]
|
||||
if use_model_types:
|
||||
indent = f"{indent} "
|
||||
lines[i] = _list_model_options(indent, config_to_class=config_to_class, use_model_types=use_model_types)
|
||||
docstrings = "\n".join(lines)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"The function {fn} should have an empty 'List options' in its docstring as placeholder, current docstring is:\n{docstrings}"
|
||||
)
|
||||
fn.__doc__ = docstrings
|
||||
return fn
|
||||
|
||||
return docstring_decorator
|
||||
|
||||
|
||||
class AutoConfig:
|
||||
r"""
|
||||
:class:`~transformers.AutoConfig` is a generic configuration class
|
||||
that will be instantiated as one of the configuration classes of the library
|
||||
when created with the :func:`~transformers.AutoConfig.from_pretrained` class method.
|
||||
This is a generic configuration class that will be instantiated as one of the configuration classes of the library
|
||||
when created with the :meth:`~transformers.AutoConfig.from_pretrained` class method.
|
||||
|
||||
The :func:`~transformers.AutoConfig.from_pretrained` method takes care of returning the correct model class instance
|
||||
This method takes care of returning the correct model class instance
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string.
|
||||
"""
|
||||
@@ -211,6 +220,7 @@ class AutoConfig:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings()
|
||||
def from_pretrained(cls, pretrained_model_name_or_path, **kwargs):
|
||||
r""" Instantiates one of the configuration classes of the library
|
||||
from a pre-trained model configuration.
|
||||
@@ -219,24 +229,7 @@ class AutoConfig:
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
- `t5`: :class:`~transformers.T5Config` (T5 model)
|
||||
- `distilbert`: :class:`~transformers.DistilBertConfig` (DistilBERT model)
|
||||
- `albert`: :class:`~transformers.AlbertConfig` (ALBERT model)
|
||||
- `camembert`: :class:`~transformers.CamembertConfig` (CamemBERT model)
|
||||
- `xlm-roberta`: :class:`~transformers.XLMRobertaConfig` (XLM-RoBERTa model)
|
||||
- `longformer`: :class:`~transformers.LongformerConfig` (Longformer model)
|
||||
- `roberta`: :class:`~transformers.RobertaConfig` (RoBERTa model)
|
||||
- `reformer`: :class:`~transformers.ReformerConfig` (Reformer model)
|
||||
- `bert`: :class:`~transformers.BertConfig` (Bert model)
|
||||
- `openai-gpt`: :class:`~transformers.OpenAIGPTConfig` (OpenAI GPT model)
|
||||
- `gpt2`: :class:`~transformers.GPT2Config` (OpenAI GPT-2 model)
|
||||
- `transfo-xl`: :class:`~transformers.TransfoXLConfig` (Transformer-XL model)
|
||||
- `xlnet`: :class:`~transformers.XLNetConfig` (XLNet model)
|
||||
- `xlm`: :class:`~transformers.XLMConfig` (XLM model)
|
||||
- `ctrl` : :class:`~transformers.CTRLConfig` (CTRL model)
|
||||
- `flaubert` : :class:`~transformers.FlaubertConfig` (Flaubert model)
|
||||
- `electra` : :class:`~transformers.ElectraConfig` (ELECTRA model)
|
||||
- `funnel`: :class:`~transformers.FunnelConfig` (Funnel Transformer model)
|
||||
List options
|
||||
|
||||
Args:
|
||||
pretrained_model_name_or_path (:obj:`string`):
|
||||
|
||||
@@ -0,0 +1,106 @@
|
||||
# coding=utf-8
|
||||
# Copyright 2020 The Google AI Language Team 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.
|
||||
""" BertGeneration model configuration """
|
||||
|
||||
from .configuration_utils import PretrainedConfig
|
||||
|
||||
|
||||
class BertGenerationConfig(PretrainedConfig):
|
||||
r"""
|
||||
This is the configuration class to store the configuration of a :class:`~transformers.BertGenerationPreTrainedModel`.
|
||||
It is used to instantiate a BertGenerationConfig model according to the specified arguments, defining the model architecture.
|
||||
|
||||
Configuration objects inherit from :class:`~transformers.PretrainedConfig` and can be used
|
||||
to control the model outputs. Read the documentation from :class:`~transformers.PretrainedConfig`
|
||||
for more information.
|
||||
|
||||
|
||||
Args:
|
||||
vocab_size (:obj:`int`, `optional`, defaults to 50358):
|
||||
Vocabulary size of the BertGeneration model. Defines the different tokens that
|
||||
can be represented by the `inputs_ids` passed to the forward method of :class:`~transformers.BertGeneration`.
|
||||
hidden_size (:obj:`int`, `optional`, defaults to 1024):
|
||||
Dimensionality of the encoder layers and the pooler layer.
|
||||
num_hidden_layers (:obj:`int`, `optional`, defaults to 24):
|
||||
Number of hidden layers in the Transformer encoder.
|
||||
num_attention_heads (:obj:`int`, `optional`, defaults to 16):
|
||||
Number of attention heads for each attention layer in the Transformer encoder.
|
||||
intermediate_size (:obj:`int`, `optional`, defaults to 3072):
|
||||
Dimensionality of the "intermediate" (i.e., feed-forward) layer in the Transformer encoder.
|
||||
hidden_act (:obj:`str` or :obj:`function`, `optional`, defaults to :obj:`"gelu"`):
|
||||
The non-linear activation function (function or string) in the encoder and pooler.
|
||||
If string, :obj:`"gelu"`, :obj:`"relu"`, :obj:`"swish"` and :obj:`"gelu_new"` are supported.
|
||||
hidden_dropout_prob (:obj:`float`, `optional`, defaults to 0.1):
|
||||
The dropout probabilitiy for all fully connected layers in the embeddings, encoder, and pooler.
|
||||
attention_probs_dropout_prob (:obj:`float`, `optional`, defaults to 0.1):
|
||||
The dropout ratio for the attention probabilities.
|
||||
max_position_embeddings (:obj:`int`, `optional`, defaults to 512):
|
||||
The maximum sequence length that this model might ever be used with.
|
||||
Typically set this to something large just in case (e.g., 512 or 1024 or 2048).
|
||||
initializer_range (:obj:`float`, `optional`, defaults to 0.02):
|
||||
The standard deviation of the truncated_normal_initializer for initializing all weight matrices.
|
||||
layer_norm_eps (:obj:`float`, `optional`, defaults to 1e-12):
|
||||
The epsilon used by the layer normalization layers.
|
||||
gradient_checkpointing (:obj:`bool`, `optional`, defaults to :obj:`False`):
|
||||
If :obj:`True`, use gradient checkpointing to save memory at the expense of slower backward pass.
|
||||
|
||||
Example::
|
||||
|
||||
>>> from transformers import BertGenerationConfig, BertGenerationEncoder
|
||||
|
||||
>>> # Initializing a BertGeneration config
|
||||
>>> configuration = BertGenerationConfig()
|
||||
|
||||
>>> # Initializing a modelfrom the config
|
||||
>>> model = BertGenerationEncoder(configuration)
|
||||
|
||||
>>> # Accessing the model configuration
|
||||
>>> configuration = model.config
|
||||
"""
|
||||
model_type = "bert-generation"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vocab_size=50358,
|
||||
hidden_size=1024,
|
||||
num_hidden_layers=24,
|
||||
num_attention_heads=16,
|
||||
intermediate_size=4096,
|
||||
hidden_act="gelu",
|
||||
hidden_dropout_prob=0.1,
|
||||
attention_probs_dropout_prob=0.1,
|
||||
max_position_embeddings=512,
|
||||
initializer_range=0.02,
|
||||
layer_norm_eps=1e-12,
|
||||
pad_token_id=0,
|
||||
bos_token_id=2,
|
||||
eos_token_id=1,
|
||||
gradient_checkpointing=False,
|
||||
**kwargs
|
||||
):
|
||||
super().__init__(pad_token_id=pad_token_id, bos_token_id=bos_token_id, eos_token_id=eos_token_id, **kwargs)
|
||||
|
||||
self.vocab_size = vocab_size
|
||||
self.hidden_size = hidden_size
|
||||
self.num_hidden_layers = num_hidden_layers
|
||||
self.num_attention_heads = num_attention_heads
|
||||
self.hidden_act = hidden_act
|
||||
self.intermediate_size = intermediate_size
|
||||
self.hidden_dropout_prob = hidden_dropout_prob
|
||||
self.attention_probs_dropout_prob = attention_probs_dropout_prob
|
||||
self.max_position_embeddings = max_position_embeddings
|
||||
self.initializer_range = initializer_range
|
||||
self.layer_norm_eps = layer_norm_eps
|
||||
self.gradient_checkpointing = gradient_checkpointing
|
||||
@@ -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
|
||||
@@ -0,0 +1,88 @@
|
||||
# coding=utf-8
|
||||
# Copyright 2020 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.
|
||||
"""Convert Seq2Seq TF Hub checkpoint."""
|
||||
|
||||
|
||||
import argparse
|
||||
|
||||
from transformers import (
|
||||
BertConfig,
|
||||
BertGenerationConfig,
|
||||
BertGenerationDecoder,
|
||||
BertGenerationEncoder,
|
||||
load_tf_weights_in_bert_generation,
|
||||
logging,
|
||||
)
|
||||
|
||||
|
||||
logging.set_verbosity_info()
|
||||
|
||||
|
||||
def convert_tf_checkpoint_to_pytorch(tf_hub_path, pytorch_dump_path, is_encoder_named_decoder, vocab_size, is_encoder):
|
||||
# Initialise PyTorch model
|
||||
bert_config = BertConfig.from_pretrained(
|
||||
"bert-large-cased",
|
||||
vocab_size=vocab_size,
|
||||
max_position_embeddings=512,
|
||||
is_decoder=True,
|
||||
add_cross_attention=True,
|
||||
)
|
||||
bert_config_dict = bert_config.to_dict()
|
||||
del bert_config_dict["type_vocab_size"]
|
||||
config = BertGenerationConfig(**bert_config_dict)
|
||||
if is_encoder:
|
||||
model = BertGenerationEncoder(config)
|
||||
else:
|
||||
model = BertGenerationDecoder(config)
|
||||
print("Building PyTorch model from configuration: {}".format(str(config)))
|
||||
|
||||
# Load weights from tf checkpoint
|
||||
load_tf_weights_in_bert_generation(
|
||||
model,
|
||||
tf_hub_path,
|
||||
model_class="bert",
|
||||
is_encoder_named_decoder=is_encoder_named_decoder,
|
||||
is_encoder=is_encoder,
|
||||
)
|
||||
|
||||
# Save pytorch-model
|
||||
print("Save PyTorch model and config to {}".format(pytorch_dump_path))
|
||||
model.save_pretrained(pytorch_dump_path)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
# Required parameters
|
||||
parser.add_argument(
|
||||
"--tf_hub_path", default=None, type=str, required=True, help="Path to the TensorFlow checkpoint path."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--pytorch_dump_path", default=None, type=str, required=True, help="Path to the output PyTorch model."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--is_encoder_named_decoder",
|
||||
action="store_true",
|
||||
help="If decoder has to be renamed to encoder in PyTorch model.",
|
||||
)
|
||||
parser.add_argument("--is_encoder", action="store_true", help="If model is an encoder.")
|
||||
parser.add_argument("--vocab_size", default=50358, type=int, help="Vocab size of model")
|
||||
args = parser.parse_args()
|
||||
convert_tf_checkpoint_to_pytorch(
|
||||
args.tf_hub_path,
|
||||
args.pytorch_dump_path,
|
||||
args.is_encoder_named_decoder,
|
||||
args.vocab_size,
|
||||
is_encoder=args.is_encoder,
|
||||
)
|
||||
@@ -149,7 +149,7 @@ class DataCollatorForLanguageModeling:
|
||||
) -> torch.Tensor:
|
||||
# In order to accept both lists of lists and lists of Tensors
|
||||
if isinstance(examples[0], (list, tuple)):
|
||||
examples = [torch.Tensor(e) for e in examples]
|
||||
examples = [torch.tensor(e, dtype=torch.long) for e in examples]
|
||||
length_of_first = examples[0].size(0)
|
||||
are_tensors_same_length = all(x.size(0) == length_of_first for x in examples)
|
||||
if are_tensors_same_length:
|
||||
@@ -198,6 +198,75 @@ class DataCollatorForLanguageModeling:
|
||||
return inputs, labels
|
||||
|
||||
|
||||
@dataclass
|
||||
class DataCollatorForSOP(DataCollatorForLanguageModeling):
|
||||
"""
|
||||
Data collator used for sentence order prediction task.
|
||||
- collates batches of tensors, honoring their tokenizer's pad_token
|
||||
- preprocesses batches for both masked language modeling and sentence order prediction
|
||||
"""
|
||||
|
||||
def __call__(self, examples: List[Dict[str, torch.Tensor]]) -> Dict[str, torch.Tensor]:
|
||||
input_ids = [example["input_ids"] for example in examples]
|
||||
input_ids = self._tensorize_batch(input_ids)
|
||||
input_ids, labels, attention_mask = self.mask_tokens(input_ids)
|
||||
|
||||
token_type_ids = [example["token_type_ids"] for example in examples]
|
||||
# size of segment_ids varied because randomness, padding zero to the end as the orignal implementation
|
||||
token_type_ids = pad_sequence(token_type_ids, batch_first=True, padding_value=self.tokenizer.pad_token_id)
|
||||
|
||||
sop_label_list = [example["sentence_order_label"] for example in examples]
|
||||
sentence_order_label = torch.stack(sop_label_list)
|
||||
|
||||
return {
|
||||
"input_ids": input_ids,
|
||||
"labels": labels,
|
||||
"attention_mask": attention_mask,
|
||||
"token_type_ids": token_type_ids,
|
||||
"sentence_order_label": sentence_order_label,
|
||||
}
|
||||
|
||||
def mask_tokens(self, inputs: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Prepare masked tokens inputs/labels/attention_mask for masked language modeling: 80% MASK, 10% random, 10% original.
|
||||
N-gram not applied yet.
|
||||
"""
|
||||
if self.tokenizer.mask_token is None:
|
||||
raise ValueError(
|
||||
"This tokenizer does not have a mask token which is necessary for masked language modeling. Remove the --mlm flag if you want to use this tokenizer."
|
||||
)
|
||||
|
||||
labels = inputs.clone()
|
||||
# We sample a few tokens in each sequence for masked-LM training (with probability args.mlm_probability defaults to 0.15 in Bert/RoBERTa)
|
||||
probability_matrix = torch.full(labels.shape, self.mlm_probability)
|
||||
special_tokens_mask = [
|
||||
self.tokenizer.get_special_tokens_mask(val, already_has_special_tokens=True) for val in labels.tolist()
|
||||
]
|
||||
probability_matrix.masked_fill_(torch.tensor(special_tokens_mask, dtype=torch.bool), value=0.0)
|
||||
if self.tokenizer._pad_token is not None:
|
||||
padding_mask = labels.eq(self.tokenizer.pad_token_id)
|
||||
probability_matrix.masked_fill_(padding_mask, value=0.0)
|
||||
masked_indices = torch.bernoulli(probability_matrix).bool()
|
||||
# probability be `1` (masked), however in albert model attention mask `0` means masked, revert the value
|
||||
attention_mask = (~masked_indices).float()
|
||||
if self.tokenizer._pad_token is not None:
|
||||
attention_padding_mask = labels.eq(self.tokenizer.pad_token_id)
|
||||
attention_mask.masked_fill_(attention_padding_mask, value=1.0)
|
||||
labels[~masked_indices] = -100 # We only compute loss on masked tokens, -100 is default for CE compute
|
||||
|
||||
# 80% of the time, we replace masked input tokens with tokenizer.mask_token ([MASK])
|
||||
indices_replaced = torch.bernoulli(torch.full(labels.shape, 0.8)).bool() & masked_indices
|
||||
inputs[indices_replaced] = self.tokenizer.convert_tokens_to_ids(self.tokenizer.mask_token)
|
||||
|
||||
# 10% of the time, we replace masked input tokens with random word
|
||||
indices_random = torch.bernoulli(torch.full(labels.shape, 0.5)).bool() & masked_indices & ~indices_replaced
|
||||
random_words = torch.randint(len(self.tokenizer), labels.shape, dtype=torch.long)
|
||||
inputs[indices_random] = random_words[indices_random]
|
||||
|
||||
# The rest of the time (10% of the time) we keep the masked input tokens unchanged
|
||||
return inputs, labels, attention_mask
|
||||
|
||||
|
||||
@dataclass
|
||||
class DataCollatorForPermutationLanguageModeling:
|
||||
"""
|
||||
|
||||
@@ -3,5 +3,10 @@
|
||||
# module, but to preserve other warnings. So, don't check this module at all.
|
||||
|
||||
from .glue import GlueDataset, GlueDataTrainingArguments
|
||||
from .language_modeling import LineByLineTextDataset, TextDataset, TextDatasetForNextSentencePrediction
|
||||
from .language_modeling import (
|
||||
LineByLineTextDataset,
|
||||
LineByLineWithSOPTextDataset,
|
||||
TextDataset,
|
||||
TextDatasetForNextSentencePrediction,
|
||||
)
|
||||
from .squad import SquadDataset, SquadDataTrainingArguments
|
||||
@@ -1,7 +1,8 @@
|
||||
import os
|
||||
import pickle
|
||||
import random
|
||||
import time
|
||||
from typing import Optional
|
||||
from typing import Dict, Optional
|
||||
|
||||
import torch
|
||||
from torch.utils.data.dataset import Dataset
|
||||
@@ -113,6 +114,147 @@ class LineByLineTextDataset(Dataset):
|
||||
return torch.tensor(self.examples[i], dtype=torch.long)
|
||||
|
||||
|
||||
class LineByLineWithSOPTextDataset(Dataset):
|
||||
"""
|
||||
Dataset for sentence order prediction task, prepare sentence pairs for SOP task
|
||||
"""
|
||||
|
||||
def __init__(self, tokenizer: PreTrainedTokenizer, file_dir: str, block_size: int):
|
||||
assert os.path.isdir(file_dir)
|
||||
logger.info(f"Creating features from dataset file folder at {file_dir}")
|
||||
self.examples = []
|
||||
# TODO: randomness could apply a random seed, ex. rng = random.Random(random_seed)
|
||||
# file path looks like ./dataset/wiki_1, ./dataset/wiki_2
|
||||
for file_name in os.listdir(file_dir):
|
||||
file_path = os.path.join(file_dir, file_name)
|
||||
assert os.path.isfile(file_path)
|
||||
article_open = False
|
||||
with open(file_path, encoding="utf-8") as f:
|
||||
original_lines = f.readlines()
|
||||
article_lines = []
|
||||
for line in original_lines:
|
||||
if "<doc id=" in line:
|
||||
article_open = True
|
||||
elif "</doc>" in line:
|
||||
article_open = False
|
||||
document = [
|
||||
tokenizer.convert_tokens_to_ids(tokenizer.tokenize(line))
|
||||
for line in article_lines[1:]
|
||||
if (len(line) > 0 and not line.isspace())
|
||||
]
|
||||
|
||||
examples = self.create_examples_from_document(document, block_size, tokenizer)
|
||||
self.examples.extend(examples)
|
||||
article_lines = []
|
||||
else:
|
||||
if article_open:
|
||||
article_lines.append(line)
|
||||
|
||||
logger.info("Dataset parse finished.")
|
||||
|
||||
def create_examples_from_document(self, document, block_size, tokenizer, short_seq_prob=0.1):
|
||||
"""Creates examples for a single document."""
|
||||
|
||||
# Account for special tokens
|
||||
max_num_tokens = block_size - tokenizer.num_special_tokens_to_add(pair=True)
|
||||
|
||||
# We *usually* want to fill up the entire sequence since we are padding
|
||||
# to `block_size` anyways, so short sequences are generally wasted
|
||||
# computation. However, we *sometimes*
|
||||
# (i.e., short_seq_prob == 0.1 == 10% of the time) want to use shorter
|
||||
# sequences to minimize the mismatch between pre-training and fine-tuning.
|
||||
# The `target_seq_length` is just a rough target however, whereas
|
||||
# `block_size` is a hard limit.
|
||||
target_seq_length = max_num_tokens
|
||||
if random.random() < short_seq_prob:
|
||||
target_seq_length = random.randint(2, max_num_tokens)
|
||||
|
||||
# We DON'T just concatenate all of the tokens from a document into a long
|
||||
# sequence and choose an arbitrary split point because this would make the
|
||||
# next sentence prediction task too easy. Instead, we split the input into
|
||||
# segments "A" and "B" based on the actual "sentences" provided by the user
|
||||
# input.
|
||||
examples = []
|
||||
current_chunk = [] # a buffer stored current working segments
|
||||
current_length = 0
|
||||
i = 0
|
||||
while i < len(document):
|
||||
segment = document[i] # get a segment
|
||||
if not segment:
|
||||
i += 1
|
||||
continue
|
||||
current_chunk.append(segment) # add a segment to current chunk
|
||||
current_length += len(segment) # overall token length
|
||||
# if current length goes to the target length or reaches the end of file, start building token a and b
|
||||
if i == len(document) - 1 or current_length >= target_seq_length:
|
||||
if current_chunk:
|
||||
# `a_end` is how many segments from `current_chunk` go into the `A` (first) sentence.
|
||||
a_end = 1
|
||||
# if current chunk has more than 2 sentences, pick part of it `A` (first) sentence
|
||||
if len(current_chunk) >= 2:
|
||||
a_end = random.randint(1, len(current_chunk) - 1)
|
||||
# token a
|
||||
tokens_a = []
|
||||
for j in range(a_end):
|
||||
tokens_a.extend(current_chunk[j])
|
||||
|
||||
# token b
|
||||
tokens_b = []
|
||||
for j in range(a_end, len(current_chunk)):
|
||||
tokens_b.extend(current_chunk[j])
|
||||
|
||||
if len(tokens_a) == 0 or len(tokens_b) == 0:
|
||||
continue
|
||||
|
||||
# switch tokens_a and tokens_b randomly
|
||||
if random.random() < 0.5:
|
||||
is_next = False
|
||||
tokens_a, tokens_b = tokens_b, tokens_a
|
||||
else:
|
||||
is_next = True
|
||||
|
||||
def truncate_seq_pair(tokens_a, tokens_b, max_num_tokens):
|
||||
"""Truncates a pair of sequences to a maximum sequence length."""
|
||||
while True:
|
||||
total_length = len(tokens_a) + len(tokens_b)
|
||||
if total_length <= max_num_tokens:
|
||||
break
|
||||
trunc_tokens = tokens_a if len(tokens_a) > len(tokens_b) else tokens_b
|
||||
assert len(trunc_tokens) >= 1
|
||||
# We want to sometimes truncate from the front and sometimes from the
|
||||
# back to add more randomness and avoid biases.
|
||||
if random.random() < 0.5:
|
||||
del trunc_tokens[0]
|
||||
else:
|
||||
trunc_tokens.pop()
|
||||
|
||||
truncate_seq_pair(tokens_a, tokens_b, max_num_tokens)
|
||||
assert len(tokens_a) >= 1
|
||||
assert len(tokens_b) >= 1
|
||||
|
||||
# add special tokens
|
||||
input_ids = tokenizer.build_inputs_with_special_tokens(tokens_a, tokens_b)
|
||||
# add token type ids, 0 for sentence a, 1 for sentence b
|
||||
token_type_ids = tokenizer.create_token_type_ids_from_sequences(tokens_a, tokens_b)
|
||||
|
||||
example = {
|
||||
"input_ids": torch.tensor(input_ids, dtype=torch.long),
|
||||
"token_type_ids": torch.tensor(token_type_ids, dtype=torch.long),
|
||||
"sentence_order_label": torch.tensor(0 if is_next else 1, dtype=torch.long),
|
||||
}
|
||||
examples.append(example)
|
||||
current_chunk = [] # clear current chunk
|
||||
current_length = 0 # reset current text length
|
||||
i += 1 # go to next line
|
||||
return examples
|
||||
|
||||
def __len__(self):
|
||||
return len(self.examples)
|
||||
|
||||
def __getitem__(self, i) -> Dict[str, torch.tensor]:
|
||||
return self.examples[i]
|
||||
|
||||
|
||||
class TextDatasetForNextSentencePrediction(Dataset):
|
||||
"""
|
||||
This will be superseded by a framework-agnostic approach
|
||||
|
||||
@@ -66,12 +66,12 @@ except (ImportError, AssertionError):
|
||||
|
||||
|
||||
try:
|
||||
import nlp # noqa: F401
|
||||
import datasets # noqa: F401
|
||||
|
||||
_nlp_available = True
|
||||
_datasets_available = True
|
||||
|
||||
except ImportError:
|
||||
_nlp_available = False
|
||||
_datasets_available = False
|
||||
|
||||
try:
|
||||
from torch.hub import _get_torch_home
|
||||
@@ -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")
|
||||
|
||||
|
||||
@@ -133,7 +142,9 @@ CONFIG_NAME = "config.json"
|
||||
MODEL_CARD_NAME = "modelcard.json"
|
||||
|
||||
|
||||
MULTIPLE_CHOICE_DUMMY_INPUTS = [[[0], [1]], [[0], [1]]]
|
||||
MULTIPLE_CHOICE_DUMMY_INPUTS = [
|
||||
[[0, 1, 0, 1], [1, 0, 0, 1]]
|
||||
] * 2 # Needs to have 0s and 1s only since XLM uses it for langs too.
|
||||
DUMMY_INPUTS = [[7, 6, 0, 0, 1], [1, 2, 3, 0, 0], [0, 0, 0, 4, 5]]
|
||||
DUMMY_MASK = [[1, 1, 1, 1, 1], [1, 1, 1, 0, 0], [0, 0, 0, 1, 1]]
|
||||
|
||||
@@ -153,8 +164,8 @@ def is_torch_tpu_available():
|
||||
return _torch_tpu_available
|
||||
|
||||
|
||||
def is_nlp_available():
|
||||
return _nlp_available
|
||||
def is_datasets_available():
|
||||
return _datasets_available
|
||||
|
||||
|
||||
def is_psutil_available():
|
||||
@@ -169,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 "")
|
||||
@@ -530,7 +545,7 @@ def add_code_sample_docstrings(*docstr, tokenizer_class=None, checkpoint=None, o
|
||||
code_sample = TF_MASKED_LM_SAMPLE if is_tf_class else PT_MASKED_LM_SAMPLE
|
||||
elif "LMHead" in model_class:
|
||||
code_sample = TF_CAUSAL_LM_SAMPLE if is_tf_class else PT_CAUSAL_LM_SAMPLE
|
||||
elif "Model" in model_class:
|
||||
elif "Model" in model_class or "Encoder" in model_class:
|
||||
code_sample = TF_BASE_MODEL_SAMPLE if is_tf_class else PT_BASE_MODEL_SAMPLE
|
||||
else:
|
||||
raise ValueError(f"Docstring can't be built for model {model_class}")
|
||||
@@ -986,6 +1001,8 @@ class ModelOutput(OrderedDict):
|
||||
setattr(self, element[0], element[1])
|
||||
if element[1] is not None:
|
||||
self[element[0]] = element[1]
|
||||
elif first_field is not None:
|
||||
self[class_fields[0].name] = first_field
|
||||
else:
|
||||
for field in class_fields:
|
||||
v = getattr(self, field.name)
|
||||
@@ -1011,6 +1028,18 @@ class ModelOutput(OrderedDict):
|
||||
else:
|
||||
return self.to_tuple()[k]
|
||||
|
||||
def __setattr__(self, name, value):
|
||||
if name in self.keys() and value is not None:
|
||||
# Don't call self.__setitem__ to avoid recursion errors
|
||||
super().__setitem__(name, value)
|
||||
super().__setattr__(name, value)
|
||||
|
||||
def __setitem__(self, key, value):
|
||||
# Will raise a KeyException if needed
|
||||
super().__setitem__(key, value)
|
||||
# Don't call self.__setattr__ to avoid recursion errors
|
||||
super().__setattr__(key, value)
|
||||
|
||||
def to_tuple(self) -> Tuple[Any]:
|
||||
"""
|
||||
Convert self to a tuple containing all the attributes/keys that are not ``None``.
|
||||
|
||||
@@ -383,7 +383,11 @@ class GenerationMixin:
|
||||
# see if BOS token can be used for decoder_start_token_id
|
||||
if bos_token_id is not None:
|
||||
decoder_start_token_id = bos_token_id
|
||||
elif hasattr(self.config, "decoder") and hasattr(self.config.decoder, "bos_token_id"):
|
||||
elif (
|
||||
hasattr(self.config, "decoder")
|
||||
and hasattr(self.config.decoder, "bos_token_id")
|
||||
and self.config.decoder.bos_token_id is not None
|
||||
):
|
||||
decoder_start_token_id = self.config.decoder.bos_token_id
|
||||
else:
|
||||
raise ValueError(
|
||||
@@ -395,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:
|
||||
|
||||
+125
-186
@@ -23,6 +23,7 @@ from .configuration_auto import (
|
||||
AutoConfig,
|
||||
BartConfig,
|
||||
BertConfig,
|
||||
BertGenerationConfig,
|
||||
CamembertConfig,
|
||||
CTRLConfig,
|
||||
DistilBertConfig,
|
||||
@@ -45,6 +46,7 @@ from .configuration_auto import (
|
||||
XLMConfig,
|
||||
XLMRobertaConfig,
|
||||
XLNetConfig,
|
||||
replace_list_option_in_docstrings,
|
||||
)
|
||||
from .configuration_marian import MarianConfig
|
||||
from .configuration_utils import PretrainedConfig
|
||||
@@ -73,6 +75,7 @@ from .modeling_bert import (
|
||||
BertLMHeadModel,
|
||||
BertModel,
|
||||
)
|
||||
from .modeling_bert_generation import BertGenerationDecoder, BertGenerationEncoder
|
||||
from .modeling_camembert import (
|
||||
CamembertForCausalLM,
|
||||
CamembertForMaskedLM,
|
||||
@@ -213,6 +216,7 @@ MODEL_MAPPING = OrderedDict(
|
||||
(ReformerConfig, ReformerModel),
|
||||
(FunnelConfig, FunnelModel),
|
||||
(LxmertConfig, LxmertModel),
|
||||
(BertGenerationConfig, BertGenerationEncoder),
|
||||
]
|
||||
)
|
||||
|
||||
@@ -284,6 +288,7 @@ MODEL_FOR_CAUSAL_LM_MAPPING = OrderedDict(
|
||||
), # XLM can be MLM and CLM => model should be split similar to BERT; leave here for now
|
||||
(CTRLConfig, CTRLLMHeadModel),
|
||||
(ReformerConfig, ReformerModelWithLMHead),
|
||||
(BertGenerationConfig, BertGenerationDecoder),
|
||||
]
|
||||
)
|
||||
|
||||
@@ -412,6 +417,7 @@ class AutoModel:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(MODEL_MAPPING, use_model_types=False)
|
||||
def from_config(cls, config):
|
||||
r"""Instantiates one of the base model classes of the library
|
||||
from a configuration.
|
||||
@@ -425,19 +431,7 @@ class AutoModel:
|
||||
config (:class:`~transformers.PretrainedConfig`):
|
||||
The model class to instantiate is selected based on the configuration class:
|
||||
|
||||
- isInstance of `distilbert` configuration class: :class:`~transformers.DistilBertModel` (DistilBERT model)
|
||||
- isInstance of `longformer` configuration class: :class:`~transformers.LongformerModel` (Longformer model)
|
||||
- isInstance of `roberta` configuration class: :class:`~transformers.RobertaModel` (RoBERTa model)
|
||||
- isInstance of `bert` configuration class: :class:`~transformers.BertModel` (Bert model)
|
||||
- isInstance of `openai-gpt` configuration class: :class:`~transformers.OpenAIGPTModel` (OpenAI GPT model)
|
||||
- isInstance of `gpt2` configuration class: :class:`~transformers.GPT2Model` (OpenAI GPT-2 model)
|
||||
- isInstance of `ctrl` configuration class: :class:`~transformers.CTRLModel` (Salesforce CTRL model)
|
||||
- isInstance of `transfo-xl` configuration class: :class:`~transformers.TransfoXLModel` (Transformer-XL model)
|
||||
- isInstance of `xlnet` configuration class: :class:`~transformers.XLNetModel` (XLNet model)
|
||||
- isInstance of `xlm` configuration class: :class:`~transformers.XLMModel` (XLM model)
|
||||
- isInstance of `flaubert` configuration class: :class:`~transformers.FlaubertModel` (Flaubert model)
|
||||
- isInstance of `electra` configuration class: :class:`~transformers.ElectraModel` (Electra model)
|
||||
- isInstance of `funnel` configuration class: :class:`~transformers.FunnelModel` (Funnel Transformer model)
|
||||
List options
|
||||
|
||||
Examples::
|
||||
|
||||
@@ -455,6 +449,7 @@ class AutoModel:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(MODEL_MAPPING)
|
||||
def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
|
||||
r"""Instantiates one of the base model classes of the library
|
||||
from a pre-trained model configuration.
|
||||
@@ -463,23 +458,7 @@ class AutoModel:
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
- `t5`: :class:`~transformers.T5Model` (T5 model)
|
||||
- `distilbert`: :class:`~transformers.DistilBertModel` (DistilBERT model)
|
||||
- `albert`: :class:`~transformers.AlbertModel` (ALBERT model)
|
||||
- `camembert`: :class:`~transformers.CamembertModel` (CamemBERT model)
|
||||
- `xlm-roberta`: :class:`~transformers.XLMRobertaModel` (XLM-RoBERTa model)
|
||||
- `longformer` :class:`~transformers.LongformerModel` (Longformer model)
|
||||
- `roberta`: :class:`~transformers.RobertaModel` (RoBERTa model)
|
||||
- `bert`: :class:`~transformers.BertModel` (Bert model)
|
||||
- `openai-gpt`: :class:`~transformers.OpenAIGPTModel` (OpenAI GPT model)
|
||||
- `gpt2`: :class:`~transformers.GPT2Model` (OpenAI GPT-2 model)
|
||||
- `transfo-xl`: :class:`~transformers.TransfoXLModel` (Transformer-XL model)
|
||||
- `xlnet`: :class:`~transformers.XLNetModel` (XLNet model)
|
||||
- `xlm`: :class:`~transformers.XLMModel` (XLM model)
|
||||
- `ctrl`: :class:`~transformers.CTRLModel` (Salesforce CTRL model)
|
||||
- `flaubert`: :class:`~transformers.FlaubertModel` (Flaubert model)
|
||||
- `electra`: :class:`~transformers.ElectraModel` (Electra model)
|
||||
- `funnel`: :class:`~transformers.FunnelModel` (Funnel Transformer model)
|
||||
List options
|
||||
|
||||
The model is set in evaluation mode by default using `model.eval()` (Dropout modules are deactivated)
|
||||
To train the model, you should first set it back in training mode with `model.train()`
|
||||
@@ -571,6 +550,7 @@ class AutoModelForPreTraining:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(MODEL_FOR_PRETRAINING_MAPPING, use_model_types=False)
|
||||
def from_config(cls, config):
|
||||
r"""Instantiates one of the base model classes of the library
|
||||
from a configuration.
|
||||
@@ -584,18 +564,7 @@ class AutoModelForPreTraining:
|
||||
config (:class:`~transformers.PretrainedConfig`):
|
||||
The model class to instantiate is selected based on the configuration class:
|
||||
|
||||
- isInstance of `distilbert` configuration class: :class:`~transformers.DistilBertForMaskedLM` (DistilBERT model)
|
||||
- isInstance of `longformer` configuration class: :class:`~transformers.LongformerForMaskedLM` (Longformer model)
|
||||
- isInstance of `roberta` configuration class: :class:`~transformers.RobertaForMaskedLM` (RoBERTa model)
|
||||
- isInstance of `bert` configuration class: :class:`~transformers.BertForPreTraining` (Bert model)
|
||||
- isInstance of `openai-gpt` configuration class: :class:`~transformers.OpenAIGPTLMHeadModel` (OpenAI GPT model)
|
||||
- isInstance of `gpt2` configuration class: :class:`~transformers.GPT2LMHeadModel` (OpenAI GPT-2 model)
|
||||
- isInstance of `ctrl` configuration class: :class:`~transformers.CTRLLMHeadModel` (Salesforce CTRL model)
|
||||
- isInstance of `transfo-xl` configuration class: :class:`~transformers.TransfoXLLMHeadModel` (Transformer-XL model)
|
||||
- isInstance of `xlnet` configuration class: :class:`~transformers.XLNetLMHeadModel` (XLNet model)
|
||||
- isInstance of `xlm` configuration class: :class:`~transformers.XLMWithLMHeadModel` (XLM model)
|
||||
- isInstance of `flaubert` configuration class: :class:`~transformers.FlaubertWithLMHeadModel` (Flaubert model)
|
||||
- isInstance of `electra` configuration class: :class:`~transformers.ElectraForPreTraining` (Electra model)
|
||||
List options
|
||||
|
||||
Examples::
|
||||
|
||||
@@ -613,6 +582,7 @@ class AutoModelForPreTraining:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(MODEL_FOR_PRETRAINING_MAPPING)
|
||||
def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
|
||||
r"""Instantiates one of the model classes of the library -with the architecture used for pretraining this model– from a pre-trained model configuration.
|
||||
|
||||
@@ -620,22 +590,7 @@ class AutoModelForPreTraining:
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
- `t5`: :class:`~transformers.T5ModelWithLMHead` (T5 model)
|
||||
- `distilbert`: :class:`~transformers.DistilBertForMaskedLM` (DistilBERT model)
|
||||
- `albert`: :class:`~transformers.AlbertForMaskedLM` (ALBERT model)
|
||||
- `camembert`: :class:`~transformers.CamembertForMaskedLM` (CamemBERT model)
|
||||
- `xlm-roberta`: :class:`~transformers.XLMRobertaForMaskedLM` (XLM-RoBERTa model)
|
||||
- `longformer`: :class:`~transformers.LongformerForMaskedLM` (Longformer model)
|
||||
- `roberta`: :class:`~transformers.RobertaForMaskedLM` (RoBERTa model)
|
||||
- `bert`: :class:`~transformers.BertForPreTraining` (Bert model)
|
||||
- `openai-gpt`: :class:`~transformers.OpenAIGPTLMHeadModel` (OpenAI GPT model)
|
||||
- `gpt2`: :class:`~transformers.GPT2LMHeadModel` (OpenAI GPT-2 model)
|
||||
- `transfo-xl`: :class:`~transformers.TransfoXLLMHeadModel` (Transformer-XL model)
|
||||
- `xlnet`: :class:`~transformers.XLNetLMHeadModel` (XLNet model)
|
||||
- `xlm`: :class:`~transformers.XLMWithLMHeadModel` (XLM model)
|
||||
- `ctrl`: :class:`~transformers.CTRLLMHeadModel` (Salesforce CTRL model)
|
||||
- `flaubert`: :class:`~transformers.FlaubertWithLMHeadModel` (Flaubert model)
|
||||
- `electra`: :class:`~transformers.ElectraForPreTraining` (Electra model)
|
||||
List options
|
||||
|
||||
The model is set in evaluation mode by default using `model.eval()` (Dropout modules are deactivated)
|
||||
To train the model, you should first set it back in training mode with `model.train()`
|
||||
@@ -722,6 +677,7 @@ class AutoModelWithLMHead:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(MODEL_WITH_LM_HEAD_MAPPING, use_model_types=False)
|
||||
def from_config(cls, config):
|
||||
r"""Instantiates one of the base model classes of the library
|
||||
from a configuration.
|
||||
@@ -735,19 +691,7 @@ class AutoModelWithLMHead:
|
||||
config (:class:`~transformers.PretrainedConfig`):
|
||||
The model class to instantiate is selected based on the configuration class:
|
||||
|
||||
- isInstance of `distilbert` configuration class: :class:`~transformers.DistilBertForMaskedLM` (DistilBERT model)
|
||||
- isInstance of `longformer` configuration class: :class:`~transformers.LongformerForMaskedLM` (Longformer model)
|
||||
- isInstance of `roberta` configuration class: :class:`~transformers.RobertaForMaskedLM` (RoBERTa model)
|
||||
- isInstance of `bert` configuration class: :class:`~transformers.BertForMaskedLM` (Bert model)
|
||||
- isInstance of `openai-gpt` configuration class: :class:`~transformers.OpenAIGPTLMHeadModel` (OpenAI GPT model)
|
||||
- isInstance of `gpt2` configuration class: :class:`~transformers.GPT2LMHeadModel` (OpenAI GPT-2 model)
|
||||
- isInstance of `ctrl` configuration class: :class:`~transformers.CTRLLMHeadModel` (Salesforce CTRL model)
|
||||
- isInstance of `transfo-xl` configuration class: :class:`~transformers.TransfoXLLMHeadModel` (Transformer-XL model)
|
||||
- isInstance of `xlnet` configuration class: :class:`~transformers.XLNetLMHeadModel` (XLNet model)
|
||||
- isInstance of `xlm` configuration class: :class:`~transformers.XLMWithLMHeadModel` (XLM model)
|
||||
- isInstance of `flaubert` configuration class: :class:`~transformers.FlaubertWithLMHeadModel` (Flaubert model)
|
||||
- isInstance of `electra` configuration class: :class:`~transformers.ElectraForMaskedLM` (Electra model)
|
||||
- isInstance of `funnel` configuration class: :class:`~transformers.FunnelForMaskedLM` (Funnel Transformer model)
|
||||
List options
|
||||
|
||||
Examples::
|
||||
|
||||
@@ -769,6 +713,7 @@ class AutoModelWithLMHead:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(MODEL_WITH_LM_HEAD_MAPPING)
|
||||
def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
|
||||
r"""Instantiates one of the language modeling model classes of the library
|
||||
from a pre-trained model configuration.
|
||||
@@ -777,23 +722,7 @@ class AutoModelWithLMHead:
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
- `t5`: :class:`~transformers.T5ForConditionalGeneration` (T5 model)
|
||||
- `distilbert`: :class:`~transformers.DistilBertForMaskedLM` (DistilBERT model)
|
||||
- `albert`: :class:`~transformers.AlbertForMaskedLM` (ALBERT model)
|
||||
- `camembert`: :class:`~transformers.CamembertForMaskedLM` (CamemBERT model)
|
||||
- `xlm-roberta`: :class:`~transformers.XLMRobertaForMaskedLM` (XLM-RoBERTa model)
|
||||
- `longformer`: :class:`~transformers.LongformerForMaskedLM` (Longformer model)
|
||||
- `roberta`: :class:`~transformers.RobertaForMaskedLM` (RoBERTa model)
|
||||
- `bert`: :class:`~transformers.BertForMaskedLM` (Bert model)
|
||||
- `openai-gpt`: :class:`~transformers.OpenAIGPTLMHeadModel` (OpenAI GPT model)
|
||||
- `gpt2`: :class:`~transformers.GPT2LMHeadModel` (OpenAI GPT-2 model)
|
||||
- `transfo-xl`: :class:`~transformers.TransfoXLLMHeadModel` (Transformer-XL model)
|
||||
- `xlnet`: :class:`~transformers.XLNetLMHeadModel` (XLNet model)
|
||||
- `xlm`: :class:`~transformers.XLMWithLMHeadModel` (XLM model)
|
||||
- `ctrl`: :class:`~transformers.CTRLLMHeadModel` (Salesforce CTRL model)
|
||||
- `flaubert`: :class:`~transformers.FlaubertWithLMHeadModel` (Flaubert model)
|
||||
- `electra`: :class:`~transformers.ElectraForMaskedLM` (Electra model)
|
||||
- `funnel`: :class:`~transformers.FunnelForMaskedLM` (Funnel Transformer model)
|
||||
List options
|
||||
|
||||
The model is set in evaluation mode by default using `model.eval()` (Dropout modules are deactivated)
|
||||
To train the model, you should first set it back in training mode with `model.train()`
|
||||
@@ -884,6 +813,7 @@ class AutoModelForCausalLM:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(MODEL_FOR_CAUSAL_LM_MAPPING, use_model_types=False)
|
||||
def from_config(cls, config):
|
||||
r"""Instantiates one of the base model classes of the library
|
||||
from a configuration.
|
||||
@@ -897,13 +827,7 @@ class AutoModelForCausalLM:
|
||||
config (:class:`~transformers.PretrainedConfig`):
|
||||
The model class to instantiate is selected based on the configuration class:
|
||||
|
||||
- isInstance of `bert` configuration class: :class:`~transformers.BertLMHeadModel` (Bert model)
|
||||
- isInstance of `openai-gpt` configuration class: :class:`~transformers.OpenAIGPTLMHeadModel` (OpenAI GPT model)
|
||||
- isInstance of `gpt2` configuration class: :class:`~transformers.GPT2LMHeadModel` (OpenAI GPT-2 model)
|
||||
- isInstance of `ctrl` configuration class: :class:`~transformers.CTRLLMHeadModel` (Salesforce CTRL model)
|
||||
- isInstance of `transfo-xl` configuration class: :class:`~transformers.TransfoXLLMHeadModel` (Transformer-XL model)
|
||||
- isInstance of `xlnet` configuration class: :class:`~transformers.XLNetLMHeadModel` (XLNet model)
|
||||
- isInstance of `reformer` configuration class: :class:`~transformers.ReformerModelWithLMHead` (Reformer model)
|
||||
List options
|
||||
|
||||
Examples::
|
||||
|
||||
@@ -921,6 +845,7 @@ class AutoModelForCausalLM:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(MODEL_FOR_CAUSAL_LM_MAPPING)
|
||||
def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
|
||||
r"""Instantiates one of the language modeling model classes of the library
|
||||
from a pre-trained model configuration.
|
||||
@@ -929,13 +854,7 @@ class AutoModelForCausalLM:
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
- `bert`: :class:`~transformers.BertLMHeadModel` (Bert model)
|
||||
- `openai-gpt`: :class:`~transformers.OpenAIGPTLMHeadModel` (OpenAI GPT model)
|
||||
- `gpt2`: :class:`~transformers.GPT2LMHeadModel` (OpenAI GPT-2 model)
|
||||
- `transfo-xl`: :class:`~transformers.TransfoXLLMHeadModel` (Transformer-XL model)
|
||||
- `xlnet`: :class:`~transformers.XLNetLMHeadModel` (XLNet model)
|
||||
- `ctrl`: :class:`~transformers.CTRLLMHeadModel` (Salesforce CTRL model)
|
||||
- `reformer`: :class:`~transformers.ReformerModelWithLMHead` (Google Reformer model)
|
||||
List options
|
||||
|
||||
The model is set in evaluation mode by default using `model.eval()` (Dropout modules are deactivated)
|
||||
To train the model, you should first set it back in training mode with `model.train()`
|
||||
@@ -1022,6 +941,7 @@ class AutoModelForMaskedLM:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(MODEL_FOR_MASKED_LM_MAPPING, use_model_types=False)
|
||||
def from_config(cls, config):
|
||||
r"""Instantiates one of the base model classes of the library
|
||||
from a configuration.
|
||||
@@ -1034,18 +954,8 @@ class AutoModelForMaskedLM:
|
||||
Args:
|
||||
config (:class:`~transformers.PretrainedConfig`):
|
||||
The model class to instantiate is selected based on the configuration class:
|
||||
- isInstance of `distilbert` configuration class: :class:`~transformers.DistilBertForMaskedLM` (DistilBERT model)
|
||||
- isInstance of `longformer` configuration class: :class:`~transformers.LongformerForMaskedLM` (Longformer model)
|
||||
- isInstance of `roberta` configuration class: :class:`~transformers.RobertaForMaskedLM` (RoBERTa model)
|
||||
- isInstance of `bert` configuration class: :class:`~transformers.BertForMaskedLM` (Bert model)
|
||||
- isInstance of `flaubert` configuration class: :class:`~transformers.FlaubertWithLMHeadModel` (Flaubert model)
|
||||
- isInstance of `xlm` configuration class: :class:`~transformers.XLMWithLMHeadModel` (XLM model)
|
||||
- isInstance of `xlm-roberta` configuration class: :class:`~transformers.XLMRobertaForMaskedLM` (XLM-Roberta model)
|
||||
- isInstance of `electra` configuration class: :class:`~transformers.ElectraForMaskedLM` (Electra model)
|
||||
- isInstance of `camembert` configuration class: :class:`~transformers.CamembertForMaskedLM` (Camembert model)
|
||||
- isInstance of `albert` configuration class: :class:`~transformers.AlbertForMaskedLM` (Albert model)
|
||||
- isInstance of `funnel` configuration class: :class:`~transformers.FunnelForMaskedLM` (Funnel Transformer model)
|
||||
|
||||
List options
|
||||
|
||||
Examples::
|
||||
|
||||
@@ -1063,6 +973,7 @@ class AutoModelForMaskedLM:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(MODEL_FOR_MASKED_LM_MAPPING)
|
||||
def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
|
||||
r"""Instantiates one of the language modeling model classes of the library
|
||||
from a pre-trained model configuration.
|
||||
@@ -1071,17 +982,7 @@ class AutoModelForMaskedLM:
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
- `distilbert`: :class:`~transformers.DistilBertForMaskedLM` (DistilBERT model)
|
||||
- `albert`: :class:`~transformers.AlbertForMaskedLM` (ALBERT model)
|
||||
- `camembert`: :class:`~transformers.CamembertForMaskedLM` (CamemBERT model)
|
||||
- `xlm-roberta`: :class:`~transformers.XLMRobertaForMaskedLM` (XLM-RoBERTa model)
|
||||
- `longformer`: :class:`~transformers.LongformerForMaskedLM` (Longformer model)
|
||||
- `roberta`: :class:`~transformers.RobertaForMaskedLM` (RoBERTa model)
|
||||
- `xlm`: :class:`~transformers.XLMWithLMHeadModel` (XLM model)
|
||||
- `flaubert`: :class:`~transformers.FlaubertWithLMHeadModel` (Flaubert model)
|
||||
- `electra`: :class:`~transformers.ElectraForMaskedLM` (Electra model)
|
||||
- `bert`: :class:`~transformers.BertLMHeadModel` (Bert model)
|
||||
- `funnel`: :class:`~transformers.FunnelForMaskedLM` (Funnel Transformer model)
|
||||
List options
|
||||
|
||||
The model is set in evaluation mode by default using `model.eval()` (Dropout modules are deactivated)
|
||||
To train the model, you should first set it back in training mode with `model.train()`
|
||||
@@ -1168,6 +1069,7 @@ class AutoModelForSeq2SeqLM:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(MODEL_FOR_SEQ_TO_SEQ_CAUSAL_LM_MAPPING, use_model_types=False)
|
||||
def from_config(cls, config):
|
||||
r"""Instantiates one of the base model classes of the library
|
||||
from a configuration.
|
||||
@@ -1181,10 +1083,7 @@ class AutoModelForSeq2SeqLM:
|
||||
config (:class:`~transformers.PretrainedConfig`):
|
||||
The model class to instantiate is selected based on the configuration class:
|
||||
|
||||
- isInstance of `t5` configuration class: :class:`~transformers.T5ForConditionalGeneration` (T5 model)
|
||||
- isInstance of `bart` configuration class: :class:`~transformers.BartForConditionalGeneration` (Bart model)
|
||||
- isInstance of `marian` configuration class: :class:`~transformers.MarianMTModel` (Marian model)
|
||||
- isInstance of `encoder-decoder` configuration class: :class:`~transformers.EncoderDecoderModel` (Encoder Decoder model)
|
||||
List options
|
||||
|
||||
Examples::
|
||||
|
||||
@@ -1204,6 +1103,7 @@ class AutoModelForSeq2SeqLM:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(MODEL_FOR_SEQ_TO_SEQ_CAUSAL_LM_MAPPING)
|
||||
def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
|
||||
r"""Instantiates one of the language modeling model classes of the library
|
||||
from a pre-trained model configuration.
|
||||
@@ -1212,10 +1112,7 @@ class AutoModelForSeq2SeqLM:
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
- `t5`: :class:`~transformers.T5ForConditionalGeneration` (T5 model)
|
||||
- `bart`: :class:`~transformers.BartForConditionalGeneration` (Bert model)
|
||||
- `marian`: :class:`~transformers.MarianMTModel` (Marian model)
|
||||
- `encoder-decoder`: :class:`~transformers.EncoderDecoderModel` (Encoder Decoder model)
|
||||
List options
|
||||
|
||||
The model is set in evaluation mode by default using `model.eval()` (Dropout modules are deactivated)
|
||||
To train the model, you should first set it back in training mode with `model.train()`
|
||||
@@ -1304,6 +1201,7 @@ class AutoModelForSequenceClassification:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(MODEL_FOR_SEQUENCE_CLASSIFICATION_MAPPING, use_model_types=False)
|
||||
def from_config(cls, config):
|
||||
r"""Instantiates one of the base model classes of the library
|
||||
from a configuration.
|
||||
@@ -1317,16 +1215,7 @@ class AutoModelForSequenceClassification:
|
||||
config (:class:`~transformers.PretrainedConfig`):
|
||||
The model class to instantiate is selected based on the configuration class:
|
||||
|
||||
- isInstance of `distilbert` configuration class: :class:`~transformers.DistilBertForSequenceClassification` (DistilBERT model)
|
||||
- isInstance of `albert` configuration class: :class:`~transformers.AlbertForSequenceClassification` (ALBERT model)
|
||||
- isInstance of `camembert` configuration class: :class:`~transformers.CamembertForSequenceClassification` (CamemBERT model)
|
||||
- isInstance of `xlm roberta` configuration class: :class:`~transformers.XLMRobertaForSequenceClassification` (XLM-RoBERTa model)
|
||||
- isInstance of `roberta` configuration class: :class:`~transformers.RobertaForSequenceClassification` (RoBERTa model)
|
||||
- isInstance of `bert` configuration class: :class:`~transformers.BertForSequenceClassification` (Bert model)
|
||||
- isInstance of `xlnet` configuration class: :class:`~transformers.XLNetForSequenceClassification` (XLNet model)
|
||||
- isInstance of `xlm` configuration class: :class:`~transformers.XLMForSequenceClassification` (XLM model)
|
||||
- isInstance of `flaubert` configuration class: :class:`~transformers.FlaubertForSequenceClassification` (Flaubert model)
|
||||
- isInstance of `funnel` configuration class: :class:`~transformers.FunnelModelForSequenceClassification` (Funnel Transformer model)
|
||||
List options
|
||||
|
||||
Examples::
|
||||
|
||||
@@ -1346,6 +1235,7 @@ class AutoModelForSequenceClassification:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(MODEL_FOR_SEQUENCE_CLASSIFICATION_MAPPING)
|
||||
def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
|
||||
r"""Instantiates one of the sequence classification model classes of the library
|
||||
from a pre-trained model configuration.
|
||||
@@ -1354,15 +1244,7 @@ class AutoModelForSequenceClassification:
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
- `distilbert`: :class:`~transformers.DistilBertForSequenceClassification` (DistilBERT model)
|
||||
- `albert`: :class:`~transformers.AlbertForSequenceClassification` (ALBERT model)
|
||||
- `camembert`: :class:`~transformers.CamembertForSequenceClassification` (CamemBERT model)
|
||||
- `xlm-roberta`: :class:`~transformers.XLMRobertaForSequenceClassification` (XLM-RoBERTa model)
|
||||
- `roberta`: :class:`~transformers.RobertaForSequenceClassification` (RoBERTa model)
|
||||
- `bert`: :class:`~transformers.BertForSequenceClassification` (Bert model)
|
||||
- `xlnet`: :class:`~transformers.XLNetForSequenceClassification` (XLNet model)
|
||||
- `flaubert`: :class:`~transformers.FlaubertForSequenceClassification` (Flaubert model)
|
||||
- `funnel`: :class:`~transformers.FunnelForSequenceClassification` (Funnel Transformer model)
|
||||
List options
|
||||
|
||||
The model is set in evaluation mode by default using `model.eval()` (Dropout modules are deactivated)
|
||||
To train the model, you should first set it back in training mode with `model.train()`
|
||||
@@ -1458,6 +1340,7 @@ class AutoModelForQuestionAnswering:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(MODEL_FOR_QUESTION_ANSWERING_MAPPING, use_model_types=False)
|
||||
def from_config(cls, config):
|
||||
r"""Instantiates one of the base model classes of the library
|
||||
from a configuration.
|
||||
@@ -1471,13 +1354,7 @@ class AutoModelForQuestionAnswering:
|
||||
config (:class:`~transformers.PretrainedConfig`):
|
||||
The model class to instantiate is selected based on the configuration class:
|
||||
|
||||
- isInstance of `distilbert` configuration class: :class:`~transformers.DistilBertForQuestionAnswering` (DistilBERT model)
|
||||
- isInstance of `albert` configuration class: :class:`~transformers.AlbertForQuestionAnswering` (ALBERT model)
|
||||
- isInstance of `bert` configuration class: :class:`~transformers.BertModelForQuestionAnswering` (Bert model)
|
||||
- isInstance of `xlnet` configuration class: :class:`~transformers.XLNetForQuestionAnswering` (XLNet model)
|
||||
- isInstance of `xlm` configuration class: :class:`~transformers.XLMForQuestionAnswering` (XLM model)
|
||||
- isInstance of `flaubert` configuration class: :class:`~transformers.FlaubertForQuestionAnswering` (XLM model)
|
||||
- isInstance of `funnel` configuration class: :class:`~transformers.FunnelForQuestionAnswering` (Funnel Transformer model)
|
||||
List options
|
||||
|
||||
Examples::
|
||||
|
||||
@@ -1498,6 +1375,7 @@ class AutoModelForQuestionAnswering:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(MODEL_FOR_QUESTION_ANSWERING_MAPPING)
|
||||
def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
|
||||
r"""Instantiates one of the question answering model classes of the library
|
||||
from a pre-trained model configuration.
|
||||
@@ -1506,13 +1384,7 @@ class AutoModelForQuestionAnswering:
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
- `distilbert`: :class:`~transformers.DistilBertForQuestionAnswering` (DistilBERT model)
|
||||
- `albert`: :class:`~transformers.AlbertForQuestionAnswering` (ALBERT model)
|
||||
- `bert`: :class:`~transformers.BertForQuestionAnswering` (Bert model)
|
||||
- `xlnet`: :class:`~transformers.XLNetForQuestionAnswering` (XLNet model)
|
||||
- `xlm`: :class:`~transformers.XLMForQuestionAnswering` (XLM model)
|
||||
- `flaubert`: :class:`~transformers.FlaubertForQuestionAnswering` (XLM model)
|
||||
- `funnel`: :class:`~transformers.FunnelForQuestionAnswering` (Funnel Transformer model)
|
||||
List options
|
||||
|
||||
The model is set in evaluation mode by default using `model.eval()` (Dropout modules are deactivated)
|
||||
To train the model, you should first set it back in training mode with `model.train()`
|
||||
@@ -1606,6 +1478,7 @@ class AutoModelForTokenClassification:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(MODEL_FOR_TOKEN_CLASSIFICATION_MAPPING, use_model_types=False)
|
||||
def from_config(cls, config):
|
||||
r"""Instantiates one of the base model classes of the library
|
||||
from a configuration.
|
||||
@@ -1619,17 +1492,7 @@ class AutoModelForTokenClassification:
|
||||
config (:class:`~transformers.PretrainedConfig`):
|
||||
The model class to instantiate is selected based on the configuration class:
|
||||
|
||||
- isInstance of `distilbert` configuration class: :class:`~transformers.DistilBertModelForTokenClassification` (DistilBERT model)
|
||||
- isInstance of `xlm` configuration class: :class:`~transformers.XLMForTokenClassification` (XLM model)
|
||||
- isInstance of `xlm roberta` configuration class: :class:`~transformers.XLMRobertaModelForTokenClassification` (XLMRoberta model)
|
||||
- isInstance of `bert` configuration class: :class:`~transformers.BertModelForTokenClassification` (Bert model)
|
||||
- isInstance of `albert` configuration class: :class:`~transformers.AlbertForTokenClassification` (AlBert model)
|
||||
- isInstance of `xlnet` configuration class: :class:`~transformers.XLNetModelForTokenClassification` (XLNet model)
|
||||
- isInstance of `flaubert` configuration class: :class:`~transformers.FlaubertForTokenClassification` (Flaubert model)
|
||||
- isInstance of `camembert` configuration class: :class:`~transformers.CamembertModelForTokenClassification` (Camembert model)
|
||||
- isInstance of `roberta` configuration class: :class:`~transformers.RobertaModelForTokenClassification` (Roberta model)
|
||||
- isInstance of `electra` configuration class: :class:`~transformers.ElectraForTokenClassification` (Electra model)
|
||||
- isInstance of `funnel` configuration class: :class:`~transformers.FunnelForTokenClassification` (Funnel Transformer model)
|
||||
List options
|
||||
|
||||
Examples::
|
||||
|
||||
@@ -1650,6 +1513,7 @@ class AutoModelForTokenClassification:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(MODEL_FOR_TOKEN_CLASSIFICATION_MAPPING)
|
||||
def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
|
||||
r"""Instantiates one of the question answering model classes of the library
|
||||
from a pre-trained model configuration.
|
||||
@@ -1658,16 +1522,7 @@ class AutoModelForTokenClassification:
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
- `distilbert`: :class:`~transformers.DistilBertForTokenClassification` (DistilBERT model)
|
||||
- `xlm`: :class:`~transformers.XLMForTokenClassification` (XLM model)
|
||||
- `xlm-roberta`: :class:`~transformers.XLMRobertaForTokenClassification` (XLM-RoBERTa?Para model)
|
||||
- `camembert`: :class:`~transformers.CamembertForTokenClassification` (Camembert model)
|
||||
- `bert`: :class:`~transformers.BertForTokenClassification` (Bert model)
|
||||
- `xlnet`: :class:`~transformers.XLNetForTokenClassification` (XLNet model)
|
||||
- `flaubert`: :class:`~transformers.FlaubertForTokenClassification` (Flaubert model)
|
||||
- `roberta`: :class:`~transformers.RobertaForTokenClassification` (Roberta model)
|
||||
- `electra`: :class:`~transformers.ElectraForTokenClassification` (Electra model)
|
||||
- `funnel`: :class:`~transformers.FunnelForTokenClassification` (Funnel Transformer model)
|
||||
List options
|
||||
|
||||
The model is set in evaluation mode by default using `model.eval()` (Dropout modules are deactivated)
|
||||
To train the model, you should first set it back in training mode with `model.train()`
|
||||
@@ -1761,7 +1616,27 @@ class AutoModelForMultipleChoice:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(MODEL_FOR_MULTIPLE_CHOICE_MAPPING, use_model_types=False)
|
||||
def from_config(cls, config):
|
||||
r"""Instantiates one of the base model classes of the library
|
||||
from a configuration.
|
||||
|
||||
Note:
|
||||
Loading a model from its configuration file does **not** load the model weights.
|
||||
It only affects the model's configuration. Use :func:`~transformers.AutoModel.from_pretrained` to load
|
||||
the model weights
|
||||
|
||||
Args:
|
||||
config (:class:`~transformers.PretrainedConfig`):
|
||||
The model class to instantiate is selected based on the configuration class:
|
||||
|
||||
List options
|
||||
|
||||
Examples::
|
||||
|
||||
config = BertConfig.from_pretrained('bert-base-uncased') # Download configuration from S3 and cache.
|
||||
model = AutoModelForMultipleChoice.from_config(config) # E.g. model was saved using `save_pretrained('./test/saved_model/')`
|
||||
"""
|
||||
for config_class, model_class in MODEL_FOR_MULTIPLE_CHOICE_MAPPING.items():
|
||||
if isinstance(config, config_class):
|
||||
return model_class(config)
|
||||
@@ -1776,7 +1651,71 @@ class AutoModelForMultipleChoice:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(MODEL_FOR_MULTIPLE_CHOICE_MAPPING)
|
||||
def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
|
||||
r"""Instantiates one of the question answering model classes of the library
|
||||
from a pre-trained model configuration.
|
||||
|
||||
The `from_pretrained()` method takes care of returning the correct model class instance
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
List options
|
||||
|
||||
The model is set in evaluation mode by default using `model.eval()` (Dropout modules are deactivated)
|
||||
To train the model, you should first set it back in training mode with `model.train()`
|
||||
|
||||
Args:
|
||||
pretrained_model_name_or_path:
|
||||
Either:
|
||||
|
||||
- a string with the `shortcut name` of a pre-trained model to load from cache or download, e.g.: ``bert-base-uncased``.
|
||||
- a path to a `directory` containing model weights saved using :func:`~transformers.PreTrainedModel.save_pretrained`, e.g.: ``./my_model_directory/``.
|
||||
- a path or url to a `tensorflow index checkpoint file` (e.g. `./tf_model/model.ckpt.index`). In this case, ``from_tf`` should be set to True and a configuration object should be provided as ``config`` argument. This loading path is slower than converting the TensorFlow checkpoint in a PyTorch model using the provided conversion scripts and loading the PyTorch model afterwards.
|
||||
|
||||
model_args: (`optional`) Sequence of positional arguments:
|
||||
All remaning positional arguments will be passed to the underlying model's ``__init__`` method
|
||||
|
||||
config: (`optional`) instance of a class derived from :class:`~transformers.PretrainedConfig`:
|
||||
Configuration for the model to use instead of an automatically loaded configuation. Configuration can be automatically loaded when:
|
||||
|
||||
- the model is a model provided by the library (loaded with the ``shortcut-name`` string of a pretrained model), or
|
||||
- the model was saved using :func:`~transformers.PreTrainedModel.save_pretrained` and is reloaded by suppling the save directory.
|
||||
- the model is loaded by suppling a local directory as ``pretrained_model_name_or_path`` and a configuration JSON file named `config.json` is found in the directory.
|
||||
|
||||
state_dict: (`optional`) dict:
|
||||
an optional state dictionary for the model to use instead of a state dictionary loaded from saved weights file.
|
||||
This option can be used if you want to create a model from a pretrained configuration but load your own weights.
|
||||
In this case though, you should check if using :func:`~transformers.PreTrainedModel.save_pretrained` and :func:`~transformers.PreTrainedModel.from_pretrained` is not a simpler option.
|
||||
|
||||
cache_dir: (`optional`) string:
|
||||
Path to a directory in which a downloaded pre-trained model
|
||||
configuration should be cached if the standard cache should not be used.
|
||||
|
||||
force_download: (`optional`) boolean, default False:
|
||||
Force to (re-)download the model weights and configuration files and override the cached versions if they exists.
|
||||
|
||||
proxies: (`optional`) dict, default None:
|
||||
A dictionary of proxy servers to use by protocol or endpoint, e.g.: {'http': 'foo.bar:3128', 'http://hostname': 'foo.bar:4012'}.
|
||||
The proxies are used on each request.
|
||||
|
||||
output_loading_info: (`optional`) boolean:
|
||||
Set to ``True`` to also return a dictionary containing missing keys, unexpected keys and error messages.
|
||||
|
||||
kwargs: (`optional`) Remaining dictionary of keyword arguments:
|
||||
These arguments will be passed to the configuration and the model.
|
||||
|
||||
Examples::
|
||||
|
||||
model = AutoModelForForMultipleChoice.from_pretrained('bert-base-uncased') # Download model and configuration from S3 and cache.
|
||||
model = AutoModelForMultipleChoice.from_pretrained('./test/bert_model/') # E.g. model was saved using `save_pretrained('./test/saved_model/')`
|
||||
model = AutoModelForMultipleChoice.from_pretrained('bert-base-uncased', output_attentions=True) # Update configuration during loading
|
||||
assert model.config.output_attentions == True
|
||||
# Loading from a TF checkpoint file instead of a PyTorch model (slower)
|
||||
config = AutoConfig.from_json_file('./tf_model/bert_tf_model_config.json')
|
||||
model = AutoModelForMultipleChoice.from_pretrained('./tf_model/bert_tf_checkpoint.ckpt.index', from_tf=True, config=config)
|
||||
|
||||
"""
|
||||
config = kwargs.pop("config", None)
|
||||
if not isinstance(config, PretrainedConfig):
|
||||
config, kwargs = AutoConfig.from_pretrained(
|
||||
|
||||
+508
-616
File diff suppressed because it is too large.
Load diff
Executable
+505
@@ -0,0 +1,505 @@
|
||||
# coding=utf-8
|
||||
# Copyright 2020 The Google AI Language Team 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.
|
||||
"""PyTorch BERT model specific for generation. """
|
||||
|
||||
|
||||
import torch
|
||||
import torch.utils.checkpoint
|
||||
from torch import nn
|
||||
from torch.nn import CrossEntropyLoss
|
||||
|
||||
from .configuration_bert_generation import BertGenerationConfig
|
||||
from .file_utils import (
|
||||
add_code_sample_docstrings,
|
||||
add_start_docstrings,
|
||||
add_start_docstrings_to_callable,
|
||||
replace_return_docstrings,
|
||||
)
|
||||
from .modeling_bert import BertEncoder
|
||||
from .modeling_outputs import BaseModelOutput, CausalLMOutput
|
||||
from .modeling_utils import PreTrainedModel
|
||||
from .utils import logging
|
||||
|
||||
|
||||
logger = logging.get_logger(__name__)
|
||||
|
||||
_CONFIG_FOR_DOC = "BertGenerationConfig"
|
||||
_TOKENIZER_FOR_DOC = "BertGenerationTokenizer"
|
||||
|
||||
|
||||
def load_tf_weights_in_bert_generation(
|
||||
model, tf_hub_path, model_class, is_encoder_named_decoder=False, is_encoder=False
|
||||
):
|
||||
try:
|
||||
import numpy as np
|
||||
import tensorflow.compat.v1 as tf
|
||||
|
||||
import tensorflow_hub as hub
|
||||
import tensorflow_text # noqa: F401
|
||||
|
||||
tf.disable_eager_execution()
|
||||
except ImportError:
|
||||
logger.error(
|
||||
"Loading a TensorFlow model in PyTorch, requires TensorFlow to be installed. Please see "
|
||||
"https://www.tensorflow.org/install/ for installation instructions."
|
||||
)
|
||||
raise
|
||||
tf_model = hub.Module(tf_hub_path)
|
||||
init = tf.global_variables_initializer()
|
||||
with tf.Session() as sess:
|
||||
init.run()
|
||||
all_variables = tf_model.variable_map
|
||||
keep_track_variables = all_variables.copy()
|
||||
for key in list(all_variables.keys()):
|
||||
if "global" in key:
|
||||
logger.info(f"Skipping {key}...")
|
||||
continue
|
||||
if not is_encoder:
|
||||
model_pointer = getattr(model, model_class)
|
||||
else:
|
||||
model_pointer = model
|
||||
is_embedding = False
|
||||
logger.info(f"Trying to match {key}...")
|
||||
# remove start_string = "module/bert/"
|
||||
sub_layers = key.split("/")[2:]
|
||||
if is_encoder_named_decoder and sub_layers[0] == "encoder":
|
||||
logger.info(f"Skipping encoder layer {key} for decoder")
|
||||
continue
|
||||
if is_encoder and sub_layers[0] == "decoder":
|
||||
logger.info(f"Skipping decoder layer {key} for encoder")
|
||||
continue
|
||||
for i, sub_layer in enumerate(sub_layers):
|
||||
if sub_layer == "embeddings":
|
||||
is_embedding = True
|
||||
elif sub_layer == "LayerNorm":
|
||||
is_embedding = False
|
||||
if "layer" in sub_layer:
|
||||
model_pointer = model_pointer.layer[int(sub_layer.split("_")[-1])]
|
||||
elif sub_layer in ["kernel", "gamma"]:
|
||||
model_pointer = model_pointer.weight
|
||||
elif sub_layer == "beta":
|
||||
model_pointer = model_pointer.bias
|
||||
elif sub_layer == "encdec":
|
||||
model_pointer = model_pointer.crossattention.self
|
||||
elif sub_layer == "encdec_output":
|
||||
model_pointer = model_pointer.crossattention.output
|
||||
elif is_encoder_named_decoder and sub_layer == "decoder":
|
||||
model_pointer = model_pointer.encoder
|
||||
else:
|
||||
if sub_layer == "attention" and "encdec" in sub_layers[i + 1]:
|
||||
continue
|
||||
try:
|
||||
model_pointer = getattr(model_pointer, sub_layer)
|
||||
except AttributeError:
|
||||
logger.info(f"Skipping to initialize {key} at {sub_layer}...")
|
||||
raise AttributeError
|
||||
|
||||
array = np.asarray(sess.run(all_variables[key]))
|
||||
if not is_embedding:
|
||||
logger.info("Transposing numpy weight of shape {} for {}".format(array.shape, key))
|
||||
array = np.transpose(array)
|
||||
else:
|
||||
model_pointer = model_pointer.weight
|
||||
|
||||
try:
|
||||
assert (
|
||||
model_pointer.shape == array.shape
|
||||
), f"Pointer shape {model_pointer.shape} and array shape {array.shape} mismatched"
|
||||
except AssertionError as e:
|
||||
e.args += (model_pointer.shape, array.shape)
|
||||
raise
|
||||
logger.info(f"Initialize PyTorch weight {key}")
|
||||
|
||||
model_pointer.data = torch.from_numpy(array.astype(np.float32))
|
||||
keep_track_variables.pop(key, None)
|
||||
|
||||
logger.info("Weights not copied to PyTorch model: {}".format(", ".join(keep_track_variables.keys())))
|
||||
return model
|
||||
|
||||
|
||||
class BertGenerationEmbeddings(nn.Module):
|
||||
"""Construct the embeddings from word, position and token_type embeddings."""
|
||||
|
||||
def __init__(self, config):
|
||||
super().__init__()
|
||||
self.word_embeddings = nn.Embedding(config.vocab_size, config.hidden_size, padding_idx=config.pad_token_id)
|
||||
self.position_embeddings = nn.Embedding(config.max_position_embeddings, config.hidden_size)
|
||||
# self.LayerNorm is not snake-cased to stick with TensorFlow model variable name and be able to load
|
||||
# any TensorFlow checkpoint file
|
||||
self.LayerNorm = torch.nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
|
||||
self.dropout = nn.Dropout(config.hidden_dropout_prob)
|
||||
|
||||
# position_ids (1, len position emb) is contiguous in memory and exported when serialized
|
||||
self.register_buffer("position_ids", torch.arange(config.max_position_embeddings).expand((1, -1)))
|
||||
|
||||
def forward(self, input_ids=None, position_ids=None, inputs_embeds=None):
|
||||
if input_ids is not None:
|
||||
input_shape = input_ids.size()
|
||||
else:
|
||||
input_shape = inputs_embeds.size()[:-1]
|
||||
|
||||
seq_length = input_shape[1]
|
||||
|
||||
if position_ids is None:
|
||||
position_ids = self.position_ids[:, :seq_length]
|
||||
|
||||
if inputs_embeds is None:
|
||||
inputs_embeds = self.word_embeddings(input_ids)
|
||||
position_embeddings = self.position_embeddings(position_ids)
|
||||
|
||||
embeddings = inputs_embeds + position_embeddings
|
||||
embeddings = self.LayerNorm(embeddings)
|
||||
embeddings = self.dropout(embeddings)
|
||||
return embeddings
|
||||
|
||||
|
||||
class BertGenerationPreTrainedModel(PreTrainedModel):
|
||||
"""An abstract class to handle weights initialization and
|
||||
a simple interface for downloading and loading pretrained models.
|
||||
"""
|
||||
|
||||
config_class = BertGenerationConfig
|
||||
base_model_prefix = "bert"
|
||||
authorized_missing_keys = [r"position_ids"]
|
||||
|
||||
def _init_weights(self, module):
|
||||
""" Initialize the weights """
|
||||
if isinstance(module, (nn.Linear, nn.Embedding)):
|
||||
# Slightly different from the TF version which uses truncated_normal for initialization
|
||||
# cf https://github.com/pytorch/pytorch/pull/5617
|
||||
module.weight.data.normal_(mean=0.0, std=self.config.initializer_range)
|
||||
elif isinstance(module, nn.LayerNorm):
|
||||
module.bias.data.zero_()
|
||||
module.weight.data.fill_(1.0)
|
||||
if isinstance(module, nn.Linear) and module.bias is not None:
|
||||
module.bias.data.zero_()
|
||||
|
||||
|
||||
BERT_GENERATION_START_DOCSTRING = r"""
|
||||
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.
|
||||
|
||||
Parameters:
|
||||
config (:class:`~transformers.BertGenerationConfig`): 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.
|
||||
"""
|
||||
|
||||
BERT_GENERATION_INPUTS_DOCSTRING = r"""
|
||||
Args:
|
||||
input_ids (:obj:`torch.LongTensor` of shape :obj:`{0}`):
|
||||
Indices of input sequence tokens in the vocabulary.
|
||||
|
||||
Indices can be obtained using :class:`transformers.BertGenerationTokenizer`.
|
||||
See :func:`transformers.PreTrainedTokenizer.encode` and
|
||||
:func:`transformers.PreTrainedTokenizer.__call__` for details.
|
||||
|
||||
`What are input IDs? <../glossary.html#input-ids>`__
|
||||
attention_mask (:obj:`torch.FloatTensor` of shape :obj:`{0}`, `optional`):
|
||||
Mask to avoid performing attention on padding token indices.
|
||||
Mask values selected in ``[0, 1]``:
|
||||
``1`` for tokens that are NOT MASKED, ``0`` for MASKED tokens.
|
||||
|
||||
`What are attention masks? <../glossary.html#attention-mask>`__
|
||||
position_ids (:obj:`torch.LongTensor` of shape :obj:`{0}`, `optional`):
|
||||
Indices of positions of each input sequence tokens in the position embeddings.
|
||||
Selected in the range ``[0, config.max_position_embeddings - 1]``.
|
||||
|
||||
`What are position IDs? <../glossary.html#position-ids>`_
|
||||
head_mask (:obj:`torch.FloatTensor` of shape :obj:`(num_heads,)` or :obj:`(num_layers, num_heads)`, `optional`):
|
||||
Mask to nullify selected heads of the self-attention modules.
|
||||
Mask values selected in ``[0, 1]``:
|
||||
:obj:`1` indicates the head is **not masked**, :obj:`0` indicates the head is **masked**.
|
||||
inputs_embeds (:obj:`torch.FloatTensor` of shape :obj:`(batch_size, sequence_length, hidden_size)`, `optional`):
|
||||
Optionally, instead of passing :obj:`input_ids` you can choose to directly pass an embedded representation.
|
||||
This is useful if you want more control over how to convert `input_ids` indices into associated vectors
|
||||
than the model's internal embedding lookup matrix.
|
||||
output_attentions (:obj:`bool`, `optional`):
|
||||
If set to ``True``, the attentions tensors of all attention layers are returned. See ``attentions`` under returned tensors for more detail.
|
||||
output_hidden_states (:obj:`bool`, `optional`):
|
||||
If set to ``True``, the hidden states of all layers are returned. See ``hidden_states`` under returned tensors for more detail.
|
||||
return_dict (:obj:`bool`, `optional`):
|
||||
If set to ``True``, the model will return a :class:`~transformers.file_utils.ModelOutput` instead of a
|
||||
plain tuple.
|
||||
"""
|
||||
|
||||
|
||||
@add_start_docstrings(
|
||||
"The bare BertGeneration model transformer outputting raw hidden-states without any specific head on top.",
|
||||
BERT_GENERATION_START_DOCSTRING,
|
||||
)
|
||||
class BertGenerationEncoder(BertGenerationPreTrainedModel):
|
||||
"""
|
||||
|
||||
The model can behave as an encoder (with only self-attention) as well
|
||||
as a decoder, in which case a layer of cross-attention is added between
|
||||
the self-attention layers, following the architecture described in `Attention is all you need <https://arxiv.org/abs/1706.03762>`__ by Ashish Vaswani,
|
||||
Noam Shazeer, Niki Parmar, Jakob Uszkoreit, Llion Jones, Aidan N. Gomez, Lukasz Kaiser and Illia Polosukhin.
|
||||
|
||||
This model should be used when leveraging Bert or Roberta checkpoints for the `EncoderDecoderModel` class as described in `Leveraging Pre-trained Checkpoints for Sequence Generation Tasks <https://arxiv.org/abs/1907.12461>`__ by Sascha Rothe, Shashi Narayan, and Aliaksei Severyn.
|
||||
|
||||
To behave as an decoder the model needs to be initialized with the
|
||||
:obj:`is_decoder` argument of the configuration set to :obj:`True`.
|
||||
To be used in a Seq2Seq model, the model needs to initialized with both :obj:`is_decoder`
|
||||
argument and :obj:`add_cross_attention` set to :obj:`True`; an
|
||||
:obj:`encoder_hidden_states` is then expected as an input to the forward pass.
|
||||
"""
|
||||
|
||||
def __init__(self, config):
|
||||
super().__init__(config)
|
||||
self.config = config
|
||||
|
||||
self.embeddings = BertGenerationEmbeddings(config)
|
||||
self.encoder = BertEncoder(config)
|
||||
|
||||
self.init_weights()
|
||||
|
||||
def get_input_embeddings(self):
|
||||
return self.embeddings.word_embeddings
|
||||
|
||||
def set_input_embeddings(self, value):
|
||||
self.embeddings.word_embeddings = value
|
||||
|
||||
def _prune_heads(self, heads_to_prune):
|
||||
"""Prunes heads of the model.
|
||||
heads_to_prune: dict of {layer_num: list of heads to prune in this layer}
|
||||
See base class PreTrainedModel
|
||||
"""
|
||||
for layer, heads in heads_to_prune.items():
|
||||
self.encoder.layer[layer].attention.prune_heads(heads)
|
||||
|
||||
@add_start_docstrings_to_callable(BERT_GENERATION_INPUTS_DOCSTRING.format("(batch_size, sequence_length)"))
|
||||
@add_code_sample_docstrings(
|
||||
tokenizer_class=_TOKENIZER_FOR_DOC,
|
||||
checkpoint="google/bert_for_seq_generation_L-24_bbc_encoder",
|
||||
output_type=BaseModelOutput,
|
||||
config_class=_CONFIG_FOR_DOC,
|
||||
)
|
||||
def forward(
|
||||
self,
|
||||
input_ids=None,
|
||||
attention_mask=None,
|
||||
position_ids=None,
|
||||
head_mask=None,
|
||||
inputs_embeds=None,
|
||||
encoder_hidden_states=None,
|
||||
encoder_attention_mask=None,
|
||||
output_attentions=None,
|
||||
output_hidden_states=None,
|
||||
return_dict=None,
|
||||
):
|
||||
r"""
|
||||
encoder_hidden_states (: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. Used in the cross-attention
|
||||
if the model is configured as a decoder.
|
||||
encoder_attention_mask (:obj:`torch.FloatTensor` of shape :obj:`(batch_size, sequence_length)`, `optional`):
|
||||
Mask to avoid performing attention on the padding token indices of the encoder input. This mask
|
||||
is used in the cross-attention if the model is configured as a decoder.
|
||||
Mask values selected in ``[0, 1]``:
|
||||
``1`` for tokens that are NOT MASKED, ``0`` for MASKED tokens.
|
||||
"""
|
||||
output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
|
||||
output_hidden_states = (
|
||||
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
|
||||
)
|
||||
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
|
||||
|
||||
if input_ids is not None and inputs_embeds is not None:
|
||||
raise ValueError("You cannot specify both input_ids and inputs_embeds at the same time")
|
||||
elif input_ids is not None:
|
||||
input_shape = input_ids.size()
|
||||
elif inputs_embeds is not None:
|
||||
input_shape = inputs_embeds.size()[:-1]
|
||||
else:
|
||||
raise ValueError("You have to specify either input_ids or inputs_embeds")
|
||||
|
||||
device = input_ids.device if input_ids is not None else inputs_embeds.device
|
||||
|
||||
if attention_mask is None:
|
||||
attention_mask = torch.ones(input_shape, device=device)
|
||||
|
||||
# We can provide a self-attention mask of dimensions [batch_size, from_seq_length, to_seq_length]
|
||||
# ourselves in which case we just need to make it broadcastable to all heads.
|
||||
extended_attention_mask: torch.Tensor = self.get_extended_attention_mask(attention_mask, input_shape, device)
|
||||
|
||||
# If a 2D or 3D attention mask is provided for the cross-attention
|
||||
# we need to make broadcastable to [batch_size, num_heads, seq_length, seq_length]
|
||||
if self.config.is_decoder and encoder_hidden_states is not None:
|
||||
encoder_batch_size, encoder_sequence_length, _ = encoder_hidden_states.size()
|
||||
encoder_hidden_shape = (encoder_batch_size, encoder_sequence_length)
|
||||
if encoder_attention_mask is None:
|
||||
encoder_attention_mask = torch.ones(encoder_hidden_shape, device=device)
|
||||
encoder_extended_attention_mask = self.invert_attention_mask(encoder_attention_mask)
|
||||
else:
|
||||
encoder_extended_attention_mask = None
|
||||
|
||||
# Prepare head mask if needed
|
||||
# 1.0 in head_mask indicate we keep the head
|
||||
# attention_probs has shape bsz x n_heads x N x N
|
||||
# input head_mask has shape [num_heads] or [num_hidden_layers x num_heads]
|
||||
# and head_mask is converted to shape [num_hidden_layers x batch x num_heads x seq_length x seq_length]
|
||||
head_mask = self.get_head_mask(head_mask, self.config.num_hidden_layers)
|
||||
|
||||
embedding_output = self.embeddings(input_ids=input_ids, position_ids=position_ids, inputs_embeds=inputs_embeds)
|
||||
|
||||
encoder_outputs = self.encoder(
|
||||
embedding_output,
|
||||
attention_mask=extended_attention_mask,
|
||||
head_mask=head_mask,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
encoder_attention_mask=encoder_extended_attention_mask,
|
||||
output_attentions=output_attentions,
|
||||
output_hidden_states=output_hidden_states,
|
||||
return_dict=return_dict,
|
||||
)
|
||||
sequence_output = encoder_outputs[0]
|
||||
|
||||
if not return_dict:
|
||||
return (sequence_output,) + encoder_outputs[1:]
|
||||
|
||||
return BaseModelOutput(
|
||||
last_hidden_state=sequence_output,
|
||||
hidden_states=encoder_outputs.hidden_states,
|
||||
attentions=encoder_outputs.attentions,
|
||||
)
|
||||
|
||||
|
||||
class BertGenerationOnlyLMHead(nn.Module):
|
||||
def __init__(self, config):
|
||||
super().__init__()
|
||||
self.decoder = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
|
||||
self.bias = nn.Parameter(torch.zeros(config.vocab_size))
|
||||
|
||||
# Need a link between the two variables so that the bias is correctly resized with `resize_token_embeddings`
|
||||
self.decoder.bias = self.bias
|
||||
|
||||
def forward(self, hidden_states):
|
||||
logits = self.decoder(hidden_states)
|
||||
return logits
|
||||
|
||||
|
||||
@add_start_docstrings(
|
||||
"""BertGeneration Model with a `language modeling` head on top for CLM fine-tuning. """,
|
||||
BERT_GENERATION_START_DOCSTRING,
|
||||
)
|
||||
class BertGenerationDecoder(BertGenerationPreTrainedModel):
|
||||
def __init__(self, config):
|
||||
super().__init__(config)
|
||||
|
||||
if not config.is_decoder:
|
||||
logger.warn("If you want to use `BertGenerationDecoder` as a standalone, add `is_decoder=True.`")
|
||||
|
||||
self.bert = BertGenerationEncoder(config)
|
||||
self.lm_head = BertGenerationOnlyLMHead(config)
|
||||
|
||||
self.init_weights()
|
||||
|
||||
def get_output_embeddings(self):
|
||||
return self.lm_head.decoder
|
||||
|
||||
@add_start_docstrings_to_callable(BERT_GENERATION_INPUTS_DOCSTRING.format("(batch_size, sequence_length)"))
|
||||
@replace_return_docstrings(output_type=CausalLMOutput, config_class=_CONFIG_FOR_DOC)
|
||||
def forward(
|
||||
self,
|
||||
input_ids=None,
|
||||
attention_mask=None,
|
||||
position_ids=None,
|
||||
head_mask=None,
|
||||
inputs_embeds=None,
|
||||
encoder_hidden_states=None,
|
||||
encoder_attention_mask=None,
|
||||
labels=None,
|
||||
output_attentions=None,
|
||||
output_hidden_states=None,
|
||||
return_dict=None,
|
||||
):
|
||||
r"""
|
||||
encoder_hidden_states (: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. Used in the cross-attention
|
||||
if the model is configured as a decoder.
|
||||
encoder_attention_mask (:obj:`torch.FloatTensor` of shape :obj:`(batch_size, sequence_length)`, `optional`):
|
||||
Mask to avoid performing attention on the padding token indices of the encoder input. This mask
|
||||
is used in the cross-attention if the model is configured as a decoder.
|
||||
Mask values selected in ``[0, 1]``:
|
||||
``1`` for tokens that are NOT MASKED, ``0`` for MASKED tokens.
|
||||
labels (:obj:`torch.LongTensor` of shape :obj:`(batch_size, sequence_length)`, `optional`):
|
||||
Labels for computing the left-to-right language modeling loss (next word prediction).
|
||||
Indices should be in ``[-100, 0, ..., config.vocab_size]`` (see ``input_ids`` docstring)
|
||||
Tokens with indices set to ``-100`` are ignored (masked), the loss is only computed for the tokens with labels
|
||||
in ``[0, ..., config.vocab_size]``
|
||||
|
||||
Returns:
|
||||
|
||||
Example::
|
||||
|
||||
>>> from transformers import BertGenerationTokenizer, BertGenerationDecoder, BertGenerationConfig
|
||||
>>> import torch
|
||||
|
||||
>>> tokenizer = BertGenerationTokenizer.from_pretrained('google/bert_for_seq_generation_L-24_bbc_encoder')
|
||||
>>> config = BertGenerationConfig.from_pretrained("google/bert_for_seq_generation_L-24_bbc_encoder")
|
||||
>>> config.is_decoder = True
|
||||
>>> model = BertGenerationDecoder.from_pretrained('google/bert_for_seq_generation_L-24_bbc_encoder', config=config, return_dict=True)
|
||||
|
||||
>>> inputs = tokenizer("Hello, my dog is cute", return_tensors="pt")
|
||||
>>> outputs = model(**inputs)
|
||||
|
||||
>>> prediction_logits = outputs.logits
|
||||
"""
|
||||
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
|
||||
|
||||
outputs = self.bert(
|
||||
input_ids,
|
||||
attention_mask=attention_mask,
|
||||
position_ids=position_ids,
|
||||
head_mask=head_mask,
|
||||
inputs_embeds=inputs_embeds,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
encoder_attention_mask=encoder_attention_mask,
|
||||
output_attentions=output_attentions,
|
||||
output_hidden_states=output_hidden_states,
|
||||
return_dict=return_dict,
|
||||
)
|
||||
|
||||
sequence_output = outputs[0]
|
||||
prediction_scores = self.lm_head(sequence_output)
|
||||
|
||||
lm_loss = None
|
||||
if labels is not None:
|
||||
# we are doing next-token prediction; shift prediction scores and input ids by one
|
||||
shifted_prediction_scores = prediction_scores[:, :-1, :].contiguous()
|
||||
labels = labels[:, 1:].contiguous()
|
||||
loss_fct = CrossEntropyLoss()
|
||||
lm_loss = loss_fct(shifted_prediction_scores.view(-1, self.config.vocab_size), labels.view(-1))
|
||||
|
||||
if not return_dict:
|
||||
output = (prediction_scores,) + outputs[1:]
|
||||
return ((lm_loss,) + output) if lm_loss is not None else output
|
||||
|
||||
return CausalLMOutput(
|
||||
loss=lm_loss,
|
||||
logits=prediction_scores,
|
||||
hidden_states=outputs.hidden_states,
|
||||
attentions=outputs.attentions,
|
||||
)
|
||||
|
||||
def prepare_inputs_for_generation(self, input_ids, attention_mask=None, **model_kwargs):
|
||||
input_shape = input_ids.shape
|
||||
|
||||
# if model is used as a decoder in encoder-decoder model, the decoder attention mask is created on the fly
|
||||
if attention_mask is None:
|
||||
attention_mask = input_ids.new_ones(input_shape)
|
||||
|
||||
return {"input_ids": input_ids, "attention_mask": attention_mask}
|
||||
@@ -238,8 +238,21 @@ class EncoderDecoderModel(PreTrainedModel):
|
||||
), "If `model` is not defined as an argument, a `encoder_pretrained_model_name_or_path` has to be defined"
|
||||
from .modeling_auto import AutoModel
|
||||
|
||||
if "config" not in kwargs_encoder:
|
||||
from .configuration_auto import AutoConfig
|
||||
|
||||
encoder_config = AutoConfig.from_pretrained(encoder_pretrained_model_name_or_path)
|
||||
if encoder_config.is_decoder is True or encoder_config.add_cross_attention is True:
|
||||
|
||||
logger.info(
|
||||
f"Initializing {encoder_pretrained_model_name_or_path} as a encoder model from a decoder model. Cross-attention and casual mask are disabled."
|
||||
)
|
||||
encoder_config.is_decoder = False
|
||||
encoder_config.add_cross_attention = False
|
||||
|
||||
kwargs_encoder["config"] = encoder_config
|
||||
|
||||
encoder = AutoModel.from_pretrained(encoder_pretrained_model_name_or_path, *model_args, **kwargs_encoder)
|
||||
encoder.config.is_decoder = False
|
||||
|
||||
decoder = kwargs_decoder.pop("model", None)
|
||||
if decoder is None:
|
||||
|
||||
@@ -425,9 +425,9 @@ def _relative_shift_gather(positional_attn, context_len, shift):
|
||||
# max_rel_len = 2 * context_len + shift -1 is the numbers of possible relative positions i-j
|
||||
|
||||
# What's next is the same as doing the following gather, which might be clearer code but less efficient.
|
||||
# idxs = context_len + torch.arange(0, context_len).unsqueeze(0) - torch.arange(0, context_len).unsqueeze(1)
|
||||
# idxs = context_len + torch.arange(0, context_len).unsqueeze(0) - torch.arange(0, seq_len).unsqueeze(1)
|
||||
# # matrix of context_len + i-j
|
||||
# return positional_attn.gather(3, idxs.expand([bs, n_head, context_len, context_len]))
|
||||
# return positional_attn.gather(3, idxs.expand([batch_size, n_head, context_len, context_len]))
|
||||
|
||||
positional_attn = torch.reshape(positional_attn, [batch_size, n_head, max_rel_len, seq_len])
|
||||
positional_attn = positional_attn[:, :, shift:, :]
|
||||
@@ -526,9 +526,9 @@ class FunnelRelMultiheadAttention(nn.Module):
|
||||
token_type_attn *= cls_mask
|
||||
return token_type_attn
|
||||
|
||||
def forward(self, query, key, value, attention_inputs, head_mask=None, output_attentions=False):
|
||||
# q has shape batch_size x seq_len x d_model
|
||||
# k and v have shapes batch_size x context_len x d_model
|
||||
def forward(self, query, key, value, attention_inputs, output_attentions=False):
|
||||
# query has shape batch_size x seq_len x d_model
|
||||
# key and value have shapes batch_size x context_len x d_model
|
||||
position_embeds, token_type_mat, attention_mask, cls_mask = attention_inputs
|
||||
|
||||
batch_size, seq_len, _ = query.shape
|
||||
@@ -598,8 +598,8 @@ class FunnelLayer(nn.Module):
|
||||
self.attention = FunnelRelMultiheadAttention(config, block_index)
|
||||
self.ffn = FunnelPositionwiseFFN(config)
|
||||
|
||||
def forward(self, q, k, v, attention_inputs, output_attentions=False):
|
||||
attn = self.attention(q, k, v, attention_inputs, output_attentions=output_attentions)
|
||||
def forward(self, query, key, value, attention_inputs, output_attentions=False):
|
||||
attn = self.attention(query, key, value, attention_inputs, output_attentions=output_attentions)
|
||||
output = self.ffn(attn[0])
|
||||
return (output, attn[1]) if output_attentions else (output,)
|
||||
|
||||
@@ -792,7 +792,7 @@ class FunnelClassificationHead(nn.Module):
|
||||
|
||||
def forward(self, hidden):
|
||||
hidden = self.linear_hidden(hidden)
|
||||
hidden = F.tanh(hidden)
|
||||
hidden = torch.tanh(hidden)
|
||||
hidden = self.dropout(hidden)
|
||||
return self.linear_out(hidden)
|
||||
|
||||
@@ -954,7 +954,7 @@ class FunnelBaseModel(FunnelPreTrainedModel):
|
||||
|
||||
|
||||
@add_start_docstrings(
|
||||
"The bare base Funnel Transformer Model transformer outputting raw hidden-states without any specific head on top.",
|
||||
"The bare Funnel Transformer Model transformer outputting raw hidden-states without any specific head on top.",
|
||||
FUNNEL_START_DOCSTRING,
|
||||
)
|
||||
class FunnelModel(FunnelPreTrainedModel):
|
||||
@@ -1099,10 +1099,10 @@ class FunnelForPreTraining(FunnelPreTrainedModel):
|
||||
>>> import torch
|
||||
|
||||
>>> tokenizer = FunnelTokenizer.from_pretrained('funnel-transformer/small')
|
||||
>>> model = FunnelForPreTraining.from_pretrained('funnel-transformer/small')
|
||||
>>> model = FunnelForPreTraining.from_pretrained('funnel-transformer/small', return_dict=True)
|
||||
|
||||
>>> input_ids = torch.tensor(tokenizer.encode("Hello, my dog is cute", add_special_tokens=True)).unsqueeze(0) # Batch size 1
|
||||
>>> logits = model(input_ids).logits
|
||||
>>> inputs = tokenizer("Hello, my dog is cute", return_tensors= "pt")
|
||||
>>> logits = model(**inputs).logits
|
||||
"""
|
||||
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -27,6 +27,7 @@ from .configuration_auto import (
|
||||
DistilBertConfig,
|
||||
ElectraConfig,
|
||||
FlaubertConfig,
|
||||
FunnelConfig,
|
||||
GPT2Config,
|
||||
LongformerConfig,
|
||||
MobileBertConfig,
|
||||
@@ -37,6 +38,7 @@ from .configuration_auto import (
|
||||
XLMConfig,
|
||||
XLMRobertaConfig,
|
||||
XLNetConfig,
|
||||
replace_list_option_in_docstrings,
|
||||
)
|
||||
from .configuration_utils import PretrainedConfig
|
||||
from .modeling_tf_albert import (
|
||||
@@ -92,6 +94,15 @@ from .modeling_tf_flaubert import (
|
||||
TFFlaubertModel,
|
||||
TFFlaubertWithLMHeadModel,
|
||||
)
|
||||
from .modeling_tf_funnel import (
|
||||
TFFunnelForMaskedLM,
|
||||
TFFunnelForMultipleChoice,
|
||||
TFFunnelForPreTraining,
|
||||
TFFunnelForQuestionAnswering,
|
||||
TFFunnelForSequenceClassification,
|
||||
TFFunnelForTokenClassification,
|
||||
TFFunnelModel,
|
||||
)
|
||||
from .modeling_tf_gpt2 import TFGPT2LMHeadModel, TFGPT2Model
|
||||
from .modeling_tf_longformer import TFLongformerForMaskedLM, TFLongformerForQuestionAnswering, TFLongformerModel
|
||||
from .modeling_tf_mobilebert import (
|
||||
@@ -163,6 +174,7 @@ TF_MODEL_MAPPING = OrderedDict(
|
||||
(XLMConfig, TFXLMModel),
|
||||
(CTRLConfig, TFCTRLModel),
|
||||
(ElectraConfig, TFElectraModel),
|
||||
(FunnelConfig, TFFunnelModel),
|
||||
]
|
||||
)
|
||||
|
||||
@@ -184,6 +196,7 @@ TF_MODEL_FOR_PRETRAINING_MAPPING = OrderedDict(
|
||||
(XLMConfig, TFXLMWithLMHeadModel),
|
||||
(CTRLConfig, TFCTRLLMHeadModel),
|
||||
(ElectraConfig, TFElectraForPreTraining),
|
||||
(FunnelConfig, TFFunnelForPreTraining),
|
||||
]
|
||||
)
|
||||
|
||||
@@ -206,6 +219,7 @@ TF_MODEL_WITH_LM_HEAD_MAPPING = OrderedDict(
|
||||
(XLMConfig, TFXLMWithLMHeadModel),
|
||||
(CTRLConfig, TFCTRLLMHeadModel),
|
||||
(ElectraConfig, TFElectraForMaskedLM),
|
||||
(FunnelConfig, TFFunnelForMaskedLM),
|
||||
]
|
||||
)
|
||||
|
||||
@@ -237,6 +251,7 @@ TF_MODEL_FOR_MASKED_LM_MAPPING = OrderedDict(
|
||||
(FlaubertConfig, TFFlaubertWithLMHeadModel),
|
||||
(XLMConfig, TFXLMWithLMHeadModel),
|
||||
(ElectraConfig, TFElectraForMaskedLM),
|
||||
(FunnelConfig, TFFunnelForMaskedLM),
|
||||
]
|
||||
)
|
||||
|
||||
@@ -255,6 +270,7 @@ TF_MODEL_FOR_SEQUENCE_CLASSIFICATION_MAPPING = OrderedDict(
|
||||
(FlaubertConfig, TFFlaubertForSequenceClassification),
|
||||
(XLMConfig, TFXLMForSequenceClassification),
|
||||
(ElectraConfig, TFElectraForSequenceClassification),
|
||||
(FunnelConfig, TFFunnelForSequenceClassification),
|
||||
]
|
||||
)
|
||||
|
||||
@@ -272,6 +288,7 @@ TF_MODEL_FOR_QUESTION_ANSWERING_MAPPING = OrderedDict(
|
||||
(FlaubertConfig, TFFlaubertForQuestionAnsweringSimple),
|
||||
(XLMConfig, TFXLMForQuestionAnsweringSimple),
|
||||
(ElectraConfig, TFElectraForQuestionAnswering),
|
||||
(FunnelConfig, TFFunnelForQuestionAnswering),
|
||||
]
|
||||
)
|
||||
|
||||
@@ -288,6 +305,7 @@ TF_MODEL_FOR_TOKEN_CLASSIFICATION_MAPPING = OrderedDict(
|
||||
(MobileBertConfig, TFMobileBertForTokenClassification),
|
||||
(XLNetConfig, TFXLNetForTokenClassification),
|
||||
(ElectraConfig, TFElectraForTokenClassification),
|
||||
(FunnelConfig, TFFunnelForTokenClassification),
|
||||
]
|
||||
)
|
||||
|
||||
@@ -304,6 +322,7 @@ TF_MODEL_FOR_MULTIPLE_CHOICE_MAPPING = OrderedDict(
|
||||
(FlaubertConfig, TFFlaubertForMultipleChoice),
|
||||
(AlbertConfig, TFAlbertForMultipleChoice),
|
||||
(ElectraConfig, TFElectraForMultipleChoice),
|
||||
(FunnelConfig, TFFunnelForMultipleChoice),
|
||||
]
|
||||
)
|
||||
|
||||
@@ -315,21 +334,6 @@ class TFAutoModel(object):
|
||||
when created with the `TFAutoModel.from_pretrained(pretrained_model_name_or_path)`
|
||||
class method.
|
||||
|
||||
The `from_pretrained()` method takes care of returning the correct model class instance
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
- `t5`: TFT5Model (T5 model)
|
||||
- `distilbert`: TFDistilBertModel (DistilBERT model)
|
||||
- `roberta`: TFRobertaModel (RoBERTa model)
|
||||
- `bert`: TFBertModel (Bert model)
|
||||
- `openai-gpt`: TFOpenAIGPTModel (OpenAI GPT model)
|
||||
- `gpt2`: TFGPT2Model (OpenAI GPT-2 model)
|
||||
- `transfo-xl`: TFTransfoXLModel (Transformer-XL model)
|
||||
- `xlnet`: TFXLNetModel (XLNet model)
|
||||
- `xlm`: TFXLMModel (XLM model)
|
||||
- `ctrl`: TFCTRLModel (CTRL model)
|
||||
|
||||
This class cannot be instantiated using `__init__()` (throws an error).
|
||||
"""
|
||||
|
||||
@@ -341,6 +345,7 @@ class TFAutoModel(object):
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(TF_MODEL_MAPPING, use_model_types=False)
|
||||
def from_config(cls, config):
|
||||
r"""Instantiates one of the base model classes of the library
|
||||
from a configuration.
|
||||
@@ -354,15 +359,7 @@ class TFAutoModel(object):
|
||||
config: (`optional`) instance of a class derived from :class:`~transformers.TFPretrainedConfig`:
|
||||
The model class to instantiate is selected based on the configuration class:
|
||||
|
||||
- isInstance of `distilbert` configuration class: TFDistilBertModel (DistilBERT model)
|
||||
- isInstance of `roberta` configuration class: TFRobertaModel (RoBERTa model)
|
||||
- isInstance of `bert` configuration class: TFBertModel (Bert model)
|
||||
- isInstance of `openai-gpt` configuration class: TFOpenAIGPTModel (OpenAI GPT model)
|
||||
- isInstance of `gpt2` configuration class: TFGPT2Model (OpenAI GPT-2 model)
|
||||
- isInstance of `ctrl` configuration class: TFCTRLModel (Salesforce CTRL model)
|
||||
- isInstance of `transfo-xl` configuration class: TFTransfoXLModel (Transformer-XL model)
|
||||
- isInstance of `xlnet` configuration class: TFXLNetModel (XLNet model)
|
||||
- isInstance of `xlm` configuration class: TFXLMModel (XLM model)
|
||||
List options
|
||||
|
||||
Examples::
|
||||
|
||||
@@ -380,6 +377,7 @@ class TFAutoModel(object):
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(TF_MODEL_MAPPING)
|
||||
def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
|
||||
r"""Instantiates one of the base model classes of the library
|
||||
from a pre-trained model configuration.
|
||||
@@ -388,15 +386,7 @@ class TFAutoModel(object):
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
- `t5`: TFT5Model (T5 model)
|
||||
- `distilbert`: TFDistilBertModel (DistilBERT model)
|
||||
- `roberta`: TFRobertaModel (RoBERTa model)
|
||||
- `bert`: TFTFBertModel (Bert model)
|
||||
- `openai-gpt`: TFOpenAIGPTModel (OpenAI GPT model)
|
||||
- `gpt2`: TFGPT2Model (OpenAI GPT-2 model)
|
||||
- `transfo-xl`: TFTransfoXLModel (Transformer-XL model)
|
||||
- `xlnet`: TFXLNetModel (XLNet model)
|
||||
- `ctrl`: TFCTRLModel (CTRL model)
|
||||
List options
|
||||
|
||||
Params:
|
||||
pretrained_model_name_or_path: either:
|
||||
@@ -492,6 +482,7 @@ class TFAutoModelForPreTraining(object):
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(TF_MODEL_FOR_PRETRAINING_MAPPING, use_model_types=False)
|
||||
def from_config(cls, config):
|
||||
r"""Instantiates one of the base model classes of the library
|
||||
from a configuration.
|
||||
@@ -505,15 +496,7 @@ class TFAutoModelForPreTraining(object):
|
||||
config (:class:`~transformers.TFPretrainedConfig`):
|
||||
The model class to instantiate is selected based on the configuration class:
|
||||
|
||||
- isInstance of `distilbert` configuration class: :class:`~transformers.TFDistilBertModelForMaskedLM` (DistilBERT model)
|
||||
- isInstance of `roberta` configuration class: :class:`~transformers.TFRobertaModelForMaskedLM` (RoBERTa model)
|
||||
- isInstance of `bert` configuration class: :class:`~transformers.TFBertForPreTraining` (Bert model)
|
||||
- isInstance of `openai-gpt` configuration class: :class:`~transformers.TFOpenAIGPTLMHeadModel` (OpenAI GPT model)
|
||||
- isInstance of `gpt2` configuration class: :class:`~transformers.TFGPT2ModelLMHeadModel` (OpenAI GPT-2 model)
|
||||
- isInstance of `ctrl` configuration class: :class:`~transformers.TFCTRLModelLMHeadModel` (Salesforce CTRL model)
|
||||
- isInstance of `transfo-xl` configuration class: :class:`~transformers.TFTransfoXLLMHeadModel` (Transformer-XL model)
|
||||
- isInstance of `xlnet` configuration class: :class:`~transformers.TFXLNetLMHeadModel` (XLNet model)
|
||||
- isInstance of `xlm` configuration class: :class:`~transformers.TFXLMWithLMHeadModel` (XLM model)
|
||||
List options
|
||||
|
||||
Examples::
|
||||
|
||||
@@ -531,6 +514,7 @@ class TFAutoModelForPreTraining(object):
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(TF_MODEL_FOR_PRETRAINING_MAPPING)
|
||||
def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
|
||||
r"""Instantiates one of the model classes of the library -with the architecture used for pretraining this model– from a pre-trained model configuration.
|
||||
|
||||
@@ -538,17 +522,7 @@ class TFAutoModelForPreTraining(object):
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
- `t5`: :class:`~transformers.TFT5ModelWithLMHead` (T5 model)
|
||||
- `distilbert`: :class:`~transformers.TFDistilBertForMaskedLM` (DistilBERT model)
|
||||
- `albert`: :class:`~transformers.TFAlbertForPreTraining` (ALBERT model)
|
||||
- `roberta`: :class:`~transformers.TFRobertaForMaskedLM` (RoBERTa model)
|
||||
- `bert`: :class:`~transformers.TFBertForPreTraining` (Bert model)
|
||||
- `openai-gpt`: :class:`~transformers.TFOpenAIGPTLMHeadModel` (OpenAI GPT model)
|
||||
- `gpt2`: :class:`~transformers.TFGPT2LMHeadModel` (OpenAI GPT-2 model)
|
||||
- `transfo-xl`: :class:`~transformers.TFTransfoXLLMHeadModel` (Transformer-XL model)
|
||||
- `xlnet`: :class:`~transformers.TFXLNetLMHeadModel` (XLNet model)
|
||||
- `xlm`: :class:`~transformers.TFXLMWithLMHeadModel` (XLM model)
|
||||
- `ctrl`: :class:`~transformers.TFCTRLLMHeadModel` (Salesforce CTRL model)
|
||||
List options
|
||||
|
||||
The model is set in evaluation mode by default using `model.eval()` (Dropout modules are deactivated)
|
||||
To train the model, you should first set it back in training mode with `model.train()`
|
||||
@@ -635,21 +609,6 @@ class TFAutoModelWithLMHead(object):
|
||||
when created with the `TFAutoModelWithLMHead.from_pretrained(pretrained_model_name_or_path)`
|
||||
class method.
|
||||
|
||||
The `from_pretrained()` method takes care of returning the correct model class instance
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
- `t5`: TFT5ForConditionalGeneration (T5 model)
|
||||
- `distilbert`: TFDistilBertForMaskedLM (DistilBERT model)
|
||||
- `roberta`: TFRobertaForMaskedLM (RoBERTa model)
|
||||
- `bert`: TFBertForMaskedLM (Bert model)
|
||||
- `openai-gpt`: TFOpenAIGPTLMHeadModel (OpenAI GPT model)
|
||||
- `gpt2`: TFGPT2LMHeadModel (OpenAI GPT-2 model)
|
||||
- `transfo-xl`: TFTransfoXLLMHeadModel (Transformer-XL model)
|
||||
- `xlnet`: TFXLNetLMHeadModel (XLNet model)
|
||||
- `xlm`: TFXLMWithLMHeadModel (XLM model)
|
||||
- `ctrl`: TFCTRLLMHeadModel (CTRL model)
|
||||
|
||||
This class cannot be instantiated using `__init__()` (throws an error).
|
||||
"""
|
||||
|
||||
@@ -661,6 +620,7 @@ class TFAutoModelWithLMHead(object):
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(TF_MODEL_WITH_LM_HEAD_MAPPING, use_model_types=False)
|
||||
def from_config(cls, config):
|
||||
r"""Instantiates one of the base model classes of the library
|
||||
from a configuration.
|
||||
@@ -674,15 +634,7 @@ class TFAutoModelWithLMHead(object):
|
||||
config: (`optional`) instance of a class derived from :class:`~transformers.TFPretrainedConfig`:
|
||||
The model class to instantiate is selected based on the configuration class:
|
||||
|
||||
- isInstance of `distilbert` configuration class: TFDistilBertModel (DistilBERT model)
|
||||
- isInstance of `roberta` configuration class: TFRobertaModel (RoBERTa model)
|
||||
- isInstance of `bert` configuration class: TFBertModel (Bert model)
|
||||
- isInstance of `openai-gpt` configuration class: OpenAIGPTModel (OpenAI GPT model)
|
||||
- isInstance of `gpt2` configuration class: TFGPT2Model (OpenAI GPT-2 model)
|
||||
- isInstance of `ctrl` configuration class: TFCTRLModel (Salesforce CTRL model)
|
||||
- isInstance of `transfo-xl` configuration class: TransfoXLModel (Transformer-XL model)
|
||||
- isInstance of `xlnet` configuration class: TFXLNetModel (XLNet model)
|
||||
- isInstance of `xlm` configuration class: TFXLMModel (XLM model)
|
||||
List options
|
||||
|
||||
Examples::
|
||||
|
||||
@@ -704,6 +656,7 @@ class TFAutoModelWithLMHead(object):
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(TF_MODEL_WITH_LM_HEAD_MAPPING)
|
||||
def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
|
||||
r"""Instantiates one of the language modeling model classes of the library
|
||||
from a pre-trained model configuration.
|
||||
@@ -712,16 +665,7 @@ class TFAutoModelWithLMHead(object):
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
- `t5`: TFT5ForConditionalGeneration (T5 model)
|
||||
- `distilbert`: TFDistilBertForMaskedLM (DistilBERT model)
|
||||
- `roberta`: TFRobertaForMaskedLM (RoBERTa model)
|
||||
- `bert`: TFBertForMaskedLM (Bert model)
|
||||
- `openai-gpt`: TFOpenAIGPTLMHeadModel (OpenAI GPT model)
|
||||
- `gpt2`: TFGPT2LMHeadModel (OpenAI GPT-2 model)
|
||||
- `transfo-xl`: TFTransfoXLLMHeadModel (Transformer-XL model)
|
||||
- `xlnet`: TFXLNetLMHeadModel (XLNet model)
|
||||
- `xlm`: TFXLMWithLMHeadModel (XLM model)
|
||||
- `ctrl`: TFCTRLLMHeadModel (CTRL model)
|
||||
List options
|
||||
|
||||
Params:
|
||||
pretrained_model_name_or_path: either:
|
||||
@@ -813,12 +757,6 @@ class TFAutoModelForMultipleChoice:
|
||||
when created with the `TFAutoModelForMultipleChoice.from_pretrained(pretrained_model_name_or_path)`
|
||||
class method.
|
||||
|
||||
The `from_pretrained()` method takes care of returning the correct model class instance
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
- `albert`: TFAlbertForMultipleChoice (Albert model)
|
||||
- `bert`: TFBertForMultipleChoice (Bert model)
|
||||
|
||||
This class cannot be instantiated using `__init__()` (throws an error).
|
||||
"""
|
||||
|
||||
@@ -830,6 +768,7 @@ class TFAutoModelForMultipleChoice:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(TF_MODEL_FOR_MULTIPLE_CHOICE_MAPPING, use_model_types=False)
|
||||
def from_config(cls, config):
|
||||
r"""Instantiates one of the base model classes of the library
|
||||
from a configuration.
|
||||
@@ -842,8 +781,8 @@ class TFAutoModelForMultipleChoice:
|
||||
Args:
|
||||
config: (`optional`) instance of a class derived from :class:`~transformers.TFPretrainedConfig`:
|
||||
The model class to instantiate is selected based on the configuration class:
|
||||
- isInstance of `albert` configuration class: TFAlbertModel (Albert model)
|
||||
- isInstance of `bert` configuration class: TFBertModel (Bert model)
|
||||
|
||||
List options
|
||||
|
||||
Examples::
|
||||
|
||||
@@ -863,6 +802,7 @@ class TFAutoModelForMultipleChoice:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(TF_MODEL_FOR_MULTIPLE_CHOICE_MAPPING)
|
||||
def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
|
||||
r"""Instantiates one of the multiple choice model classes of the library
|
||||
from a pre-trained model configuration.
|
||||
@@ -871,8 +811,7 @@ class TFAutoModelForMultipleChoice:
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
- `albert`: TFRobertaForMultiple (Albert model)
|
||||
- `bert`: TFBertForMultipleChoice (Bert model)
|
||||
List options
|
||||
|
||||
The model is set in evaluation mode by default using `model.eval()` (Dropout modules are deactivated)
|
||||
To train the model, you should first set it back in training mode with `model.train()`
|
||||
@@ -974,6 +913,7 @@ class TFAutoModelForCausalLM:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(TF_MODEL_FOR_CAUSAL_LM_MAPPING, use_model_types=False)
|
||||
def from_config(cls, config):
|
||||
r"""Instantiates one of the base model classes of the library
|
||||
from a configuration.
|
||||
@@ -987,12 +927,7 @@ class TFAutoModelForCausalLM:
|
||||
config (:class:`~transformers.TFPretrainedConfig`):
|
||||
The model class to instantiate is selected based on the configuration class:
|
||||
|
||||
- isInstance of `bert` configuration class: :class:`~transformers.TFBertLMHeadModel` (Bert model)
|
||||
- isInstance of `openai-gpt` configuration class: :class:`~transformers.TFOpenAIGPTLMHeadModel` (OpenAI GPT model)
|
||||
- isInstance of `gpt2` configuration class: :class:`~transformers.TFGPT2LMHeadModel` (OpenAI GPT-2 model)
|
||||
- isInstance of `ctrl` configuration class: :class:`~transformers.TFCTRLLMHeadModel` (Salesforce CTRL model)
|
||||
- isInstance of `transfo-xl` configuration class: :class:`~transformers.TFTransfoXLLMHeadModel` (Transformer-XL model)
|
||||
- isInstance of `xlnet` configuration class: :class:`~transformers.TFXLNetLMHeadModel` (XLNet model)
|
||||
List options
|
||||
|
||||
Examples::
|
||||
|
||||
@@ -1010,6 +945,7 @@ class TFAutoModelForCausalLM:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(TF_MODEL_FOR_CAUSAL_LM_MAPPING)
|
||||
def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
|
||||
r"""Instantiates one of the language modeling model classes of the library
|
||||
from a pre-trained model configuration.
|
||||
@@ -1018,12 +954,7 @@ class TFAutoModelForCausalLM:
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
- `bert`: :class:`~transformers.TFBertLMHeadModel` (Bert model)
|
||||
- `openai-gpt`: :class:`~transformers.TFOpenAIGPTLMHeadModel` (OpenAI GPT model)
|
||||
- `gpt2`: :class:`~transformers.TFGPT2LMHeadModel` (OpenAI GPT-2 model)
|
||||
- `transfo-xl`: :class:`~transformers.TFTransfoXLLMHeadModel` (Transformer-XL model)
|
||||
- `xlnet`: :class:`~transformers.TFXLNetLMHeadModel` (XLNet model)
|
||||
- `ctrl`: :class:`~transformers.TFCTRLLMHeadModel` (Salesforce CTRL model)
|
||||
List options
|
||||
|
||||
The model is set in evaluation mode by default using `model.eval()` (Dropout modules are deactivated)
|
||||
To train the model, you should first set it back in training mode with `model.train()`
|
||||
@@ -1110,6 +1041,7 @@ class TFAutoModelForMaskedLM:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(TF_MODEL_FOR_MASKED_LM_MAPPING, use_model_types=False)
|
||||
def from_config(cls, config):
|
||||
r"""Instantiates one of the base model classes of the library
|
||||
from a configuration.
|
||||
@@ -1122,16 +1054,8 @@ class TFAutoModelForMaskedLM:
|
||||
Args:
|
||||
config (:class:`~transformers.TFPretrainedConfig`):
|
||||
The model class to instantiate is selected based on the configuration class:
|
||||
- isInstance of `distilbert` configuration class: :class:`~transformers.TFDistilBertForMaskedLM` (DistilBERT model)
|
||||
- isInstance of `roberta` configuration class: :class:`~transformers.TFRobertaForMaskedLM` (RoBERTa model)
|
||||
- isInstance of `bert` configuration class: :class:`~transformers.TFBertForMaskedLM` (Bert model)
|
||||
- isInstance of `flaubert` configuration class: :class:`~transformers.TFFlaubertWithLMHeadModel` (Flaubert model)
|
||||
- isInstance of `xlm` configuration class: :class:`~transformers.TFXLMWithLMHeadModel` (XLM model)
|
||||
- isInstance of `xlm-roberta` configuration class: :class:`~transformers.TFXLMRobertaForMaskedLM` (XLM-Roberta model)
|
||||
- isInstance of `electra` configuration class: :class:`~transformers.TFElectraForMaskedLM` (Electra model)
|
||||
- isInstance of `camembert` configuration class: :class:`~transformers.TFCamembertForMaskedLM` (Camembert model)
|
||||
- isInstance of `albert` configuration class: :class:`~transformers.TFAlbertForMaskedLM` (Albert model)
|
||||
|
||||
List options
|
||||
|
||||
Examples::
|
||||
|
||||
@@ -1149,6 +1073,7 @@ class TFAutoModelForMaskedLM:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(TF_MODEL_FOR_MASKED_LM_MAPPING)
|
||||
def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
|
||||
r"""Instantiates one of the language modeling model classes of the library
|
||||
from a pre-trained model configuration.
|
||||
@@ -1157,16 +1082,7 @@ class TFAutoModelForMaskedLM:
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
- `distilbert`: :class:`~transformers.TFDistilBertForMaskedLM` (DistilBERT model)
|
||||
- `albert`: :class:`~transformers.TFAlbertForMaskedLM` (ALBERT model)
|
||||
- `camembert`: :class:`~transformers.TFCamembertForMaskedLM` (CamemBERT model)
|
||||
- `xlm-roberta`: :class:`~transformers.TFXLMRobertaForMaskedLM` (XLM-RoBERTa model)
|
||||
- `longformer`: :class:`~transformers.TFLongformerForMaskedLM` (Longformer model)
|
||||
- `roberta`: :class:`~transformers.TFRobertaForMaskedLM` (RoBERTa model)
|
||||
- `xlm`: :class:`~transformers.TFXLMWithLMHeadModel` (XLM model)
|
||||
- `flaubert`: :class:`~transformers.TFFlaubertWithLMHeadModel` (Flaubert model)
|
||||
- `electra`: :class:`~transformers.TFElectraForMaskedLM` (Electra model)
|
||||
- `bert`: :class:`~transformers.TFBertLMHeadModel` (Bert model)
|
||||
List options
|
||||
|
||||
The model is set in evaluation mode by default using `model.eval()` (Dropout modules are deactivated)
|
||||
To train the model, you should first set it back in training mode with `model.train()`
|
||||
@@ -1253,6 +1169,7 @@ class TFAutoModelForSeq2SeqLM:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(TF_MODEL_FOR_SEQ_TO_SEQ_CAUSAL_LM_MAPPING, use_model_types=False)
|
||||
def from_config(cls, config):
|
||||
r"""Instantiates one of the base model classes of the library
|
||||
from a configuration.
|
||||
@@ -1266,7 +1183,7 @@ class TFAutoModelForSeq2SeqLM:
|
||||
config (:class:`~transformers.TFPretrainedConfig`):
|
||||
The model class to instantiate is selected based on the configuration class:
|
||||
|
||||
- isInstance of `t5` configuration class: :class:`~transformers.TFT5ForConditionalGeneration` (T5 model)
|
||||
List options
|
||||
|
||||
Examples::
|
||||
|
||||
@@ -1286,6 +1203,7 @@ class TFAutoModelForSeq2SeqLM:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(TF_MODEL_FOR_SEQ_TO_SEQ_CAUSAL_LM_MAPPING, use_model_types=False)
|
||||
def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
|
||||
r"""Instantiates one of the language modeling model classes of the library
|
||||
from a pre-trained model configuration.
|
||||
@@ -1294,7 +1212,7 @@ class TFAutoModelForSeq2SeqLM:
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
- `t5`: :class:`~transformers.TFT5ForConditionalGeneration` (T5 model)
|
||||
List options
|
||||
|
||||
The model is set in evaluation mode by default using `model.eval()` (Dropout modules are deactivated)
|
||||
To train the model, you should first set it back in training mode with `model.train()`
|
||||
@@ -1372,16 +1290,6 @@ class TFAutoModelForSequenceClassification(object):
|
||||
when created with the `TFAutoModelForSequenceClassification.from_pretrained(pretrained_model_name_or_path)`
|
||||
class method.
|
||||
|
||||
The `from_pretrained()` method takes care of returning the correct model class instance
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
- `distilbert`: TFDistilBertForSequenceClassification (DistilBERT model)
|
||||
- `roberta`: TFRobertaForSequenceClassification (RoBERTa model)
|
||||
- `bert`: TFBertForSequenceClassification (Bert model)
|
||||
- `xlnet`: TFXLNetForSequenceClassification (XLNet model)
|
||||
- `xlm`: TFXLMForSequenceClassification (XLM model)
|
||||
|
||||
This class cannot be instantiated using `__init__()` (throws an error).
|
||||
"""
|
||||
|
||||
@@ -1393,6 +1301,7 @@ class TFAutoModelForSequenceClassification(object):
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(TF_MODEL_FOR_SEQUENCE_CLASSIFICATION_MAPPING, use_model_types=False)
|
||||
def from_config(cls, config):
|
||||
r"""Instantiates one of the base model classes of the library
|
||||
from a configuration.
|
||||
@@ -1406,11 +1315,7 @@ class TFAutoModelForSequenceClassification(object):
|
||||
config: (`optional`) instance of a class derived from :class:`~transformers.TFPretrainedConfig`:
|
||||
The model class to instantiate is selected based on the configuration class:
|
||||
|
||||
- isInstance of `distilbert` configuration class: DistilBertModel (DistilBERT model)
|
||||
- isInstance of `roberta` configuration class: RobertaModel (RoBERTa model)
|
||||
- isInstance of `bert` configuration class: BertModel (Bert model)
|
||||
- isInstance of `xlnet` configuration class: XLNetModel (XLNet model)
|
||||
- isInstance of `xlm` configuration class: XLMModel (XLM model)
|
||||
List options
|
||||
|
||||
Examples::
|
||||
|
||||
@@ -1430,6 +1335,7 @@ class TFAutoModelForSequenceClassification(object):
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(TF_MODEL_FOR_SEQUENCE_CLASSIFICATION_MAPPING)
|
||||
def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
|
||||
r"""Instantiates one of the sequence classification model classes of the library
|
||||
from a pre-trained model configuration.
|
||||
@@ -1438,11 +1344,7 @@ class TFAutoModelForSequenceClassification(object):
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
- `distilbert`: TFDistilBertForSequenceClassification (DistilBERT model)
|
||||
- `roberta`: TFRobertaForSequenceClassification (RoBERTa model)
|
||||
- `bert`: TFBertForSequenceClassification (Bert model)
|
||||
- `xlnet`: TFXLNetForSequenceClassification (XLNet model)
|
||||
- `xlm`: TFXLMForSequenceClassification (XLM model)
|
||||
List options
|
||||
|
||||
The model is set in evaluation mode by default using `model.eval()` (Dropout modules are deactivated)
|
||||
To train the model, you should first set it back in training mode with `model.train()`
|
||||
@@ -1533,17 +1435,6 @@ class TFAutoModelForQuestionAnswering(object):
|
||||
when created with the `TFAutoModelForQuestionAnswering.from_pretrained(pretrained_model_name_or_path)`
|
||||
class method.
|
||||
|
||||
The `from_pretrained()` method takes care of returning the correct model class instance
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
- `distilbert`: TFDistilBertForQuestionAnswering (DistilBERT model)
|
||||
- `albert`: TFAlbertForQuestionAnswering (ALBERT model)
|
||||
- `roberta`: TFRobertaForQuestionAnswering (RoBERTa model)
|
||||
- `bert`: TFBertForQuestionAnswering (Bert model)
|
||||
- `xlnet`: TFXLNetForQuestionAnswering (XLNet model)
|
||||
- `xlm`: TFXLMForQuestionAnswering (XLM model)
|
||||
|
||||
This class cannot be instantiated using `__init__()` (throws an error).
|
||||
"""
|
||||
|
||||
@@ -1555,6 +1446,7 @@ class TFAutoModelForQuestionAnswering(object):
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(TF_MODEL_FOR_QUESTION_ANSWERING_MAPPING, use_model_types=False)
|
||||
def from_config(cls, config):
|
||||
r"""Instantiates one of the base model classes of the library
|
||||
from a configuration.
|
||||
@@ -1568,12 +1460,7 @@ class TFAutoModelForQuestionAnswering(object):
|
||||
config: (`optional`) instance of a class derived from :class:`~transformers.TFPretrainedConfig`:
|
||||
The model class to instantiate is selected based on the configuration class:
|
||||
|
||||
- isInstance of `distilbert` configuration class: DistilBertModel (DistilBERT model)
|
||||
- isInstance of `albert` configuration class: AlbertModel (ALBERT model)
|
||||
- isInstance of `roberta` configuration class: RobertaModel (RoBERTa model)
|
||||
- isInstance of `bert` configuration class: BertModel (Bert model)
|
||||
- isInstance of `xlnet` configuration class: XLNetModel (XLNet model)
|
||||
- isInstance of `xlm` configuration class: XLMModel (XLM model)
|
||||
List options
|
||||
|
||||
Examples::
|
||||
|
||||
@@ -1593,6 +1480,7 @@ class TFAutoModelForQuestionAnswering(object):
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(TF_MODEL_FOR_QUESTION_ANSWERING_MAPPING)
|
||||
def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
|
||||
r"""Instantiates one of the question answering model classes of the library
|
||||
from a pre-trained model configuration.
|
||||
@@ -1601,12 +1489,7 @@ class TFAutoModelForQuestionAnswering(object):
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
- `distilbert`: TFDistilBertForQuestionAnswering (DistilBERT model)
|
||||
- `albert`: TFAlbertForQuestionAnswering (ALBERT model)
|
||||
- `roberta`: TFRobertaForQuestionAnswering (RoBERTa model)
|
||||
- `bert`: TFBertForQuestionAnswering (Bert model)
|
||||
- `xlnet`: TFXLNetForQuestionAnswering (XLNet model)
|
||||
- `xlm`: TFXLMForQuestionAnswering (XLM model)
|
||||
List options
|
||||
|
||||
The model is set in evaluation mode by default using `model.eval()` (Dropout modules are deactivated)
|
||||
To train the model, you should first set it back in training mode with `model.train()`
|
||||
@@ -1699,6 +1582,7 @@ class TFAutoModelForTokenClassification:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(TF_MODEL_FOR_TOKEN_CLASSIFICATION_MAPPING, use_model_types=False)
|
||||
def from_config(cls, config):
|
||||
r"""Instantiates one of the base model classes of the library
|
||||
from a configuration.
|
||||
@@ -1712,10 +1596,7 @@ class TFAutoModelForTokenClassification:
|
||||
config: (`optional`) instance of a class derived from :class:`~transformers.TFPretrainedConfig`:
|
||||
The model class to instantiate is selected based on the configuration class:
|
||||
|
||||
- isInstance of `bert` configuration class: BertModel (Bert model)
|
||||
- isInstance of `xlnet` configuration class: XLNetModel (XLNet model)
|
||||
- isInstance of `distilbert` configuration class: DistilBertModel (DistilBert model)
|
||||
- isInstance of `roberta` configuration class: RobteraModel (Roberta model)
|
||||
List options
|
||||
|
||||
Examples::
|
||||
|
||||
@@ -1735,6 +1616,7 @@ class TFAutoModelForTokenClassification:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(TF_MODEL_FOR_TOKEN_CLASSIFICATION_MAPPING)
|
||||
def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
|
||||
r"""Instantiates one of the question answering model classes of the library
|
||||
from a pre-trained model configuration.
|
||||
@@ -1743,10 +1625,7 @@ class TFAutoModelForTokenClassification:
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
- `bert`: BertForTokenClassification (Bert model)
|
||||
- `xlnet`: XLNetForTokenClassification (XLNet model)
|
||||
- `distilbert`: DistilBertForTokenClassification (DistilBert model)
|
||||
- `roberta`: RobertaForTokenClassification (Roberta model)
|
||||
List options
|
||||
|
||||
The model is set in evaluation mode by default using `model.eval()` (Dropout modules are deactivated)
|
||||
To train the model, you should first set it back in training mode with `model.train()`
|
||||
|
||||
File diff suppressed because it is too large.
Load diff
@@ -148,7 +148,7 @@ def load_pytorch_weights_in_tf2_model(tf_model, pt_state_dict, tf_inputs=None, a
|
||||
tf_loaded_numel = 0
|
||||
weight_value_tuples = []
|
||||
all_pytorch_weights = set(list(pt_state_dict.keys()))
|
||||
unexpected_keys = []
|
||||
missing_keys = []
|
||||
for symbolic_weight in symbolic_weights:
|
||||
sw_name = symbolic_weight.name
|
||||
name, transpose = convert_tf_weight_name_to_pt_weight_name(
|
||||
@@ -158,7 +158,7 @@ def load_pytorch_weights_in_tf2_model(tf_model, pt_state_dict, tf_inputs=None, a
|
||||
# Find associated numpy array in pytorch model state dict
|
||||
if name not in pt_state_dict:
|
||||
if allow_missing_keys:
|
||||
unexpected_keys.append(name)
|
||||
missing_keys.append(name)
|
||||
continue
|
||||
|
||||
raise AttributeError("{} not found in PyTorch model".format(name))
|
||||
@@ -192,28 +192,28 @@ def load_pytorch_weights_in_tf2_model(tf_model, pt_state_dict, tf_inputs=None, a
|
||||
|
||||
logger.info("Loaded {:,} parameters in the TF 2.0 model.".format(tf_loaded_numel))
|
||||
|
||||
missing_keys = list(all_pytorch_weights)
|
||||
unexpected_keys = list(all_pytorch_weights)
|
||||
|
||||
if len(unexpected_keys) > 0:
|
||||
logger.warning(
|
||||
f"Some weights of the PyTorch model were not used when "
|
||||
f"initializing the TF 2.0 model {tf_model.__class__.__name__}: {unexpected_keys}\n"
|
||||
f"- This IS expected if you are initializing {tf_model.__class__.__name__} from a TF 2.0 model trained on another task "
|
||||
f"or with another architecture (e.g. initializing a BertForSequenceClassification model from a TFBertForPretraining model).\n"
|
||||
f"- This IS NOT expected if you are initializing {tf_model.__class__.__name__} from a TF 2.0 model that you expect "
|
||||
f"to be exactly identical (e.g. initializing a BertForSequenceClassification model from a TFBertForSequenceClassification model)."
|
||||
f"- This IS expected if you are initializing {tf_model.__class__.__name__} from a PyTorch model trained on another task "
|
||||
f"or with another architecture (e.g. initializing a TFBertForSequenceClassification model from a BertForPretraining model).\n"
|
||||
f"- This IS NOT expected if you are initializing {tf_model.__class__.__name__} from a PyTorch model that you expect "
|
||||
f"to be exactly identical (e.g. initializing a TFBertForSequenceClassification model from a BertForSequenceClassification model)."
|
||||
)
|
||||
else:
|
||||
logger.warning(f"All PyTorch model weights were used when initializing {tf_model.__class__.__name__}.\n")
|
||||
if len(missing_keys) > 0:
|
||||
logger.warning(
|
||||
f"Some weights or buffers of the PyTorch model {tf_model.__class__.__name__} were not initialized from the TF 2.0 model "
|
||||
f"Some weights or buffers of the TF 2.0 model {tf_model.__class__.__name__} were not initialized from the PyTorch model "
|
||||
f"and are newly initialized: {missing_keys}\n"
|
||||
f"You should probably TRAIN this model on a down-stream task to be able to use it for predictions and inference."
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
f"All the weights of {tf_model.__class__.__name__} were initialized from the TF 2.0 model.\n"
|
||||
f"All the weights of {tf_model.__class__.__name__} were initialized from the PyTorch model.\n"
|
||||
f"If your task is similar to the task the model of the ckeckpoint was trained on, "
|
||||
f"you can already use {tf_model.__class__.__name__} for predictions without further training."
|
||||
)
|
||||
|
||||
@@ -17,8 +17,9 @@
|
||||
import inspect
|
||||
import os
|
||||
import re
|
||||
import warnings
|
||||
from dataclasses import dataclass
|
||||
from typing import Callable, Dict, List, Optional, Set, Tuple, Union
|
||||
from typing import Any, Callable, Dict, List, Optional, Set, Tuple, Union
|
||||
|
||||
import torch
|
||||
from torch import Tensor, device, dtype, nn
|
||||
@@ -45,7 +46,6 @@ from .utils import logging
|
||||
|
||||
logger = logging.get_logger(__name__)
|
||||
|
||||
|
||||
try:
|
||||
from torch.nn import Identity
|
||||
except ImportError:
|
||||
@@ -91,20 +91,6 @@ class ModuleUtilsMixin:
|
||||
A few utilities for :obj:`torch.nn.Modules`, to be used as a mixin.
|
||||
"""
|
||||
|
||||
def num_parameters(self, only_trainable: bool = False) -> int:
|
||||
"""
|
||||
Get the number of (optionally, trainable) parameters in the model.
|
||||
|
||||
Args:
|
||||
only_trainable (:obj:`bool`, `optional`, defaults to :obj:`False`):
|
||||
Whether or not to return only the number of trainable parameters
|
||||
|
||||
Returns:
|
||||
:obj:`int`: The number of parameters.
|
||||
"""
|
||||
params = filter(lambda x: x.requires_grad, self.parameters()) if only_trainable else self.parameters()
|
||||
return sum(p.numel() for p in params)
|
||||
|
||||
@staticmethod
|
||||
def _hook_rss_memory_pre_forward(module, *args, **kwargs):
|
||||
try:
|
||||
@@ -172,10 +158,10 @@ class ModuleUtilsMixin:
|
||||
first_tuple = next(gen)
|
||||
return first_tuple[1].device
|
||||
|
||||
# TorchScript does not support, so add non-property option
|
||||
def get_dtype(self) -> dtype:
|
||||
@property
|
||||
def dtype(self) -> dtype:
|
||||
"""
|
||||
Get torch.dtype from module, assuming that the whole module has one dtype.
|
||||
:obj:`torch.dtype`: The dtype of the module (assuming that all the module parameters have the same dtype).
|
||||
"""
|
||||
try:
|
||||
return next(self.parameters()).dtype
|
||||
@@ -190,10 +176,6 @@ class ModuleUtilsMixin:
|
||||
first_tuple = next(gen)
|
||||
return first_tuple[1].dtype
|
||||
|
||||
@property
|
||||
def dtype(self):
|
||||
return self.get_dtype()
|
||||
|
||||
def invert_attention_mask(self, encoder_attention_mask: Tensor) -> Tensor:
|
||||
"""
|
||||
Invert an attention mask (e.g., switches 0. and 1.).
|
||||
@@ -204,43 +186,31 @@ class ModuleUtilsMixin:
|
||||
Returns:
|
||||
:obj:`torch.Tensor`: The inverted attention mask.
|
||||
"""
|
||||
encoder_extended_attention_mask: Optional[Tensor] = None
|
||||
if encoder_attention_mask.dim() == 3:
|
||||
encoder_extended_attention_mask = encoder_attention_mask[:, None, :, :]
|
||||
if encoder_attention_mask.dim() == 2:
|
||||
encoder_extended_attention_mask = encoder_attention_mask[:, None, None, :]
|
||||
assert encoder_extended_attention_mask is not None
|
||||
# T5 has a mask that can compare sequence ids, we can simulate this here with this transposition
|
||||
# Cf. https://github.com/tensorflow/mesh/blob/8d2465e9bc93129b913b5ccc6a59aa97abd96ec6/mesh_tensorflow
|
||||
# /transformer/transformer_layers.py#L270
|
||||
# encoder_extended_attention_mask = (encoder_extended_attention_mask ==
|
||||
# encoder_extended_attention_mask.transpose(-1, -2))
|
||||
encoder_extended_attention_mask = encoder_extended_attention_mask.to(
|
||||
dtype=self.get_dtype()
|
||||
) # fp16 compatibility
|
||||
encoder_extended_attention_mask = encoder_extended_attention_mask.to(dtype=self.dtype) # fp16 compatibility
|
||||
|
||||
if self.get_dtype() == torch.float16:
|
||||
if self.dtype == torch.float16:
|
||||
encoder_extended_attention_mask = (1.0 - encoder_extended_attention_mask) * -1e4
|
||||
elif self.get_dtype() == torch.float32:
|
||||
elif self.dtype == torch.float32:
|
||||
encoder_extended_attention_mask = (1.0 - encoder_extended_attention_mask) * -1e9
|
||||
else:
|
||||
raise ValueError(
|
||||
"{} not recognized. `dtype` should be set to either `torch.float32` or `torch.float16`".format(
|
||||
self.get_dtype()
|
||||
self.dtype
|
||||
)
|
||||
)
|
||||
|
||||
return encoder_extended_attention_mask
|
||||
|
||||
def get_is_decoder(self):
|
||||
if hasattr(self, "is_decoder"):
|
||||
return self.is_decoder
|
||||
else:
|
||||
return self.config.is_decoder
|
||||
|
||||
def get_extended_attention_mask(
|
||||
self, attention_mask: Tensor, input_shape: Tuple[int, int], device: device
|
||||
) -> Tensor:
|
||||
def get_extended_attention_mask(self, attention_mask: Tensor, input_shape: Tuple[int], device: device) -> Tensor:
|
||||
"""
|
||||
Makes broadcastable attention and causal masks so that future and masked tokens are ignored.
|
||||
|
||||
@@ -263,7 +233,7 @@ class ModuleUtilsMixin:
|
||||
# Provided a padding mask of dimensions [batch_size, seq_length]
|
||||
# - if the model is a decoder, apply a causal mask in addition to the padding mask
|
||||
# - if the model is an encoder, make the mask broadcastable to [batch_size, num_heads, seq_length, seq_length]
|
||||
if self.get_is_decoder():
|
||||
if self.config.is_decoder:
|
||||
batch_size, seq_length = input_shape
|
||||
seq_ids = torch.arange(seq_length, device=device)
|
||||
causal_mask = seq_ids[None, None, :].repeat(batch_size, seq_length, 1) <= seq_ids[None, :, None]
|
||||
@@ -284,7 +254,7 @@ class ModuleUtilsMixin:
|
||||
# positions we want to attend and -10000.0 for masked positions.
|
||||
# Since we are adding it to the raw scores before the softmax, this is
|
||||
# effectively the same as removing these entirely.
|
||||
extended_attention_mask = extended_attention_mask.to(dtype=self.get_dtype()) # fp16 compatibility
|
||||
extended_attention_mask = extended_attention_mask.to(dtype=self.dtype) # fp16 compatibility
|
||||
extended_attention_mask = (1.0 - extended_attention_mask) * -10000.0
|
||||
return extended_attention_mask
|
||||
|
||||
@@ -306,34 +276,16 @@ class ModuleUtilsMixin:
|
||||
:obj:`torch.Tensor` with shape :obj:`[num_hidden_layers x batch x num_heads x seq_length x seq_length]`
|
||||
or list with :obj:`[None]` for each layer.
|
||||
"""
|
||||
if head_mask is None:
|
||||
return [None] * num_hidden_layers
|
||||
else:
|
||||
return self.get_scriptable_head_mask(head_mask, num_hidden_layers, is_attention_chunked)
|
||||
|
||||
def get_scriptable_head_mask(
|
||||
self, head_mask: Optional[Tensor], num_hidden_layers: int, is_attention_chunked: bool = False
|
||||
) -> Optional[Tensor]:
|
||||
"""
|
||||
# Prepare head mask if needed
|
||||
# 1.0 in head_mask indicate we keep the head
|
||||
attention_probs has shape bsz x n_heads x N x N
|
||||
Arguments:
|
||||
head_mask: torch.Tensor or None: has shape [num_heads] or [num_hidden_layers x num_heads]
|
||||
num_hidden_layers: int
|
||||
Returns:
|
||||
Tensor of shape shape [num_hidden_layers x batch x num_heads x seq_length x seq_length]
|
||||
or None
|
||||
"""
|
||||
if head_mask is None:
|
||||
return None
|
||||
else:
|
||||
if head_mask is not None:
|
||||
head_mask = self._convert_head_mask_to_5d(head_mask, num_hidden_layers)
|
||||
if is_attention_chunked is True:
|
||||
head_mask = head_mask.unsqueeze(-1)
|
||||
return head_mask
|
||||
else:
|
||||
head_mask = [None] * num_hidden_layers
|
||||
|
||||
def _convert_head_mask_to_5d(self, head_mask, num_hidden_layers: int):
|
||||
return head_mask
|
||||
|
||||
def _convert_head_mask_to_5d(self, head_mask, num_hidden_layers):
|
||||
"""-> [num_hidden_layers x batch x num_heads x seq_length x seq_length]"""
|
||||
if head_mask.dim() == 1:
|
||||
head_mask = head_mask.unsqueeze(0).unsqueeze(0).unsqueeze(-1).unsqueeze(-1)
|
||||
@@ -341,9 +293,77 @@ class ModuleUtilsMixin:
|
||||
elif head_mask.dim() == 2:
|
||||
head_mask = head_mask.unsqueeze(1).unsqueeze(-1).unsqueeze(-1) # We can specify head_mask for each layer
|
||||
assert head_mask.dim() == 5, f"head_mask.dim != 5, instead {head_mask.dim()}"
|
||||
head_mask = head_mask.to(dtype=self.get_dtype()) # switch to fload if need + fp16 compatibility
|
||||
head_mask = head_mask.to(dtype=self.dtype) # switch to float if need + fp16 compatibility
|
||||
return head_mask
|
||||
|
||||
def num_parameters(self, only_trainable: bool = False, exclude_embeddings: bool = False) -> int:
|
||||
"""
|
||||
Get number of (optionally, trainable or non-embeddings) parameters in the module.
|
||||
|
||||
Args:
|
||||
only_trainable (:obj:`bool`, `optional`, defaults to :obj:`False`):
|
||||
Whether or not to return only the number of trainable parameters
|
||||
|
||||
exclude_embeddings (:obj:`bool`, `optional`, defaults to :obj:`False`):
|
||||
Whether or not to return only the number of non-embeddings parameters
|
||||
|
||||
Returns:
|
||||
:obj:`int`: The number of parameters.
|
||||
"""
|
||||
|
||||
def parameter_filter(x):
|
||||
return (x.requires_grad or not only_trainable) and not (
|
||||
isinstance(x, torch.nn.Embedding) and exclude_embeddings
|
||||
)
|
||||
|
||||
params = filter(parameter_filter, self.parameters()) if only_trainable else self.parameters()
|
||||
return sum(p.numel() for p in params)
|
||||
|
||||
def estimate_tokens(self, input_dict: Dict[str, Union[torch.Tensor, Any]]) -> int:
|
||||
"""
|
||||
Helper function to estimate the total number of tokens from the model inputs.
|
||||
|
||||
Args:
|
||||
inputs (:obj:`dict`): The model inputs.
|
||||
|
||||
Returns:
|
||||
:obj:`int`: The total number of tokens.
|
||||
"""
|
||||
token_inputs = [tensor for key, tensor in input_dict.items() if "input" in key]
|
||||
if token_inputs:
|
||||
return sum([token_input.numel() for token_input in token_inputs])
|
||||
else:
|
||||
warnings.warn(
|
||||
"Could not estimate the number of tokens of the input, floating-point operations will not be computed"
|
||||
)
|
||||
return 0
|
||||
|
||||
def floating_point_ops(
|
||||
self, input_dict: Dict[str, Union[torch.Tensor, Any]], exclude_embeddings: bool = True
|
||||
) -> int:
|
||||
"""
|
||||
Get number of (optionally, non-embeddings) floating-point operations for the forward and backward passes of a
|
||||
batch with this transformer model. Default approximation neglects the quadratic dependency on the number of
|
||||
tokens (valid if :obj:`12 * d_model << sequence_length`) as laid out in `this paper <https://arxiv.org/pdf/2001.08361.pdf>`__ section
|
||||
2.1. Should be overriden for transformers with parameter re-use e.g. Albert or Universal Transformers, or
|
||||
if doing long-range modeling with very high sequence lengths.
|
||||
|
||||
Args:
|
||||
batch_size (:obj:`int`):
|
||||
The batch size for the forward pass.
|
||||
|
||||
sequence_length (:obj:`int`):
|
||||
The number of tokens in each line of the batch.
|
||||
|
||||
exclude_embeddings (:obj:`bool`, `optional`, defaults to :obj:`True`):
|
||||
Whether or not to count embedding and softmax operations.
|
||||
|
||||
Returns:
|
||||
:obj:`int`: The number of floating-point operations.
|
||||
"""
|
||||
|
||||
return 6 * self.estimate_tokens(input_dict) * self.num_parameters(exclude_embeddings=exclude_embeddings)
|
||||
|
||||
|
||||
class PreTrainedModel(nn.Module, ModuleUtilsMixin, GenerationMixin):
|
||||
r"""
|
||||
|
||||
@@ -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)
|
||||
@@ -122,6 +122,20 @@ def require_multigpu(test_case):
|
||||
return test_case
|
||||
|
||||
|
||||
def require_non_multigpu(test_case):
|
||||
"""
|
||||
Decorator marking a test that requires 0 or 1 GPU setup (in PyTorch).
|
||||
"""
|
||||
if not _torch_available:
|
||||
return unittest.skip("test requires PyTorch")(test_case)
|
||||
|
||||
import torch
|
||||
|
||||
if torch.cuda.device_count() > 1:
|
||||
return unittest.skip("test requires 0 or 1 GPU")(test_case)
|
||||
return test_case
|
||||
|
||||
|
||||
def require_torch_tpu(test_case):
|
||||
"""
|
||||
Decorator marking a test that requires a TPU (in PyTorch).
|
||||
|
||||
@@ -22,10 +22,12 @@ from .configuration_auto import (
|
||||
AutoConfig,
|
||||
BartConfig,
|
||||
BertConfig,
|
||||
BertGenerationConfig,
|
||||
CamembertConfig,
|
||||
CTRLConfig,
|
||||
DistilBertConfig,
|
||||
ElectraConfig,
|
||||
EncoderDecoderConfig,
|
||||
FlaubertConfig,
|
||||
FunnelConfig,
|
||||
GPT2Config,
|
||||
@@ -44,11 +46,13 @@ from .configuration_auto import (
|
||||
XLMConfig,
|
||||
XLMRobertaConfig,
|
||||
XLNetConfig,
|
||||
replace_list_option_in_docstrings,
|
||||
)
|
||||
from .configuration_utils import PretrainedConfig
|
||||
from .tokenization_albert import AlbertTokenizer
|
||||
from .tokenization_bart import BartTokenizer, BartTokenizerFast
|
||||
from .tokenization_bert import BertTokenizer, BertTokenizerFast
|
||||
from .tokenization_bert_generation import BertGenerationTokenizer
|
||||
from .tokenization_bert_japanese import BertJapaneseTokenizer
|
||||
from .tokenization_camembert import CamembertTokenizer
|
||||
from .tokenization_ctrl import CTRLTokenizer
|
||||
@@ -105,9 +109,12 @@ TOKENIZER_MAPPING = OrderedDict(
|
||||
(FlaubertConfig, (FlaubertTokenizer, None)),
|
||||
(XLMConfig, (XLMTokenizer, None)),
|
||||
(CTRLConfig, (CTRLTokenizer, None)),
|
||||
(BertGenerationConfig, (BertGenerationTokenizer, None)),
|
||||
]
|
||||
)
|
||||
|
||||
SLOW_TOKENIZER_MAPPING = {k: v[0] for k, v in TOKENIZER_MAPPING.items()}
|
||||
|
||||
|
||||
class AutoTokenizer:
|
||||
r""":class:`~transformers.AutoTokenizer` is a generic tokenizer class
|
||||
@@ -115,28 +122,6 @@ class AutoTokenizer:
|
||||
when created with the `AutoTokenizer.from_pretrained(pretrained_model_name_or_path)`
|
||||
class method.
|
||||
|
||||
The `from_pretrained()` method takes care of returning the correct tokenizer class instance
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
- `t5`: T5Tokenizer (T5 model)
|
||||
- `distilbert`: DistilBertTokenizer (DistilBert model)
|
||||
- `albert`: AlbertTokenizer (ALBERT model)
|
||||
- `camembert`: CamembertTokenizer (CamemBERT model)
|
||||
- `xlm-roberta`: XLMRobertaTokenizer (XLM-RoBERTa model)
|
||||
- `longformer`: LongformerTokenizer (AllenAI Longformer model)
|
||||
- `roberta`: RobertaTokenizer (RoBERTa model)
|
||||
- `bert`: BertTokenizer (Bert model)
|
||||
- `openai-gpt`: OpenAIGPTTokenizer (OpenAI GPT model)
|
||||
- `gpt2`: GPT2Tokenizer (OpenAI GPT-2 model)
|
||||
- `transfo-xl`: TransfoXLTokenizer (Transformer-XL model)
|
||||
- `xlnet`: XLNetTokenizer (XLNet model)
|
||||
- `xlm`: XLMTokenizer (XLM model)
|
||||
- `ctrl`: CTRLTokenizer (Salesforce CTRL model)
|
||||
- `electra`: ElectraTokenizer (Google ELECTRA model)
|
||||
- `funnel`: FunnelTokenizer (Funnel Transformer model)
|
||||
- `lxmert`: LxmertTokenizer (Lxmert model)
|
||||
|
||||
This class cannot be instantiated using `__init__()` (throw an error).
|
||||
"""
|
||||
|
||||
@@ -147,6 +132,7 @@ class AutoTokenizer:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings(SLOW_TOKENIZER_MAPPING)
|
||||
def from_pretrained(cls, pretrained_model_name_or_path, *inputs, **kwargs):
|
||||
r"""Instantiate one of the tokenizer classes of the library
|
||||
from a pre-trained model vocabulary.
|
||||
@@ -155,24 +141,7 @@ class AutoTokenizer:
|
||||
based on the `model_type` property of the config object, or when it's missing,
|
||||
falling back to using pattern matching on the `pretrained_model_name_or_path` string:
|
||||
|
||||
- `t5`: T5Tokenizer (T5 model)
|
||||
- `distilbert`: DistilBertTokenizer (DistilBert model)
|
||||
- `albert`: AlbertTokenizer (ALBERT model)
|
||||
- `camembert`: CamembertTokenizer (CamemBERT model)
|
||||
- `xlm-roberta`: XLMRobertaTokenizer (XLM-RoBERTa model)
|
||||
- `longformer`: LongformerTokenizer (AllenAI Longformer model)
|
||||
- `roberta`: RobertaTokenizer (RoBERTa model)
|
||||
- `bert-base-japanese`: BertJapaneseTokenizer (Bert model)
|
||||
- `bert`: BertTokenizer (Bert model)
|
||||
- `openai-gpt`: OpenAIGPTTokenizer (OpenAI GPT model)
|
||||
- `gpt2`: GPT2Tokenizer (OpenAI GPT-2 model)
|
||||
- `transfo-xl`: TransfoXLTokenizer (Transformer-XL model)
|
||||
- `xlnet`: XLNetTokenizer (XLNet model)
|
||||
- `xlm`: XLMTokenizer (XLM model)
|
||||
- `ctrl`: CTRLTokenizer (Salesforce CTRL model)
|
||||
- `electra`: ElectraTokenizer (Google ELECTRA model)
|
||||
- `funnel`: FunnelTokenizer (Funnel Transformer model)
|
||||
- `lxmert`: LxmertTokenizer (Lxmert model)
|
||||
List options
|
||||
|
||||
Params:
|
||||
pretrained_model_name_or_path: either:
|
||||
@@ -230,9 +199,19 @@ class AutoTokenizer:
|
||||
tokenizer_class_candidate = config.tokenizer_class
|
||||
tokenizer_class = globals().get(tokenizer_class_candidate)
|
||||
if tokenizer_class is None:
|
||||
raise ValueError("Tokenizer class {} does not exist or is not currently imported.")
|
||||
raise ValueError(
|
||||
"Tokenizer class {} does not exist or is not currently imported.".format(tokenizer_class_candidate)
|
||||
)
|
||||
return tokenizer_class.from_pretrained(pretrained_model_name_or_path, *inputs, **kwargs)
|
||||
|
||||
# if model is an encoder decoder, the encoder tokenizer class is used by default
|
||||
if isinstance(config, EncoderDecoderConfig):
|
||||
if type(config.decoder) is not type(config.encoder): # noqa: E721
|
||||
logger.warn(
|
||||
f"The encoder model config class: {config.encoder.__class__} is different from the decoder model config class: {config.decoder.__class}. It is not recommended to use the `AutoTokenizer.from_pretrained(..)` method in this case. Please use the encoder and decoder specific tokenizer classes."
|
||||
)
|
||||
config = config.encoder
|
||||
|
||||
for config_class, (tokenizer_class_py, tokenizer_class_fast) in TOKENIZER_MAPPING.items():
|
||||
if isinstance(config, config_class):
|
||||
if tokenizer_class_fast and use_fast:
|
||||
|
||||
@@ -0,0 +1,156 @@
|
||||
# coding=utf-8
|
||||
# Copyright (c) 2020, NVIDIA CORPORATION. All rights reserved.
|
||||
#
|
||||
# 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 class for model BertGeneration."""
|
||||
|
||||
|
||||
import os
|
||||
from shutil import copyfile
|
||||
from typing import List
|
||||
|
||||
from .tokenization_utils import PreTrainedTokenizer
|
||||
from .utils import logging
|
||||
|
||||
|
||||
logger = logging.get_logger(__name__)
|
||||
|
||||
VOCAB_FILES_NAMES = {"vocab_file": "spiece.model"}
|
||||
|
||||
tokenizer_url = (
|
||||
"https://s3.amazonaws.com/models.huggingface.co/bert/google/bert_for_seq_generation_L-24_bbc_encoder/spiece.model"
|
||||
)
|
||||
|
||||
|
||||
class BertGenerationTokenizer(PreTrainedTokenizer):
|
||||
"""
|
||||
Constructs a BertGenerationTokenizer tokenizer. Based on `SentencePiece <https://github.com/google/sentencepiece>`__ .
|
||||
|
||||
This tokenizer inherits from :class:`~transformers.PreTrainedTokenizer` which contains most of the methods. Users
|
||||
should refer to the superclass for more information regarding methods.
|
||||
|
||||
Args:
|
||||
vocab_file (:obj:`string`):
|
||||
`SentencePiece <https://github.com/google/sentencepiece>`__ file (generally has a `.spm` extension) that
|
||||
contains the vocabulary necessary to instantiate a tokenizer.
|
||||
eos_token (:obj:`string`, `optional`, defaults to :obj:`"</s>"`):
|
||||
The end of sequence token.
|
||||
bos_token (:obj:`string`, `optional`, defaults to :obj:`"<s>"`):
|
||||
The begin of sequence token.
|
||||
unk_token (:obj:`string`, `optional`, defaults to :obj:`"<unk>"`):
|
||||
The unknown token. A token that is not in the vocabulary cannot be converted to an ID and is set to be this
|
||||
token instead.
|
||||
pad_token (:obj:`string`, `optional`, defaults to :obj:`"<pad>"`):
|
||||
The token used for padding, for example when batching sequences of different lengths.
|
||||
"""
|
||||
|
||||
vocab_files_names = VOCAB_FILES_NAMES
|
||||
prefix_tokens: List[int] = []
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vocab_file,
|
||||
bos_token="<s>",
|
||||
eos_token="</s>",
|
||||
unk_token="<unk>",
|
||||
pad_token="<pad>",
|
||||
sep_token="<::::>",
|
||||
**kwargs
|
||||
):
|
||||
# Add extra_ids to the special token list
|
||||
super().__init__(
|
||||
bos_token=bos_token,
|
||||
eos_token=eos_token,
|
||||
unk_token=unk_token,
|
||||
pad_token=pad_token,
|
||||
sep_token=sep_token,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
try:
|
||||
import sentencepiece as spm
|
||||
except ImportError:
|
||||
logger.warning(
|
||||
"You need to install SentencePiece to use T5Tokenizer:"
|
||||
"https://github.com/google/sentencepiece"
|
||||
"pip install sentencepiece"
|
||||
)
|
||||
raise
|
||||
|
||||
self.vocab_file = vocab_file
|
||||
|
||||
self.sp_model = spm.SentencePieceProcessor()
|
||||
self.sp_model.Load(vocab_file)
|
||||
|
||||
@property
|
||||
def vocab_size(self):
|
||||
return self.sp_model.get_piece_size()
|
||||
|
||||
def get_vocab(self):
|
||||
vocab = {self.convert_ids_to_tokens(i): i for i in range(self.vocab_size)}
|
||||
vocab.update(self.added_tokens_encoder)
|
||||
return vocab
|
||||
|
||||
def __getstate__(self):
|
||||
state = self.__dict__.copy()
|
||||
state["sp_model"] = None
|
||||
return state
|
||||
|
||||
def __setstate__(self, d):
|
||||
self.__dict__ = d
|
||||
try:
|
||||
import sentencepiece as spm
|
||||
except ImportError:
|
||||
logger.warning(
|
||||
"You need to install SentencePiece to use BertGenerationTokenizer: https://github.com/google/sentencepiece"
|
||||
"pip install sentencepiece"
|
||||
)
|
||||
raise
|
||||
self.sp_model = spm.SentencePieceProcessor()
|
||||
self.sp_model.Load(self.vocab_file)
|
||||
|
||||
def _tokenize(self, text, sample=False):
|
||||
"""Take as input a string and return a list of strings (tokens) for words/sub-words"""
|
||||
if not sample:
|
||||
pieces = self.sp_model.EncodeAsPieces(text)
|
||||
else:
|
||||
pieces = self.sp_model.SampleEncodeAsPieces(text, 64, 0.1)
|
||||
return pieces
|
||||
|
||||
def _convert_token_to_id(self, token):
|
||||
""" Converts a token (str) in an id using the vocab. """
|
||||
return self.sp_model.piece_to_id(token)
|
||||
|
||||
def _convert_id_to_token(self, index):
|
||||
"""Converts an index (integer) in a token (str) using the vocab."""
|
||||
token = self.sp_model.IdToPiece(index)
|
||||
return token
|
||||
|
||||
def convert_tokens_to_string(self, tokens):
|
||||
""" Converts a sequence of tokens (string) in a single string. """
|
||||
out_string = self.sp_model.decode_pieces(tokens)
|
||||
return out_string
|
||||
|
||||
def save_vocabulary(self, save_directory):
|
||||
"""Save the sentencepiece vocabulary (copy original file) and special tokens file
|
||||
to a directory.
|
||||
"""
|
||||
if not os.path.isdir(save_directory):
|
||||
logger.error("Vocabulary path ({}) should be a directory".format(save_directory))
|
||||
return
|
||||
out_vocab_file = os.path.join(save_directory, VOCAB_FILES_NAMES["vocab_file"])
|
||||
|
||||
if os.path.abspath(self.vocab_file) != os.path.abspath(out_vocab_file):
|
||||
copyfile(self.vocab_file, out_vocab_file)
|
||||
|
||||
return (out_vocab_file,)
|
||||
@@ -22,7 +22,6 @@ from typing import List, Optional
|
||||
import sentencepiece as spm
|
||||
|
||||
from .tokenization_utils import PreTrainedTokenizer
|
||||
from .tokenization_xlnet import SPIECE_UNDERLINE
|
||||
from .utils import logging
|
||||
|
||||
|
||||
@@ -47,6 +46,8 @@ SHARED_MODEL_IDENTIFIERS = [
|
||||
"Musixmatch/umberto-wikipedia-uncased-v1",
|
||||
]
|
||||
|
||||
SPIECE_UNDERLINE = "▁"
|
||||
|
||||
|
||||
class CamembertTokenizer(PreTrainedTokenizer):
|
||||
"""
|
||||
@@ -253,7 +254,7 @@ class CamembertTokenizer(PreTrainedTokenizer):
|
||||
import sentencepiece as spm
|
||||
except ImportError:
|
||||
logger.warning(
|
||||
"You need to install SentencePiece to use AlbertTokenizer: https://github.com/google/sentencepiece"
|
||||
"You need to install SentencePiece to use CamembertTokenizer: https://github.com/google/sentencepiece"
|
||||
"pip install sentencepiece"
|
||||
)
|
||||
raise
|
||||
|
||||
@@ -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
|
||||
@@ -27,8 +27,6 @@ from .utils import logging
|
||||
|
||||
logger = logging.get_logger(__name__)
|
||||
|
||||
SPIECE_UNDERLINE = "▁"
|
||||
|
||||
####################################################
|
||||
# Mapping from the keyword arguments names of Tokenizer `__init__`
|
||||
# to file names for serializing Tokenizer instances
|
||||
@@ -98,8 +96,6 @@ class T5Tokenizer(PreTrainedTokenizer):
|
||||
max_model_input_sizes = PRETRAINED_POSITIONAL_EMBEDDINGS_SIZES
|
||||
model_input_names = ["attention_mask"]
|
||||
|
||||
prefix_tokens: List[int] = []
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vocab_file,
|
||||
@@ -212,10 +208,10 @@ class T5Tokenizer(PreTrainedTokenizer):
|
||||
"""
|
||||
token_ids_0 = self._add_eos_if_not_present(token_ids_0)
|
||||
if token_ids_1 is None:
|
||||
return self.prefix_tokens + token_ids_0
|
||||
return token_ids_0
|
||||
else:
|
||||
token_ids_1 = self._add_eos_if_not_present(token_ids_1)
|
||||
return self.prefix_tokens + token_ids_0 + token_ids_1
|
||||
return token_ids_0 + token_ids_1
|
||||
|
||||
def __getstate__(self):
|
||||
state = self.__dict__.copy()
|
||||
@@ -345,7 +341,6 @@ class T5Tokenizer(PreTrainedTokenizer):
|
||||
"""
|
||||
if max_length is None:
|
||||
max_length = self.max_len
|
||||
self.prefix_tokens = []
|
||||
model_inputs = self(
|
||||
src_texts,
|
||||
add_special_tokens=True,
|
||||
@@ -360,8 +355,6 @@ class T5Tokenizer(PreTrainedTokenizer):
|
||||
# Process tgt_texts
|
||||
if max_target_length is None:
|
||||
max_target_length = max_length
|
||||
# set prefix_tokens for target text
|
||||
self.prefix_tokens = [self.pad_token_id]
|
||||
labels_and_decoder_mask = self(
|
||||
tgt_texts,
|
||||
add_special_tokens=True,
|
||||
@@ -372,5 +365,4 @@ class T5Tokenizer(PreTrainedTokenizer):
|
||||
**kwargs,
|
||||
)
|
||||
model_inputs["labels"] = labels_and_decoder_mask["input_ids"]
|
||||
self.prefix_tokens = []
|
||||
return model_inputs
|
||||
@@ -79,37 +79,37 @@ PRETRAINED_INIT_CONFIGURATION = {
|
||||
"xlm-mlm-en-2048": {"do_lowercase_and_remove_accent": True},
|
||||
"xlm-mlm-ende-1024": {
|
||||
"do_lowercase_and_remove_accent": True,
|
||||
"id2lang": {"0": "de", "1": "en"},
|
||||
"id2lang": {0: "de", 1: "en"},
|
||||
"lang2id": {"de": 0, "en": 1},
|
||||
},
|
||||
"xlm-mlm-enfr-1024": {
|
||||
"do_lowercase_and_remove_accent": True,
|
||||
"id2lang": {"0": "en", "1": "fr"},
|
||||
"id2lang": {0: "en", 1: "fr"},
|
||||
"lang2id": {"en": 0, "fr": 1},
|
||||
},
|
||||
"xlm-mlm-enro-1024": {
|
||||
"do_lowercase_and_remove_accent": True,
|
||||
"id2lang": {"0": "en", "1": "ro"},
|
||||
"id2lang": {0: "en", 1: "ro"},
|
||||
"lang2id": {"en": 0, "ro": 1},
|
||||
},
|
||||
"xlm-mlm-tlm-xnli15-1024": {
|
||||
"do_lowercase_and_remove_accent": True,
|
||||
"id2lang": {
|
||||
"0": "ar",
|
||||
"1": "bg",
|
||||
"2": "de",
|
||||
"3": "el",
|
||||
"4": "en",
|
||||
"5": "es",
|
||||
"6": "fr",
|
||||
"7": "hi",
|
||||
"8": "ru",
|
||||
"9": "sw",
|
||||
"10": "th",
|
||||
"11": "tr",
|
||||
"12": "ur",
|
||||
"13": "vi",
|
||||
"14": "zh",
|
||||
0: "ar",
|
||||
1: "bg",
|
||||
2: "de",
|
||||
3: "el",
|
||||
4: "en",
|
||||
5: "es",
|
||||
6: "fr",
|
||||
7: "hi",
|
||||
8: "ru",
|
||||
9: "sw",
|
||||
10: "th",
|
||||
11: "tr",
|
||||
12: "ur",
|
||||
13: "vi",
|
||||
14: "zh",
|
||||
},
|
||||
"lang2id": {
|
||||
"ar": 0,
|
||||
@@ -132,21 +132,21 @@ PRETRAINED_INIT_CONFIGURATION = {
|
||||
"xlm-mlm-xnli15-1024": {
|
||||
"do_lowercase_and_remove_accent": True,
|
||||
"id2lang": {
|
||||
"0": "ar",
|
||||
"1": "bg",
|
||||
"2": "de",
|
||||
"3": "el",
|
||||
"4": "en",
|
||||
"5": "es",
|
||||
"6": "fr",
|
||||
"7": "hi",
|
||||
"8": "ru",
|
||||
"9": "sw",
|
||||
"10": "th",
|
||||
"11": "tr",
|
||||
"12": "ur",
|
||||
"13": "vi",
|
||||
"14": "zh",
|
||||
0: "ar",
|
||||
1: "bg",
|
||||
2: "de",
|
||||
3: "el",
|
||||
4: "en",
|
||||
5: "es",
|
||||
6: "fr",
|
||||
7: "hi",
|
||||
8: "ru",
|
||||
9: "sw",
|
||||
10: "th",
|
||||
11: "tr",
|
||||
12: "ur",
|
||||
13: "vi",
|
||||
14: "zh",
|
||||
},
|
||||
"lang2id": {
|
||||
"ar": 0,
|
||||
@@ -168,34 +168,34 @@ PRETRAINED_INIT_CONFIGURATION = {
|
||||
},
|
||||
"xlm-clm-enfr-1024": {
|
||||
"do_lowercase_and_remove_accent": True,
|
||||
"id2lang": {"0": "en", "1": "fr"},
|
||||
"id2lang": {0: "en", 1: "fr"},
|
||||
"lang2id": {"en": 0, "fr": 1},
|
||||
},
|
||||
"xlm-clm-ende-1024": {
|
||||
"do_lowercase_and_remove_accent": True,
|
||||
"id2lang": {"0": "de", "1": "en"},
|
||||
"id2lang": {0: "de", 1: "en"},
|
||||
"lang2id": {"de": 0, "en": 1},
|
||||
},
|
||||
"xlm-mlm-17-1280": {
|
||||
"do_lowercase_and_remove_accent": False,
|
||||
"id2lang": {
|
||||
"0": "ar",
|
||||
"1": "de",
|
||||
"2": "en",
|
||||
"3": "es",
|
||||
"4": "fr",
|
||||
"5": "hi",
|
||||
"6": "it",
|
||||
"7": "ja",
|
||||
"8": "ko",
|
||||
"9": "nl",
|
||||
"10": "pl",
|
||||
"11": "pt",
|
||||
"12": "ru",
|
||||
"13": "sv",
|
||||
"14": "tr",
|
||||
"15": "vi",
|
||||
"16": "zh",
|
||||
0: "ar",
|
||||
1: "de",
|
||||
2: "en",
|
||||
3: "es",
|
||||
4: "fr",
|
||||
5: "hi",
|
||||
6: "it",
|
||||
7: "ja",
|
||||
8: "ko",
|
||||
9: "nl",
|
||||
10: "pl",
|
||||
11: "pt",
|
||||
12: "ru",
|
||||
13: "sv",
|
||||
14: "tr",
|
||||
15: "vi",
|
||||
16: "zh",
|
||||
},
|
||||
"lang2id": {
|
||||
"ar": 0,
|
||||
@@ -220,106 +220,106 @@ PRETRAINED_INIT_CONFIGURATION = {
|
||||
"xlm-mlm-100-1280": {
|
||||
"do_lowercase_and_remove_accent": False,
|
||||
"id2lang": {
|
||||
"0": "af",
|
||||
"1": "als",
|
||||
"2": "am",
|
||||
"3": "an",
|
||||
"4": "ang",
|
||||
"5": "ar",
|
||||
"6": "arz",
|
||||
"7": "ast",
|
||||
"8": "az",
|
||||
"9": "bar",
|
||||
"10": "be",
|
||||
"11": "bg",
|
||||
"12": "bn",
|
||||
"13": "br",
|
||||
"14": "bs",
|
||||
"15": "ca",
|
||||
"16": "ceb",
|
||||
"17": "ckb",
|
||||
"18": "cs",
|
||||
"19": "cy",
|
||||
"20": "da",
|
||||
"21": "de",
|
||||
"22": "el",
|
||||
"23": "en",
|
||||
"24": "eo",
|
||||
"25": "es",
|
||||
"26": "et",
|
||||
"27": "eu",
|
||||
"28": "fa",
|
||||
"29": "fi",
|
||||
"30": "fr",
|
||||
"31": "fy",
|
||||
"32": "ga",
|
||||
"33": "gan",
|
||||
"34": "gl",
|
||||
"35": "gu",
|
||||
"36": "he",
|
||||
"37": "hi",
|
||||
"38": "hr",
|
||||
"39": "hu",
|
||||
"40": "hy",
|
||||
"41": "ia",
|
||||
"42": "id",
|
||||
"43": "is",
|
||||
"44": "it",
|
||||
"45": "ja",
|
||||
"46": "jv",
|
||||
"47": "ka",
|
||||
"48": "kk",
|
||||
"49": "kn",
|
||||
"50": "ko",
|
||||
"51": "ku",
|
||||
"52": "la",
|
||||
"53": "lb",
|
||||
"54": "lt",
|
||||
"55": "lv",
|
||||
"56": "mk",
|
||||
"57": "ml",
|
||||
"58": "mn",
|
||||
"59": "mr",
|
||||
"60": "ms",
|
||||
"61": "my",
|
||||
"62": "nds",
|
||||
"63": "ne",
|
||||
"64": "nl",
|
||||
"65": "nn",
|
||||
"66": "no",
|
||||
"67": "oc",
|
||||
"68": "pl",
|
||||
"69": "pt",
|
||||
"70": "ro",
|
||||
"71": "ru",
|
||||
"72": "scn",
|
||||
"73": "sco",
|
||||
"74": "sh",
|
||||
"75": "si",
|
||||
"76": "simple",
|
||||
"77": "sk",
|
||||
"78": "sl",
|
||||
"79": "sq",
|
||||
"80": "sr",
|
||||
"81": "sv",
|
||||
"82": "sw",
|
||||
"83": "ta",
|
||||
"84": "te",
|
||||
"85": "th",
|
||||
"86": "tl",
|
||||
"87": "tr",
|
||||
"88": "tt",
|
||||
"89": "uk",
|
||||
"90": "ur",
|
||||
"91": "uz",
|
||||
"92": "vi",
|
||||
"93": "war",
|
||||
"94": "wuu",
|
||||
"95": "yi",
|
||||
"96": "zh",
|
||||
"97": "zh_classical",
|
||||
"98": "zh_min_nan",
|
||||
"99": "zh_yue",
|
||||
0: "af",
|
||||
1: "als",
|
||||
2: "am",
|
||||
3: "an",
|
||||
4: "ang",
|
||||
5: "ar",
|
||||
6: "arz",
|
||||
7: "ast",
|
||||
8: "az",
|
||||
9: "bar",
|
||||
10: "be",
|
||||
11: "bg",
|
||||
12: "bn",
|
||||
13: "br",
|
||||
14: "bs",
|
||||
15: "ca",
|
||||
16: "ceb",
|
||||
17: "ckb",
|
||||
18: "cs",
|
||||
19: "cy",
|
||||
20: "da",
|
||||
21: "de",
|
||||
22: "el",
|
||||
23: "en",
|
||||
24: "eo",
|
||||
25: "es",
|
||||
26: "et",
|
||||
27: "eu",
|
||||
28: "fa",
|
||||
29: "fi",
|
||||
30: "fr",
|
||||
31: "fy",
|
||||
32: "ga",
|
||||
33: "gan",
|
||||
34: "gl",
|
||||
35: "gu",
|
||||
36: "he",
|
||||
37: "hi",
|
||||
38: "hr",
|
||||
39: "hu",
|
||||
40: "hy",
|
||||
41: "ia",
|
||||
42: "id",
|
||||
43: "is",
|
||||
44: "it",
|
||||
45: "ja",
|
||||
46: "jv",
|
||||
47: "ka",
|
||||
48: "kk",
|
||||
49: "kn",
|
||||
50: "ko",
|
||||
51: "ku",
|
||||
52: "la",
|
||||
53: "lb",
|
||||
54: "lt",
|
||||
55: "lv",
|
||||
56: "mk",
|
||||
57: "ml",
|
||||
58: "mn",
|
||||
59: "mr",
|
||||
60: "ms",
|
||||
61: "my",
|
||||
62: "nds",
|
||||
63: "ne",
|
||||
64: "nl",
|
||||
65: "nn",
|
||||
66: "no",
|
||||
67: "oc",
|
||||
68: "pl",
|
||||
69: "pt",
|
||||
70: "ro",
|
||||
71: "ru",
|
||||
72: "scn",
|
||||
73: "sco",
|
||||
74: "sh",
|
||||
75: "si",
|
||||
76: "simple",
|
||||
77: "sk",
|
||||
78: "sl",
|
||||
79: "sq",
|
||||
80: "sr",
|
||||
81: "sv",
|
||||
82: "sw",
|
||||
83: "ta",
|
||||
84: "te",
|
||||
85: "th",
|
||||
86: "tl",
|
||||
87: "tr",
|
||||
88: "tt",
|
||||
89: "uk",
|
||||
90: "ur",
|
||||
91: "uz",
|
||||
92: "vi",
|
||||
93: "war",
|
||||
94: "wuu",
|
||||
95: "yi",
|
||||
96: "zh",
|
||||
97: "zh_classical",
|
||||
98: "zh_min_nan",
|
||||
99: "zh_yue",
|
||||
},
|
||||
"lang2id": {
|
||||
"af": 0,
|
||||
|
||||
+30
-23
@@ -20,7 +20,7 @@ from torch.utils.data.sampler import RandomSampler, Sampler, SequentialSampler
|
||||
from tqdm.auto import tqdm, trange
|
||||
|
||||
from .data.data_collator import DataCollator, DataCollatorWithPadding, default_data_collator
|
||||
from .file_utils import is_nlp_available, is_torch_tpu_available
|
||||
from .file_utils import is_datasets_available, is_torch_tpu_available
|
||||
from .integrations import (
|
||||
default_hp_search_backend,
|
||||
is_comet_available,
|
||||
@@ -65,8 +65,8 @@ else:
|
||||
_use_native_amp = True
|
||||
from torch.cuda.amp import autocast
|
||||
|
||||
if is_nlp_available():
|
||||
import nlp
|
||||
if is_datasets_available():
|
||||
import datasets
|
||||
|
||||
if is_torch_tpu_available():
|
||||
import torch_xla.core.xla_model as xm
|
||||
@@ -179,10 +179,10 @@ class Trainer:
|
||||
:obj:`eval_dataset`. Will default to :func:`~transformers.default_data_collator` if no ``tokenizer`` is
|
||||
provided, an instance of :func:`~transformers.DataCollatorWithPadding` otherwise.
|
||||
train_dataset (:obj:`torch.utils.data.dataset.Dataset`, `optional`):
|
||||
The dataset to use for training. If it is an :obj:`nlp.Dataset`, columns not accepted by the
|
||||
The dataset to use for training. If it is an :obj:`datasets.Dataset`, columns not accepted by the
|
||||
``model.forward()`` method are automatically removed.
|
||||
eval_dataset (:obj:`torch.utils.data.dataset.Dataset`, `optional`):
|
||||
The dataset to use for evaluation. If it is an :obj:`nlp.Dataset`, columns not accepted by the
|
||||
The dataset to use for evaluation. If it is an :obj:`datasets.Dataset`, columns not accepted by the
|
||||
``model.forward()`` method are automatically removed.
|
||||
tokenizer (:class:`PreTrainedTokenizerBase`, `optional`):
|
||||
The tokenizer used to preprocess the data. If provided, will be used to automatically pad the inputs the
|
||||
@@ -280,10 +280,10 @@ class Trainer:
|
||||
FutureWarning,
|
||||
)
|
||||
|
||||
if is_nlp_available():
|
||||
if isinstance(train_dataset, nlp.Dataset):
|
||||
if is_datasets_available():
|
||||
if isinstance(train_dataset, datasets.Dataset):
|
||||
self._remove_unused_columns(self.train_dataset, description="training")
|
||||
if isinstance(eval_dataset, nlp.Dataset):
|
||||
if isinstance(eval_dataset, datasets.Dataset):
|
||||
self._remove_unused_columns(self.eval_dataset, description="evaluation")
|
||||
|
||||
self.global_step = None
|
||||
@@ -294,7 +294,7 @@ class Trainer:
|
||||
self.hp_search_backend = None
|
||||
self.use_tune_checkpoints = False
|
||||
|
||||
def _remove_unused_columns(self, dataset: "nlp.Dataset", description: Optional[str] = None):
|
||||
def _remove_unused_columns(self, dataset: "datasets.Dataset", description: Optional[str] = None):
|
||||
if not self.args.remove_unused_columns:
|
||||
return
|
||||
# Inspect model forward signature to keep only the arguments it accepts.
|
||||
@@ -364,12 +364,12 @@ class Trainer:
|
||||
|
||||
Args:
|
||||
eval_dataset (:obj:`torch.utils.data.dataset.Dataset`, `optional`):
|
||||
If provided, will override :obj:`self.eval_dataset`. If it is an :obj:`nlp.Dataset`, columns not
|
||||
If provided, will override :obj:`self.eval_dataset`. If it is an :obj:`datasets.Dataset`, columns not
|
||||
accepted by the ``model.forward()`` method are automatically removed.
|
||||
"""
|
||||
if eval_dataset is None and self.eval_dataset is None:
|
||||
raise ValueError("Trainer: evaluation requires an eval_dataset.")
|
||||
elif eval_dataset is not None and is_nlp_available() and isinstance(eval_dataset, nlp.Dataset):
|
||||
elif eval_dataset is not None and is_datasets_available() and isinstance(eval_dataset, datasets.Dataset):
|
||||
self._remove_unused_columns(eval_dataset, description="evaluation")
|
||||
eval_dataset = eval_dataset if eval_dataset is not None else self.eval_dataset
|
||||
eval_sampler = self._get_eval_sampler(eval_dataset)
|
||||
@@ -393,10 +393,10 @@ class Trainer:
|
||||
|
||||
Args:
|
||||
eval_dataset (:obj:`torch.utils.data.dataset.Dataset`, `optional`):
|
||||
The test dataset to use. If it is an :obj:`nlp.Dataset`, columns not accepted by the
|
||||
The test dataset to use. If it is an :obj:`datasets.Dataset`, columns not accepted by the
|
||||
``model.forward()`` method are automatically removed.
|
||||
"""
|
||||
if is_nlp_available() and isinstance(test_dataset, nlp.Dataset):
|
||||
if is_datasets_available() and isinstance(test_dataset, datasets.Dataset):
|
||||
self._remove_unused_columns(test_dataset, description="test")
|
||||
test_sampler = self._get_eval_sampler(test_dataset)
|
||||
|
||||
@@ -1024,15 +1024,9 @@ class Trainer:
|
||||
|
||||
if self.args.fp16 and _use_native_amp:
|
||||
with autocast():
|
||||
outputs = model(**inputs)
|
||||
loss = outputs[0]
|
||||
loss = self.compute_loss(model, inputs)
|
||||
else:
|
||||
outputs = model(**inputs)
|
||||
# We don't use .loss here since the model may return tuples instead of ModelOutput.
|
||||
loss = outputs[0]
|
||||
|
||||
if self.args.past_index >= 0:
|
||||
self._past = outputs[self.args.past_index]
|
||||
loss = self.compute_loss(model, inputs)
|
||||
|
||||
if self.args.n_gpu > 1:
|
||||
loss = loss.mean() # mean() to average on multi-gpu parallel training
|
||||
@@ -1050,6 +1044,19 @@ class Trainer:
|
||||
|
||||
return loss.detach()
|
||||
|
||||
def compute_loss(self, model, inputs):
|
||||
"""
|
||||
How the loss is computed by Trainer. By default, all models return the loss in the first element.
|
||||
|
||||
Subclass and override for custom behavior.
|
||||
"""
|
||||
outputs = model(**inputs)
|
||||
# Save past state if it exists
|
||||
if self.args.past_index >= 0:
|
||||
self._past = outputs[self.args.past_index]
|
||||
# We don't use .loss here since the model may return tuples instead of ModelOutput.
|
||||
return outputs[0]
|
||||
|
||||
def is_local_master(self) -> bool:
|
||||
"""
|
||||
Whether or not this process is the local (e.g., on one machine if training in a distributed fashion on
|
||||
@@ -1200,7 +1207,7 @@ class Trainer:
|
||||
|
||||
Args:
|
||||
eval_dataset (:obj:`Dataset`, `optional`):
|
||||
Pass a dataset if you wish to override :obj:`self.eval_dataset`. If it is an :obj:`nlp.Dataset`,
|
||||
Pass a dataset if you wish to override :obj:`self.eval_dataset`. If it is an :obj:`datasets.Dataset`,
|
||||
columns not accepted by the ``model.forward()`` method are automatically removed.
|
||||
|
||||
Returns:
|
||||
@@ -1227,7 +1234,7 @@ class Trainer:
|
||||
|
||||
Args:
|
||||
test_dataset (:obj:`Dataset`):
|
||||
Dataset to run the predictions on. If it is an :obj:`nlp.Dataset`, columns not accepted by the
|
||||
Dataset to run the predictions on. If it is an :obj:`datasets.Dataset`, columns not accepted by the
|
||||
``model.forward()`` method are automatically removed.
|
||||
|
||||
Returns:
|
||||
|
||||
@@ -19,9 +19,6 @@
|
||||
# In this template, replace all the XXX (various casings) with your model name
|
||||
####################################################
|
||||
|
||||
|
||||
import logging
|
||||
|
||||
import tensorflow as tf
|
||||
|
||||
from .configuration_xxx import XxxConfig
|
||||
@@ -47,12 +44,14 @@ from .modeling_tf_utils import (
|
||||
TFSequenceClassificationLoss,
|
||||
TFTokenClassificationLoss,
|
||||
get_initializer,
|
||||
keras_serializable,
|
||||
shape_list,
|
||||
)
|
||||
from .tokenization_utils import BatchEncoding
|
||||
from .utils import logging
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
logger = logging.get_logger(__name__)
|
||||
|
||||
_CONFIG_FOR_DOC = "XXXConfig"
|
||||
_TOKENIZER_FOR_DOC = "XxxTokenizer"
|
||||
@@ -115,15 +114,20 @@ class TFXxxLayer(tf.keras.layers.Layer):
|
||||
# The full model without a specific pretrained or finetuning head is
|
||||
# provided as a tf.keras.layers.Layer usually called "TFXxxMainLayer"
|
||||
####################################################
|
||||
@keras_serializable
|
||||
class TFXxxMainLayer(tf.keras.layers.Layer):
|
||||
def __init__(self, config, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
|
||||
def _resize_token_embeddings(self, new_num_tokens):
|
||||
raise NotImplementedError # Not implemented yet in the library fr TF 2.0 models
|
||||
def get_input_embeddings(self):
|
||||
return self.embeddings
|
||||
|
||||
def set_input_embeddings(self, value):
|
||||
self.embeddings.word_embeddings = value
|
||||
self.embeddings.vocab_size = value.shape[0]
|
||||
|
||||
def _prune_heads(self, heads_to_prune):
|
||||
raise NotImplementedError # Not implemented yet in the library fr TF 2.0 models
|
||||
raise NotImplementedError # Not implemented yet in the library for TF 2.0 models
|
||||
|
||||
def call(
|
||||
self,
|
||||
|
||||
+251
@@ -0,0 +1,251 @@
|
||||
<doc id="12" url="https://en.wikipedia.org/wiki?curid=12" title="Anarchism">
|
||||
Anarchism
|
||||
|
||||
Anarchism is a political philosophy and movement that rejects all involuntary, coercive forms of hierarchy. It radically calls for the abolition of the state which it holds to be undesirable, unnecessary, and harmful.
|
||||
|
||||
The history of anarchism stretches back to prehistory, when humans lived in anarchistic societies long before the establishment of formal states, realms or empires. With the rise of organised hierarchical bodies, skepticism toward authority also rose, but it was not until the 19th century that a self-conscious political movement emerged. During the latter half of the 19th and the first decades of the 20th century, the anarchist movement flourished in most parts of the world and had a significant role in worker's struggles for emancipation. Various anarchist schools of thought formed during this period.
|
||||
|
||||
Anarchists took part in several revolutions, most notably in the Spanish Civil War, where they were crushed along with the alliance to restore the Second Republic by the fascist forces of the Nationalist faction and its foreign allies in Nazi Germany, Fascist Italy, Portuguese Dictatorship and the Catholic Church in 1939, marking the end of the classical era of anarchism. In the last decades of the 20th century and into the 21st century, the anarchist movement has been resurgent once more.
|
||||
|
||||
Anarchism employs various tactics in order to meet its ideal ends; these can be broadly separated into revolutionary and evolutionary tactics. There is significant overlap between the two, which are merely descriptive. Revolutionary tactics aim to bring down authority and state, and have taken a violent turn in the past. Evolutionary tactics aim to prefigure what an anarchist society would be like. Anarchist thought, criticism, and praxis has played a part in diverse areas of human society.
|
||||
|
||||
The etymological origin of "anarchism" is from the Ancient Greek "anarkhia", meaning "without a ruler", composed of the prefix "an-" (i.e. "without") and the word "arkhos" (i.e. "leader" or "ruler"). The suffix "-ism" denotes the ideological current that favours anarchy. "Anarchism" appears in English from 1642 as "anarchisme" and "anarchy" from 1539. Various factions within the French Revolution labelled their opponents as "anarchists", although few such accused shared many views with later anarchists. Many revolutionaries of the 19th century such as William Godwin (1756–1836) and Wilhelm Weitling (1808–1871) would contribute to the anarchist doctrines of the next generation, but they did not use "anarchist" or "anarchism" in describing themselves or their beliefs.
|
||||
|
||||
The first political philosopher to call himself an "anarchist" () was Pierre-Joseph Proudhon (1809–1865), marking the formal birth of anarchism in the mid-19th century. Since the 1890s and beginning in France, "libertarianism" has often been used as a synonym for anarchism and its use as a synonym is still common outside the United States. On the other hand, some use "libertarianism" to refer to individualistic free-market philosophy only, referring to free-market anarchism as "libertarian anarchism".
|
||||
|
||||
While opposition to the state is central to anarchist thought, defining anarchism is not an easy task as there is a lot of discussion among scholars and anarchists on the matter and various currents perceive anarchism slightly differently. Hence, it might be true to say that anarchism is a cluster of political philosophies opposing authority and hierarchical organization (including the state, capitalism, nationalism and all associated institutions) in the conduct of all human relations in favour of a society based on voluntary association, on freedom and on decentralisation, but this definition has the same shortcomings as the definition based on etymology (which is simply a negation of a ruler), or based on anti-statism (anarchism is much more than that) or even the anti-authoritarian (which is an "a posteriori" conclusion). Nonetheless, major elements of the definition of anarchism include the following:
|
||||
|
||||
During the prehistoric era of mankind, an established authority did not exist. It was after the creation of towns and cities that institutions of authority were established and anarchistic ideas espoused as a reaction. Most notable precursors to anarchism in the ancient world were in China and Greece. In China, philosophical anarchism (i.e. the discussion on the legitimacy of the state) was delineated by Taoist philosophers Zhuang Zhou and Laozi.
|
||||
|
||||
Likewise, anarchic attitudes were articulated by tragedians and philosophers in Greece. Aeschylus and Sophocles used the myth of Antigone to illustrate the conflict between rules set by the state and personal autonomy. Socrates questioned Athenian authorities constantly and insisted to the right of individual freedom of consciousness. Cynics dismissed human law ("nomos") and associated authorities while trying to live according to nature ("physis"). Stoics were supportive of a society based on unofficial and friendly relations among its citizens without the presence of a state.
|
||||
|
||||
During the Middle Ages, there was no anarchistic activity except some ascetic religious movements in the Muslim world or in Christian Europe. This kind of tradition later gave birth to religious anarchism. In the Sasanian Empire, Mazdak called for an egalitarian society and the abolition of monarchy, only to be soon executed by Emperor Kavad I.
|
||||
|
||||
In Basra, religious sects preached against the state. In Europe, various sects developed anti-state and libertarian tendencies. Libertarian ideas further emerged during the Renaissance with the spread of reasoning and humanism through Europe. Novelists fictionalised ideal societies that were based not on coercion but voluntarism. The Enlightenment further pushed towards anarchism with the optimism for social progress.
|
||||
|
||||
During the French Revolution, partisan groups such as the Enragés and the saw a turning point in the fermentation of anti-state and federalist sentiments. The first anarchist currents developed throughout the 18th century—William Godwin espoused philosophical anarchism in England, morally delegitimizing the state, Max Stirner's thinking paved the way to individualism, and Pierre-Joseph Proudhon's theory of mutualism found fertile soil in France. This era of classical anarchism lasted until the end of the Spanish Civil War of 1936 and is considered the golden age of anarchism.
|
||||
Drawing from mutualism, Mikhail Bakunin founded collectivist anarchism and entered the International Workingmen's Association, a class worker union later known as the First International that formed in 1864 to unite diverse revolutionary currents. The International became a significant political force, with Karl Marx being a leading figure and a member of its General Council. Bakunin's faction (the Jura Federation) and Proudhon's followers (the mutualists) opposed Marxist state socialism, advocating political abstentionism and small property holdings. After bitter disputes, the Bakuninists were expelled from the International by the Marxists at the 1872 Hague Congress. Bakunin famously predicted that if revolutionaries gained power by Marx's terms, they would end up the new tyrants of workers. After being expelled, anarchists formed the St. Imier International. Under the influence of Peter Kropotkin, a Russian philosopher and scientist, anarcho-communism overlapped with collectivism. Anarcho-communists, who drew inspiration from the 1871 Paris Commune, advocated for free federation and for the distribution of goods according to one's needs.
|
||||
|
||||
At the turn of the century, anarchism had spread all over the world. In China, small groups of students imported the humanistic pro-science version of anarcho-communism. Tokyo was a hotspot for rebellious youth from countries of the far east, travelling to the Japanese capital to study. In Latin America, Argentina was a stronghold for anarcho-syndicalism, where it became the most prominent left-wing ideology. During this time, a minority of anarchists adopted tactics of revolutionary political violence. This strategy became known as propaganda of the deed. The dismemberment of the French socialist movement into many groups, and the execution and exile of many Communards to penal colonies following the suppression of the Paris Commune, favoured individualist political expression and acts. Even though many anarchists distanced themselves from these terrorist acts, infamy came upon the movement. Illegalism was another strategy which some anarchists adopted during this period.
|
||||
Anarchists enthusiastically participated in the Russian Revolution—despite concerns—in opposition to the Whites. However, they met harsh suppression after the Bolshevik government was stabilized. Several anarchists from Petrograd and Moscow fled to Ukraine, notably leading to the Kronstadt rebellion and Nestor Makhno's struggle in the Free Territory. With the anarchists being crushed in Russia, two new antithetical currents emerged, namely platformism and synthesis anarchism. The former sought to create a coherent group that would push for revolution while the latter were against anything that would resemble a political party. Seeing the victories of the Bolsheviks in the October Revolution and the resulting Russian Civil War, many workers and activists turned to communist parties, which grew at the expense of anarchism and other socialist movements. In France and the United States, members of major syndicalist movements, the General Confederation of Labour and Industrial Workers of the World, left their organisations and joined the Communist International.
|
||||
|
||||
In the Spanish Civil War, anarchists and syndicalists (CNT and FAI) once again allied themselves with various currents of leftists. A long tradition of Spanish anarchism led to anarchists playing a pivotal role in the war. In response to the army rebellion, an anarchist-inspired movement of peasants and workers, supported by armed militias, took control of Barcelona and of large areas of rural Spain, where they collectivised the land. The Soviet Union provided some limited assistance at the beginning of the war, but the result was a bitter fight among communists and anarchists at a series of events named May Days as Joseph Stalin tried to seize control of the Republicans.
|
||||
|
||||
At the end of World War II, the anarchist movement was severely weakened. However, the 1960s witnessed a revival of anarchism likely caused by a perceived failure of Marxism–Leninism and tensions built by the Cold War. During this time, anarchism took root in other movements critical towards both the state and capitalism, such as the anti-nuclear, environmental and pacifist movements, the New Left, and the counterculture of the 1960s. Anarchism became associated with punk subculture, as exemplified by bands such as Crass and the Sex Pistols, and the established feminist tendencies of anarcha-feminism returned with vigour during the second wave of feminism.
|
||||
|
||||
Around the turn of the 21st century, anarchism grew in popularity and influence within anti-war, anti-capitalist, and anti-globalisation movements. Anarchists became known for their involvement in protests against the World Trade Organization, the Group of Eight and the World Economic Forum. During the protests, "ad hoc" leaderless anonymous cadres known as black blocs engaged in rioting, property destruction, and violent confrontations with the police. Other organisational tactics pioneered in this time include security culture, affinity groups, and the use of decentralised technologies such as the internet. A significant event of this period was the confrontations at the WTO conference in Seattle in 1999. Anarchist ideas have been influential in the development of the Zapatistas in Mexico and the Democratic Federation of Northern Syria, more commonly known as Rojava, a "de facto" autonomous region in northern Syria.
|
||||
|
||||
Anarchist schools of thought have been generally grouped into two main historical traditions, social anarchism and individualist anarchism, owing to their different origins, values and evolution. The individualist current emphasises negative liberty in opposing restraints upon the free individual, while the social current emphasises positive liberty in aiming to achieve the free potential of society through equality and social ownership. In a chronological sense, anarchism can be segmented by the classical currents of the late 19th century, and the post-classical currents (such as anarcha-feminism, green anarchism and post-anarchism) developed thereafter.
|
||||
|
||||
Beyond the specific factions of anarchist movements which constitute political anarchism lies philosophical anarchism, which holds that the state lacks moral legitimacy, without necessarily accepting the imperative of revolution to eliminate it. A component especially of individualist anarchism, philosophical anarchism may tolerate the existence of a minimal state, but argues that citizens have no moral obligation to obey government when it conflicts with individual autonomy. Anarchism pays significant attention to moral arguments since ethics have a central role in anarchist philosophy.
|
||||
|
||||
One reaction against sectarianism within the anarchist milieu was anarchism without adjectives, a call for toleration and unity among anarchists first adopted by Fernando Tarrida del Mármol in 1889 in response to the bitter debates of anarchist theory at the time. Despite separation, the various anarchist schools of thought are not seen as distinct entities, but as tendencies that intermingle.
|
||||
|
||||
Anarchism is usually placed on the far-left of the political spectrum. Much of its economics and legal philosophy reflect anti-authoritarian, anti-statist, and libertarian interpretations of the radical left-wing and socialist politics of collectivism, communism, individualism, mutualism, and syndicalism, among other libertarian socialist economic theories. As anarchism does not offer a fixed body of doctrine from a single particular worldview, many anarchist types and traditions exist, and varieties of anarchy diverge widely.
|
||||
|
||||
Inceptive currents among classical anarchist currents were mutualism and individualism. They were followed by the major currents of social anarchism (collectivist, communist, and syndicalist). They differ on organizational and economic aspects of their ideal society.
|
||||
|
||||
Mutualism is an 18th-century economic theory that was developed into anarchist theory by Pierre-Joseph Proudhon. Its aims include reciprocity, free association, voluntary contract, federation, and credit and currency reform that would be regulated by a bank of the people. Mutualism has been retrospectively characterised as ideologically situated between individualist and collectivist forms of anarchism. Proudhon first characterised his goal as a "third form of society, the synthesis of communism and property".
|
||||
|
||||
Collectivist anarchism, also known as anarchist collectivism or anarcho-collectivism, is a revolutionary socialist form of anarchism commonly associated with Mikhail Bakunin. Collectivist anarchists advocate collective ownership of the means of production, theorised to be achieved through violent revolution, and that workers be paid according to time worked, rather than goods being distributed according to need as in communism. Collectivist anarchism arose alongside Marxism, but rejected the dictatorship of the proletariat despite the stated Marxist goal of a collectivist stateless society. Anarcho-communism, also known as anarchist-communism, communist anarchism, and libertarian communism, is a theory of anarchism that advocates a communist society with common ownership of the means of production, direct democracy, and a horizontal network of voluntary associations and workers' councils with production and consumption based on the guiding principle: "From each according to his ability, to each according to his need". Anarcho-communism developed from radical socialist currents after the French Revolution, but it was first formulated as such in the Italian section of the First International. It was later expanded upon in the theoretical work of Peter Kropotkin.
|
||||
|
||||
Anarcho-syndicalism, also referred to as revolutionary syndicalism, is a branch of anarchism that views labour syndicates as a potential force for revolutionary social change, replacing capitalism and the state with a new society democratically self-managed by workers. The basic principles of anarcho-syndicalism are workers' solidarity, direct action, and workers' self-management.
|
||||
|
||||
Individualist anarchism refers to several traditions of thought within the anarchist movement that emphasise the individual and their will over any kinds of external determinants. Early influences on individualist forms of anarchism include William Godwin, Max Stirner and Henry David Thoreau. Through many countries, individualist anarchism attracted a small yet diverse following of Bohemian artists and intellectuals as well as young anarchist outlaws in what became known as illegalism and individual reclamation.
|
||||
|
||||
Anarchist principles undergird contemporary radical social movements of the left. Interest in the anarchist movement developed alongside momentum in the anti-globalization movement, whose leading activist networks were anarchist in orientation. As the movement shaped 21st century radicalism, wider embrace of anarchist principles signaled a revival of interest. Contemporary news coverage which emphasizes black bloc demonstrations has reinforced anarchism's historical association with chaos and violence, although its publicity has also led more scholars to engage with the anarchist movement. Anarchism has continued to generate many philosophies and movements—at times eclectic, drawing upon various sources, and syncretic, combining disparate concepts to create new philosophical approaches. The anti-capitalist tradition of classical anarchism has remained prominent within contemporary currents.
|
||||
|
||||
Various anarchist groups, tendencies, and schools of thought exist today, making it difficult to describe contemporary anarchist movement. While theorists and activists have established "relatively stable constellations of anarchist principles", there is no consensus on which principles are core. As a result, commentators describe multiple "anarchisms" (rather than a singular "anarchism") in which common principles are shared between schools of anarchism while each group prioritizes those principles differently. For example, gender equality can be a common principle but ranks as a higher priority to anarcha-feminists than anarchist communists. Anarchists are generally committed against coercive authority in all forms, namely "all centralized and hierarchical forms of government (e.g., monarchy, representative democracy, state socialism, etc.), economic class systems (e.g., capitalism, Bolshevism, feudalism, slavery, etc.), autocratic religions (e.g., fundamentalist Islam, Roman Catholicism, etc.), patriarchy, heterosexism, white supremacy, and imperialism". However, anarchist schools disagree on the methods by which these forms should be opposed.
|
||||
|
||||
Anarchists' tactics take various forms but in general serve two major goals—first, to oppose the Establishment; and second, to promote anarchist ethics and reflect an anarchist vision of society, illustrating the unity of means and ends. A broad categorization can be made between aims to destroy oppressive states and institutions by revolutionary means, and aims to change society through evolutionary means. Evolutionary tactics reject violence and take a gradual approach to anarchist aims, though there is significant overlap between the two.
|
||||
|
||||
Anarchist tactics have shifted during the course of the last century. Anarchists during the early 20th century focused more on strikes and militancy, while contemporary anarchists use a broader array of approaches.
|
||||
|
||||
During the classical era, anarchists had a militant tendency. Not only did they confront state armed forces (as in Spain and Ukraine) but some of them also employed terrorism as propaganda of the deed. Assassination attempts were carried out against heads of state, some of which were successful. Anarchists also took part in revolutions. Anarchist perspectives towards violence have always been perplexing and controversial. On one hand, anarcho-pacifists point out the unity of means and ends. On the other hand, other anarchist groups advocate direct action, a tactic which can include acts of sabotage or even acts of terrorism. This attitude was quite prominent a century ago; seeing the state as a tyrant, some anarchists believed that they had every right to oppose its oppression by any means possible. Emma Goldman and Errico Malatesta, who were proponents of limited use of violence, argued that violence is merely a reaction to state violence as a necessary evil.
|
||||
|
||||
Anarchists took an active role in strikes, although they tended to be antipathetic to formal syndicalism, seeing it as reformist. They saw it as a part of the movement which sought to overthrow the state and capitalism. Anarchists also reinforced their propaganda within the arts, some of whom practiced nudism. They also built communities which were based on friendship. They were also involved in the press.
|
||||
|
||||
In the current era, Italian anarchist Alfredo Bonanno, a proponent of insurrectionary anarchism, has reinstated the debate on violence by rejecting the nonviolence tactic adopted since the late 19th century by Kropotkin and other prominent anarchists afterwards. Both Bonanno and the French group The Invisible Committee advocate for small, informal affiliation groups, where each member is responsible for their own actions but works together to bring down oppression utilizing sabotage and other violent means against state, capitalism and other enemies. Members of The Invisible Committee were arrested in 2008 on various charges, terrorism included.
|
||||
|
||||
Overall, today's anarchists are much less violent and militant than their ideological ancestors. They mostly engage in confronting the police during demonstrations and riots, especially in countries like Canada, Mexico or Greece. Μilitant black bloc protest groups are known for clashing with the police. However, anarchists not only clash with state operators; they also engage in the struggle against fascists and racists, taking anti-fascist action and mobilizing to prevent hate rallies from happening.
|
||||
|
||||
Anarchists commonly employ direct action. This can take the form of disrupting and protesting against unjust hierarchy, or the form of self-managing their lives through the creation of counter-institutions such as communes and non-hierarchical collectives. Often, decision-making is handled in an anti-authoritarian way, with everyone having equal say in each decision, an approach known as horizontalism. Contemporary-era anarchists have been engaging with various grassroots movements that are not explicitly anarchist but are more or less based on horizontalism, respecting personal autonomy, and participating in mass activism such as strikes and demonstrations. The newly coined term "small-a anarchism", in contrast with the "big-A anarchism" of the classical era, signals their tendency not to base their thoughts and actions on classical-era anarchism or to refer to Kropotkin or Proudhon to justify their opinions. They would rather base their thought and praxis on their own experience, which they will later theorize.
|
||||
|
||||
The decision-making process of small affinity anarchist groups play a significant tactical role. Anarchists have employed various methods in order to build a rough consensus among members of their group, without the need of a leader or a leading group. One way is for an individual from the group to play the role of facilitator to help achieve a consensus without taking part in the discussion themselves or promoting a specific point. Minorities usually accept rough consensus, except when they feel the proposal contradicts anarchist goals, values, or ethics. Anarchists usually form small groups (5–20 individuals) to enhance autonomy and friendships among their members. These kind of groups more often than not interconnect with each other, forming larger networks. Anarchists still support and participate in strikes, especially wildcat strikes; these are leaderless strikes not organised centrally by a syndicate.
|
||||
|
||||
Anarchists have gone online to spread their message. As in the past, newspapers and journals are used; however, because of distributional and other difficulties, anarchists have found it easier to create websites, hosting electronic libraries and other portals. Anarchists were also involved in developing various software that are available for free. The way these hacktivists work to develop and distribute resembles the anarchist ideals, especially when it comes to preserving user's privacy from state surveillance.
|
||||
|
||||
Anarchists organize themselves to squat and reclaim public spaces. During important events such as protests and when spaces are being occupied, they are often called Temporary Autonomous Zones (TAZ), spaces where surrealism, poetry and art are blended to display the anarchist ideal. As seen by anarchists, squatting is a way to regain urban space from the capitalist market, serving pragmatical needs, and is also seen an exemplary direct action. Acquiring space enables anarchists to experiment with their ideas and build social bonds. Adding up these tactics, and having in mind that not all anarchists share the same attitudes towards them, along with various forms of protesting at highly symbolic events, make up a carnivalesque atmosphere that is part of contemporary anarchist vividity.
|
||||
|
||||
As anarchism is a philosophy that embodies many diverse attitudes, tendencies, and schools of thought, and disagreement over questions of values, ideology, and tactics is common, its diversity has led to widely different uses of identical terms among different anarchist traditions, which has created a number of definitional concerns in anarchist theory. For instance, the compatibility of capitalism, nationalism and religion with anarchism is widely disputed. Similarly, anarchism enjoys complex relationships with ideologies such as Marxism, communism, collectivism and trade unionism. Anarchists may be motivated by humanism, divine authority, enlightened self-interest, veganism, or any number of alternative ethical doctrines. Phenomena such as civilisation, technology (e.g. within anarcho-primitivism) and the democratic process may be sharply criticised within some anarchist tendencies and simultaneously lauded in others.
|
||||
|
||||
Gender and sexuality carry along them dynamics of hierarchy; anarchism is obliged to address, analyse and oppose the suppression of one's autonomy because of the dynamics that gender roles traditionally impose.
|
||||
|
||||
A historical current that arose and flourished during 1890 and 1920 within anarchism was free love; in contemporary anarchism, this current survives as a tendency to support polyamory and queer anarchism. Free love advocates were against marriage, which they saw as a way of men imposing authority over women, largely because marriage law greatly favoured the power of men. The notion of free love, though, was much broader; it included critique of the established order that limited women's sexual freedom and pleasure. Such free love movements contributed to the establishment of communal houses, where large groups of travelers, anarchists, and other activists slept in beds together. Free love had roots both in Europe and the United States. Some anarchists, however, struggled with the jealousy that arose from free love. Anarchist feminists were advocates of free love, against marriage, were pro-choice (utilizing a contemporary term) and had a likewise agenda. Anarchist and non-anarchist feminists differed on suffrage, but were nonetheless supportive of one another.
|
||||
|
||||
During the second half of the 20th century, anarchism intermingled with the second wave of feminism, radicalizing some currents of the feminist movement (and being influenced as well). By the latest decades of the 20th century, anarchists and feminists were advocating for the rights and autonomy of women, gays, queers and other marginalized groups, with some feminist thinkers suggesting a fusion of the two currents. With the third wave of feminism, sexual identity and compulsory heterosexuality became a subject of study for anarchists, which yielded a post-structuralist critique of sexual normality. However, some anarchists distanced themselves from this line of thinking, suggesting that it leaned towards individualism and was, therefore, dropping the cause of social liberation.
|
||||
|
||||
The interest of anarchists in education stretches back to the first emergence of classical anarchism. Anarchists consider 'proper' education, which sets the foundations of the future autonomy of the individual and the society, to be an act of mutual aid. Anarchist writers such as Willian Godwin and Max Stirner attacked both state education and private education as another means by which the ruling class replicate their privileges.
|
||||
|
||||
In 1901, Catalan anarchist and free thinker Francisco Ferrer established the Escuela Moderna in Barcelona as an opposition to the established education system, which was dictated largely by the Catholic Church. Ferrer's approach was secular, rejecting both state and church involvement in the educational process, and gave pupils large amounts of autonomy in planning their work and attendance. Ferrer aimed to educate the working class and explicitly sought to foster class consciousness among students. The school closed after constant harassment by the state and Ferrer was later arrested. His ideas, however, formed the inspiration for a series of modern schools around the world. Christian anarchist Leo Tolstoy also established a similar school, with its founding principle, according to Tolstoy, being that "for education to be effective it had to be free". In a similar token, A. S. Neill founding what became Summerhill School in 1921, also declaring being free from coercion.
|
||||
|
||||
Anarchist education is based largely on the idea that a child's right to develop freely, without manipulation, ought to be respected, and that rationality will lead children to morally good conclusions. However, there has been little consensus among anarchist figures as to what constitutes manipulation; Ferrer, for example, believed that moral indoctrination was necessary and explicitly taught pupils that equality, liberty, and social justice were not possible under capitalism (along with other critiques of nationalism and government).
|
||||
|
||||
Late 20th century and contemporary anarchist writers (such as Colin Ward, Herbert Read and Paul Goodman) intensified and expanded the anarchist critique of state education, largely focusing on the need for a system that focuses on children's creativity rather than on their ability to attain a career or participate in consumer society. Contemporary anarchists, such as Colin Ward, have further argued that state education serves to perpetuate socio-economic inequality.
|
||||
|
||||
While few anarchist education institutions have survived to the modern day, major tenets of anarchist schools, such as respect for child autonomy and relying on reasoning rather than indoctrination as a teaching method, have spread among mainstream educational institutions.
|
||||
|
||||
Objection to the state and its institutions is a "sine qua non" of anarchism. Anarchists consider the state as a tool of domination and believe it to be illegitimate regardless of its political tendencies. Instead of people being able to control the aspects of their life, major decisions are taken by a small elite. Authority ultimately rests solely on power, regardless of whether that power is open or transparent, as it still has the ability to coerce people. Another anarchist argument against states is that the people constituting a government, even the most altruistic among officials, will unavoidably seek to gain more power, leading to corruption. Anarchists consider the idea that the state is the collective will of the people to be an unachievable fiction, due to the fact that the ruling class is distinct from the rest of society.
|
||||
|
||||
The connection between anarchism and art was quite profound during the classical era of anarchism, especially among artistic currents that were developing during that era, such as futurists, surrealists, and others, while in literature anarchism was mostly associated with the New Apocalyptics and the Neo-romanticism movement. In music, anarchism has been associated with music scenes such as Punk. Anarchists such as Leo Tolstoy and Herbert Read argued that the border between the artist and the non-artist, what separates art from a daily act, is a construct produced by the alienation caused by capitalism, and it prevents humans from living a joyful life.
|
||||
|
||||
Other anarchists advocated for or used art as a means to achieve anarchist ends. In his book Breaking the Spell: A History of Anarchist Filmmakers, Videotape Guerrillas, and Digital Ninjas Chris Robé claims that "anarchist-inflected practices have increasingly structured movement-based video activism."
|
||||
|
||||
Three overlapping properties made art useful to anarchists: It could depict a critique of existing society and hierarchies; it could serve as a prefigurative tool to reflect the anarchist ideal society, and also it could turn into a means of direct action, in protests for example. As it appeals to both emotion and reason, art could appeal to the "whole human" and have a powerful effect.
|
||||
|
||||
Philosophy lecturer Andrew G. Fiala has listed five main arguments against anarchism. Firstly, he notes that anarchism is related to violence and destruction, not only in the pragmatic world (i.e. at protests) but in the world of ethics as well. The second argument is that it is impossible for a society to function without a state or something like a state, acting to protect citizens from criminality. Fiala takes "Leviathan" from Thomas Hobbes and the night-watchman state from philosopher Robert Nozick as examples. Thirdly, anarchism is evaluated as unfeasible or utopian since the state can not be defeated practically; this line of arguments most often calls for political action within the system to reform it. The fourth argument is that anarchism is self-contradictory since while it advocates for no-one to "archiei", if accepted by the many, then anarchism will turn into the ruling political theory. In this line of criticism also comes the self contradiction that anarchist calls for collective action while anarchism endorses the autonomy of the individual and hence no collective action can be taken. Lastly, Fiala mentions a critique towards philosophical anarchism, of being ineffective (all talk and thoughts) and in the meantime capitalism and bourgeois class remains strong.
|
||||
|
||||
Philosophical anarchism has met the criticism of members of academia, following the release of pro-anarchist books such as A. John Simmons' "Moral Principles and Political Obligations" (1979). Law professor William A. Edmundson authored an essay arguing against three major philosophical anarchist principles, which he finds fallacious; Edmundson claims that while the individual does not owe a normal state a duty of obedience, this does not imply that anarchism is the inevitable conclusion, and the state is still morally legitimate.
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
</doc>
|
||||
<doc id="25" url="https://en.wikipedia.org/wiki?curid=25" title="Autism">
|
||||
Autism
|
||||
|
||||
Autism is a developmental disorder characterized by difficulties with social interaction and communication, and by restricted and repetitive behavior. Parents often notice signs during the first three years of their child's life. These signs often develop gradually, though some children with autism experience worsening in their communication and social skills after reaching developmental milestones at a normal pace.
|
||||
Autism is associated with a combination of genetic and environmental factors. Risk factors during pregnancy include certain infections, such as rubella, toxins including valproic acid, alcohol, cocaine, pesticides, lead, and air pollution, fetal growth restriction, and autoimmune diseases. Controversies surround other proposed environmental causes; for example, the vaccine hypothesis, which has been disproven. Autism affects information processing in the brain and how nerve cells and their synapses connect and organize; how this occurs is not well understood. The Diagnostic and Statistical Manual of Mental Disorders (DSM-5), combines autism and less severe forms of the condition, including Asperger syndrome and pervasive developmental disorder not otherwise specified (PDD-NOS) into the diagnosis of autism spectrum disorder (ASD).
|
||||
Early behavioral interventions or speech therapy can help children with autism gain self-care, social, and communication skills. Although there is no known cure, there have been cases of children who recovered. Some autistic adults are unable to live independently. An autistic culture has developed, with some individuals seeking a cure and others believing autism should be accepted as a difference to be accommodated instead of cured.
|
||||
Globally, autism is estimated to affect 24.8 million people . In the 2000s, the number of people affected was estimated at 1–2 per 1,000 people worldwide. In the developed countries, about 1.5% of children are diagnosed with ASD , from 0.7% in 2000 in the United States. It occurs four-to-five times more often in males than females. The number of people diagnosed has increased dramatically since the 1960s, which may be partly due to changes in diagnostic practice. The question of whether actual rates have increased is unresolved.
|
||||
Autism is a highly variable, neurodevelopmental disorder whose symptoms first appears during infancy or childhood, and generally follows a steady course without remission. People with autism may be severely impaired in some respects but average, or even superior, in others. Overt symptoms gradually begin after the age of six months, become established by age two or three years and tend to continue through adulthood, although often in more muted form. It is distinguished by a characteristic triad of symptoms: impairments in social interaction, impairments in communication, and repetitive behavior. Other aspects, such as atypical eating, are also common but are not essential for diagnosis. Individual symptoms of autism occur in the general population and appear not to associate highly, without a sharp line separating pathologically severe from common traits.
|
||||
|
||||
Social deficits distinguish autism and the related autism spectrum disorders (ASD; see Classification) from other developmental disorders. People with autism have social impairments and often lack the intuition about others that many people take for granted. Noted autistic Temple Grandin described her inability to understand the social communication of neurotypicals, or people with typical neural development, as leaving her feeling "like an anthropologist on Mars".
|
||||
|
||||
Unusual social development becomes apparent early in childhood. Autistic infants show less attention to social stimuli, smile and look at others less often, and respond less to their own name. Autistic toddlers differ more strikingly from social norms; for example, they have less eye contact and turn-taking, and do not have the ability to use simple movements to express themselves, such as pointing at things. Three- to five-year-old children with autism are less likely to exhibit social understanding, approach others spontaneously, imitate and respond to emotions, communicate nonverbally, and take turns with others. However, they do form attachments to their primary caregivers. Most children with autism display moderately less attachment security than neurotypical children, although this difference disappears in children with higher mental development or less pronounced autistic traits. Older children and adults with ASD perform worse on tests of face and emotion recognition although this may be partly due to a lower ability to define a person's own emotions.
|
||||
|
||||
Children with high-functioning autism have more intense and frequent loneliness compared to non-autistic peers, despite the common belief that children with autism prefer to be alone. Making and maintaining friendships often proves to be difficult for those with autism. For them, the quality of friendships, not the number of friends, predicts how lonely they feel. Functional friendships, such as those resulting in invitations to parties, may affect the quality of life more deeply.
|
||||
There are many anecdotal reports, but few systematic studies, of aggression and violence in individuals with ASD. The limited data suggest that, in children with intellectual disability, autism is associated with aggression, destruction of property, and meltdowns.
|
||||
|
||||
About a third to a half of individuals with autism do not develop enough natural speech to meet their daily communication needs. Differences in communication may be present from the first year of life, and may include delayed onset of babbling, unusual gestures, diminished responsiveness, and vocal patterns that are not synchronized with the caregiver. In the second and third years, children with autism have less frequent and less diverse babbling, consonants, words, and word combinations; their gestures are less often integrated with words. Children with autism are less likely to make requests or share experiences, and are more likely to simply repeat others' words (echolalia) or reverse pronouns. Joint attention seems to be necessary for functional speech, and deficits in joint attention seem to distinguish infants with ASD. For example, they may look at a pointing hand instead of the pointed-at object, and they consistently fail to point at objects in order to comment on or share an experience. Children with autism may have difficulty with imaginative play and with developing symbols into language.
|
||||
|
||||
In a pair of studies, high-functioning children with autism aged 8–15 performed equally well as, and as adults better than, individually matched controls at basic language tasks involving vocabulary and spelling. Both autistic groups performed worse than controls at complex language tasks such as figurative language, comprehension and inference. As people are often sized up initially from their basic language skills, these studies suggest that people speaking to autistic individuals are more likely to overestimate what their audience comprehends.
|
||||
|
||||
Autistic individuals can display many forms of repetitive or restricted behavior, which the Repetitive Behavior Scale-Revised (RBS-R) categorizes as follows.
|
||||
|
||||
|
||||
No single repetitive or self-injurious behavior seems to be specific to autism, but autism appears to have an elevated pattern of occurrence and severity of these behaviors.
|
||||
|
||||
Autistic individuals may have symptoms that are independent of the diagnosis, but that can affect the individual or the family.
|
||||
An estimated 0.5% to 10% of individuals with ASD show unusual abilities, ranging from splinter skills such as the memorization of trivia to the extraordinarily rare talents of prodigious autistic savants. Many individuals with ASD show superior skills in perception and attention, relative to the general population. Sensory abnormalities are found in over 90% of those with autism, and are considered core features by some, although there is no good evidence that sensory symptoms differentiate autism from other developmental disorders. Differences are greater for under-responsivity (for example, walking into things) than for over-responsivity (for example, distress from loud noises) or for sensation seeking (for example, rhythmic movements). An estimated 60–80% of autistic people have motor signs that include poor muscle tone, poor motor planning, and toe walking; deficits in motor coordination are pervasive across ASD and are greater in autism proper. Unusual eating behavior occurs in about three-quarters of children with ASD, to the extent that it was formerly a diagnostic indicator. Selectivity is the most common problem, although eating rituals and food refusal also occur.
|
||||
|
||||
There is tentative evidence that autism occurs more frequently in people with gender dysphoria.
|
||||
|
||||
Gastrointestinal problems are one of the most commonly associated medical disorders in people with autism. These are linked to greater social impairment, irritability, behavior and sleep problems, language impairments and mood changes.
|
||||
|
||||
Parents of children with ASD have higher levels of stress. Siblings of children with ASD report greater admiration of and less conflict with the affected sibling than siblings of unaffected children and were similar to siblings of children with Down syndrome in these aspects of the sibling relationship. However, they reported lower levels of closeness and intimacy than siblings of children with Down syndrome; siblings of individuals with ASD have greater risk of negative well-being and poorer sibling relationships as adults.
|
||||
|
||||
It has long been presumed that there is a common cause at the genetic, cognitive, and neural levels for autism's characteristic triad of symptoms. However, there is increasing suspicion that autism is instead a complex disorder whose core aspects have distinct causes that often co-occur.
|
||||
Autism has a strong genetic basis, although the genetics of autism are complex and it is unclear whether ASD is explained more by rare mutations with major effects, or by rare multigene interactions of common genetic variants. Complexity arises due to interactions among multiple genes, the environment, and epigenetic factors which do not change DNA sequencing but are heritable and influence gene expression. Many genes have been associated with autism through sequencing the genomes of affected individuals and their parents. Studies of twins suggest that heritability is 0.7 for autism and as high as 0.9 for ASD, and siblings of those with autism are about 25 times more likely to be autistic than the general population. However, most of the mutations that increase autism risk have not been identified. Typically, autism cannot be traced to a Mendelian (single-gene) mutation or to a single chromosome abnormality, and none of the genetic syndromes associated with ASDs have been shown to selectively cause ASD. Numerous candidate genes have been located, with only small effects attributable to any particular gene. Most loci individually explain less than 1% of cases of autism. The large number of autistic individuals with unaffected family members may result from spontaneous structural variation—such as deletions, duplications or inversions in genetic material during meiosis. Hence, a substantial fraction of autism cases may be traceable to genetic causes that are highly heritable but not inherited: that is, the mutation that causes the autism is not present in the parental genome. Autism may be underdiagnosed in women and girls due to an assumption that it is primarily a male condition, but genetic phenomena such as imprinting and X linkage have the ability to raise the frequency and severity of conditions in males, and theories have been put forward for a genetic reason why males are diagnosed more often, such as the imprinted brain theory and the extreme male brain theory.
|
||||
|
||||
Maternal nutrition and inflammation during preconception and pregnancy influences fetal neurodevelopment. Intrauterine growth restriction is associated with ASD, in both term and preterm infants. Maternal inflammatory and autoimmune diseases may damage fetal tissues, aggravating a genetic problem or damaging the nervous system.
|
||||
|
||||
Exposure to air pollution during pregnancy, especially heavy metals and particulates, may increase the risk of autism. Environmental factors that have been claimed without evidence to contribute to or exacerbate autism include certain foods, infectious diseases, solvents, PCBs, phthalates and phenols used in plastic products, pesticides, brominated flame retardants, alcohol, smoking, illicit drugs, vaccines, and prenatal stress. Some, such as the MMR vaccine, have been completely disproven.
|
||||
|
||||
Parents may first become aware of autistic symptoms in their child around the time of a routine vaccination. This has led to unsupported theories blaming vaccine "overload", a vaccine preservative, or the MMR vaccine for causing autism. The latter theory was supported by a litigation-funded study that has since been shown to have been "an elaborate fraud". Although these theories lack convincing scientific evidence and are biologically implausible, parental concern about a potential vaccine link with autism has led to lower rates of childhood immunizations, outbreaks of previously controlled childhood diseases in some countries, and the preventable deaths of several children.
|
||||
|
||||
Autism's symptoms result from maturation-related changes in various systems of the brain. How autism occurs is not well understood. Its mechanism can be divided into two areas: the pathophysiology of brain structures and processes associated with autism, and the neuropsychological linkages between brain structures and behaviors. The behaviors appear to have multiple pathophysiologies.
|
||||
|
||||
There is evidence that gut–brain axis abnormalities may be involved. A 2015 review proposed that immune dysregulation, gastrointestinal inflammation, malfunction of the autonomic nervous system, gut flora alterations, and food metabolites may cause brain neuroinflammation and dysfunction. A 2016 review concludes that enteric nervous system abnormalities might play a role in neurological disorders such as autism. Neural connections and the immune system are a pathway that may allow diseases originated in the intestine to spread to the brain.
|
||||
|
||||
Several lines of evidence point to synaptic dysfunction as a cause of autism. Some rare mutations may lead to autism by disrupting some synaptic pathways, such as those involved with cell adhesion. Gene replacement studies in mice suggest that autistic symptoms are closely related to later developmental steps that depend on activity in synapses and on activity-dependent changes. All known teratogens (agents that cause birth defects) related to the risk of autism appear to act during the first eight weeks from conception, and though this does not exclude the possibility that autism can be initiated or affected later, there is strong evidence that autism arises very early in development.
|
||||
|
||||
Diagnosis is based on behavior, not cause or mechanism. Under the DSM-5, autism is characterized by persistent deficits in social communication and interaction across multiple contexts, as well as restricted, repetitive patterns of behavior, interests, or activities. These deficits are present in early childhood, typically before age three, and lead to clinically significant functional impairment. Sample symptoms include lack of social or emotional reciprocity, stereotyped and repetitive use of language or idiosyncratic language, and persistent preoccupation with unusual objects. The disturbance must not be better accounted for by Rett syndrome, intellectual disability or global developmental delay. ICD-10 uses essentially the same definition.
|
||||
|
||||
Several diagnostic instruments are available. Two are commonly used in autism research: the Autism Diagnostic Interview-Revised (ADI-R) is a semistructured parent interview, and the Autism Diagnostic Observation Schedule (ADOS) uses observation and interaction with the child. The Childhood Autism Rating Scale (CARS) is used widely in clinical environments to assess severity of autism based on observation of children. The Diagnostic interview for social and communication disorders (DISCO) may also be used.
|
||||
|
||||
A pediatrician commonly performs a preliminary investigation by taking developmental history and physically examining the child. If warranted, diagnosis and evaluations are conducted with help from ASD specialists, observing and assessing cognitive, communication, family, and other factors using standardized tools, and taking into account any associated medical conditions. A pediatric neuropsychologist is often asked to assess behavior and cognitive skills, both to aid diagnosis and to help recommend educational interventions. A differential diagnosis for ASD at this stage might also consider intellectual disability, hearing impairment, and a specific language impairment such as Landau–Kleffner syndrome. The presence of autism can make it harder to diagnose coexisting psychiatric disorders such as depression.
|
||||
|
||||
Clinical genetics evaluations are often done once ASD is diagnosed, particularly when other symptoms already suggest a genetic cause. Although genetic technology allows clinical geneticists to link an estimated 40% of cases to genetic causes, consensus guidelines in the US and UK are limited to high-resolution chromosome and fragile X testing. A genotype-first model of diagnosis has been proposed, which would routinely assess the genome's copy number variations. As new genetic tests are developed several ethical, legal, and social issues will emerge. Commercial availability of tests may precede adequate understanding of how to use test results, given the complexity of autism's genetics. Metabolic and neuroimaging tests are sometimes helpful, but are not routine.
|
||||
|
||||
ASD can sometimes be diagnosed by age 14 months, although diagnosis becomes increasingly stable over the first three years of life: for example, a one-year-old who meets diagnostic criteria for ASD is less likely than a three-year-old to continue to do so a few years later. In the UK the National Autism Plan for Children recommends at most 30 weeks from first concern to completed diagnosis and assessment, though few cases are handled that quickly in practice. Although the symptoms of autism and ASD begin early in childhood, they are sometimes missed; years later, adults may seek diagnoses to help them or their friends and family understand themselves, to help their employers make adjustments, or in some locations to claim disability living allowances or other benefits. Girls are often diagnosed later than boys.
|
||||
|
||||
Underdiagnosis and overdiagnosis are problems in marginal cases, and much of the recent increase in the number of reported ASD cases is likely due to changes in diagnostic practices. The increasing popularity of drug treatment options and the expansion of benefits has given providers incentives to diagnose ASD, resulting in some overdiagnosis of children with uncertain symptoms. Conversely, the cost of screening and diagnosis and the challenge of obtaining payment can inhibit or delay diagnosis. It is particularly hard to diagnose autism among the visually impaired, partly because some of its diagnostic criteria depend on vision, and partly because autistic symptoms overlap with those of common blindness syndromes or blindisms.
|
||||
|
||||
Autism is one of the five pervasive developmental disorders (PDD), which are characterized by widespread abnormalities of social interactions and communication, and severely restricted interests and highly repetitive behavior. These symptoms do not imply sickness, fragility, or emotional disturbance.
|
||||
|
||||
Of the five PDD forms, Asperger syndrome is closest to autism in signs and likely causes; Rett syndrome and childhood disintegrative disorder share several signs with autism, but may have unrelated causes; PDD not otherwise specified (PDD-NOS; also called "atypical autism") is diagnosed when the criteria are not met for a more specific disorder. Unlike with autism, people with Asperger syndrome have no substantial delay in language development. The terminology of autism can be bewildering, with autism, Asperger syndrome and PDD-NOS often called the "autism spectrum disorders" (ASD) or sometimes the "autistic disorders", whereas autism itself is often called "autistic disorder", "childhood autism", or "infantile autism". In this article, "autism" refers to the classic autistic disorder; in clinical practice, though, "autism", "ASD", and "PDD" are often used interchangeably. ASD, in turn, is a subset of the broader autism phenotype, which describes individuals who may not have ASD but do have autistic-like traits, such as avoiding eye contact.
|
||||
|
||||
Autism can also be divided into syndromal and non-syndromal autism; the syndromal autism is associated with severe or profound intellectual disability or a congenital syndrome with physical symptoms, such as tuberous sclerosis. Although individuals with Asperger syndrome tend to perform better cognitively than those with autism, the extent of the overlap between Asperger syndrome, HFA, and non-syndromal autism is unclear.
|
||||
|
||||
Some studies have reported diagnoses of autism in children due to a loss of language or social skills, as opposed to a failure to make progress, typically from 15 to 30 months of age. The validity of this distinction remains controversial; it is possible that regressive autism is a specific subtype, or that there is a continuum of behaviors between autism with and without regression.
|
||||
|
||||
Research into causes has been hampered by the inability to identify biologically meaningful subgroups within the autistic population and by the traditional boundaries between the disciplines of psychiatry, psychology, neurology and pediatrics. Newer technologies such as fMRI and diffusion tensor imaging can help identify biologically relevant phenotypes (observable traits) that can be viewed on brain scans, to help further neurogenetic studies of autism; one example is lowered activity in the fusiform face area of the brain, which is associated with impaired perception of people versus objects. It has been proposed to classify autism using genetics as well as behavior.
|
||||
|
||||
Autism has long been thought to cover a wide spectrum, ranging from individuals with severe impairments—who may be silent, developmentally disabled, and prone to frequent repetitive behavior such as hand flapping and rocking—to high functioning individuals who may have active but distinctly odd social approaches, narrowly focused interests, and verbose, pedantic communication. Because the behavior spectrum is continuous, boundaries between diagnostic categories are necessarily somewhat arbitrary. Sometimes the syndrome is divided into low-, medium- or high-functioning autism (LFA, MFA, and HFA), based on IQ thresholds. Some people have called for an end to the terms "high-functioning" and "low-functioning" due to lack of nuance and the potential for a person's needs or abilities to be overlooked.
|
||||
|
||||
About half of parents of children with ASD notice their child's unusual behaviors by age 18 months, and about four-fifths notice by age 24 months. According to an article, failure to meet any of the following milestones "is an absolute indication to proceed with further evaluations. Delay in referral for such testing may delay early diagnosis and treatment and affect the long-term outcome".
|
||||
|
||||
The United States Preventive Services Task Force in 2016 found it was unclear if screening was beneficial or harmful among children in whom there is no concerns. The Japanese practice is to screen all children for ASD at 18 and 24 months, using autism-specific formal screening tests. In contrast, in the UK, children whose families or doctors recognize possible signs of autism are screened. It is not known which approach is more effective. Screening tools include the Modified Checklist for Autism in Toddlers (M-CHAT), the Early Screening of Autistic Traits Questionnaire, and the First Year Inventory; initial data on M-CHAT and its predecessor, the Checklist for Autism in Toddlers (CHAT), on children aged 18–30 months suggests that it is best used in a clinical setting and that it has low sensitivity (many false-negatives) but good specificity (few false-positives). It may be more accurate to precede these tests with a broadband screener that does not distinguish ASD from other developmental disorders. Screening tools designed for one culture's norms for behaviors like eye contact may be inappropriate for a different culture. Although genetic screening for autism is generally still impractical, it can be considered in some cases, such as children with neurological symptoms and dysmorphic features.
|
||||
|
||||
While infection with rubella during pregnancy causes fewer than 1% of cases of autism, vaccination against rubella can prevent many of those cases.
|
||||
|
||||
The main goals when treating children with autism are to lessen associated deficits and family distress, and to increase quality of life and functional independence. In general, higher IQs are correlated with greater responsiveness to treatment and improved treatment outcomes. No single treatment is best and treatment is typically tailored to the child's needs. Families and the educational system are the main resources for treatment. Services should be carried out by behavior analysts, special education teachers, speech pathologists, and licensed psychologists. Studies of interventions have methodological problems that prevent definitive conclusions about efficacy. However, the development of evidence-based interventions has advanced in recent years. Although many psychosocial interventions have some positive evidence, suggesting that some form of treatment is preferable to no treatment, the methodological quality of systematic reviews of these studies has generally been poor, their clinical results are mostly tentative, and there is little evidence for the relative effectiveness of treatment options. Intensive, sustained special education programs and behavior therapy early in life can help children acquire self-care, communication, and job skills, and often improve functioning and decrease symptom severity and maladaptive behaviors; claims that intervention by around age three years is crucial are not substantiated. While medications have not been found to help with core symptoms, they may be used for associated symptoms, such as irritability, inattention, or repetitive behavior patterns.
|
||||
|
||||
Educational interventions often used include applied behavior analysis (ABA), developmental models, structured teaching, speech and language therapy, social skills therapy, and occupational therapy. Among these approaches, interventions either treat autistic features comprehensively, or focalize treatment on a specific area of deficit. The quality of research for early intensive behavioral intervention (EIBI)—a treatment procedure incorporating over thirty hours per week of the structured type of ABA that is carried out with very young children—is currently low, and more vigorous research designs with larger sample sizes are needed. Two theoretical frameworks outlined for early childhood intervention include structured and naturalistic ABA interventions, and developmental social pragmatic models (DSP). One interventional strategy utilizes a parent training model, which teaches parents how to implement various ABA and DSP techniques, allowing for parents to disseminate interventions themselves. Various DSP programs have been developed to explicitly deliver intervention systems through at-home parent implementation. Despite the recent development of parent training models, these interventions have demonstrated effectiveness in numerous studies, being evaluated as a probable efficacious mode of treatment.
|
||||
|
||||
Early, intensive ABA therapy has demonstrated effectiveness in enhancing communication and adaptive functioning in preschool children; it is also well-established for improving the intellectual performance of that age group. Similarly, a teacher-implemented intervention that utilizes a more naturalistic form of ABA combined with a developmental social pragmatic approach has been found to be beneficial in improving social-communication skills in young children, although there is less evidence in its treatment of global symptoms. Neuropsychological reports are often poorly communicated to educators, resulting in a gap between what a report recommends and what education is provided. It is not known whether treatment programs for children lead to significant improvements after the children grow up, and the limited research on the effectiveness of adult residential programs shows mixed results. The appropriateness of including children with varying severity of autism spectrum disorders in the general education population is a subject of current debate among educators and researchers.
|
||||
|
||||
Medications may be used to treat ASD symptoms that interfere with integrating a child into home or school when behavioral treatment fails. They may also be used for associated health problems, such as ADHD or anxiety. More than half of US children diagnosed with ASD are prescribed psychoactive drugs or anticonvulsants, with the most common drug classes being antidepressants, stimulants, and antipsychotics. The atypical antipsychotic drugs risperidone and aripiprazole are FDA-approved for treating associated aggressive and self-injurious behaviors. However, their side effects must be weighed against their potential benefits, and people with autism may respond atypically. Side effects, for example, may include weight gain, tiredness, drooling, and aggression. SSRI antidepressants, such as fluoxetine and fluvoxamine, have been shown to be effective in reducing repetitive and ritualistic behaviors, while the stimulant medication methylphenidate is beneficial for some children with co-morbid inattentiveness or hyperactivity. There is scant reliable research about the effectiveness or safety of drug treatments for adolescents and adults with ASD. No known medication relieves autism's core symptoms of social and communication impairments. Experiments in mice have reversed or reduced some symptoms related to autism by replacing or modulating gene function, suggesting the possibility of targeting therapies to specific rare mutations known to cause autism.
|
||||
|
||||
Although many alternative therapies and interventions are available, few are supported by scientific studies. Treatment approaches have little empirical support in quality-of-life contexts, and many programs focus on success measures that lack predictive validity and real-world relevance. Some alternative treatments may place the child at risk. The preference that children with autism have for unconventional foods can lead to reduction in bone cortical thickness with this being greater in those on casein-free diets, as a consequence of the low intake of calcium and vitamin D; however, suboptimal bone development in ASD has also been associated with lack of exercise and gastrointestinal disorders. In 2005, botched chelation therapy killed a five-year-old child with autism. Chelation is not recommended for people with ASD since the associated risks outweigh any potential benefits. Another alternative medicine practice with no evidence is CEASE therapy, a mixture of homeopathy, supplements, and 'vaccine detoxing'.
|
||||
|
||||
Although popularly used as an alternative treatment for people with autism, as of 2018 there is no good evidence to recommend a gluten- and casein-free diet as a standard treatment. A 2018 review concluded that it may be a therapeutic option for specific groups of children with autism, such as those with known food intolerances or allergies, or with food intolerance markers. The authors analyzed the prospective trials conducted to date that studied the efficacy of the gluten- and casein-free diet in children with ASD (4 in total). All of them compared gluten- and casein-free diet versus normal diet with a control group (2 double-blind randomized controlled trials, 1 double-blind crossover trial, 1 single-blind trial). In two of the studies, whose duration was 12 and 24 months, a significant improvement in ASD symptoms (efficacy rate 50%) was identified. In the other two studies, whose duration was 3 months, no significant effect was observed. The authors concluded that a longer duration of the diet may be necessary to achieve the improvement of the ASD symptoms. Other problems documented in the trials carried out include transgressions of the diet, small sample size, the heterogeneity of the participants and the possibility of a placebo effect.
|
||||
|
||||
In the subset of people who have gluten sensitivity there is limited evidence that suggests that a gluten-free diet may improve some autistic behaviors.
|
||||
|
||||
There is tentative evidence that music therapy may improve social interactions, verbal communication, and non-verbal communication skills. There has been early research looking at hyperbaric treatments in children with autism. Studies on pet therapy have shown positive effects.
|
||||
|
||||
There is no known cure. The degree of symptoms can decrease, occasionally to the extent that people lose their diagnosis of ASD; this occurs sometimes after intensive treatment and sometimes not. It is not known how often recovery happens; reported rates in unselected samples have ranged from 3% to 25%. Most children with autism acquire language by age five or younger, though a few have developed communication skills in later years. Many children with autism lack social support, future employment opportunities or self-determination. Although core difficulties tend to persist, symptoms often become less severe with age.
|
||||
|
||||
Few high-quality studies address long-term prognosis. Some adults show modest improvement in communication skills, but a few decline; no study has focused on autism after midlife. Acquiring language before age six, having an IQ above 50, and having a marketable skill all predict better outcomes; independent living is unlikely with severe autism.
|
||||
|
||||
Many individuals with autism face significant obstacles in transitioning to adulthood. Compared to the general population individuals with autism are more likely to be unemployed and to have never had a job. About half of people in their 20s with autism are not employed.
|
||||
|
||||
Most recent reviews tend to estimate a prevalence of 1–2 per 1,000 for autism and close to 6 per 1,000 for ASD as of 2007. A 2016 survey in the United States reported a rate of 25 per 1,000 children for ASD. Globally, autism affects an estimated 24.8 million people , while Asperger syndrome affects a further 37.2 million. In 2012, the NHS estimated that the overall prevalence of autism among adults aged 18 years and over in the UK was 1.1%. Rates of PDD-NOS's has been estimated at 3.7 per 1,000, Asperger syndrome at roughly 0.6 per 1,000, and childhood disintegrative disorder at 0.02 per 1,000. CDC estimates about 1 out of 59 (1.7%) for 2014, an increase from 1 out of every 68 children (1.5%) for 2010.
|
||||
|
||||
The number of reported cases of autism increased dramatically in the 1990s and early 2000s. This increase is largely attributable to changes in diagnostic practices, referral patterns, availability of services, age at diagnosis, and public awareness, though unidentified environmental risk factors cannot be ruled out. The available evidence does not rule out the possibility that autism's true prevalence has increased; a real increase would suggest directing more attention and funding toward changing environmental factors instead of continuing to focus on genetics.
|
||||
|
||||
Boys are at higher risk for ASD than girls. The sex ratio averages 4.3:1 and is greatly modified by cognitive impairment: it may be close to 2:1 with intellectual disability and more than 5.5:1 without. Several theories about the higher prevalence in males have been investigated, but the cause of the difference is unconfirmed; one theory is that females are underdiagnosed.
|
||||
|
||||
Although the evidence does not implicate any single pregnancy-related risk factor as a cause of autism, the risk of autism is associated with advanced age in either parent, and with diabetes, bleeding, and use of psychiatric drugs in the mother during pregnancy. The risk is greater with older fathers than with older mothers; two potential explanations are the known increase in mutation burden in older sperm, and the hypothesis that men marry later if they carry genetic liability and show some signs of autism. Most professionals believe that race, ethnicity, and socioeconomic background do not affect the occurrence of autism.
|
||||
|
||||
Several other conditions are common in children with autism. They include:
|
||||
|
||||
A few examples of autistic symptoms and treatments were described long before autism was named. The "Table Talk" of Martin Luther, compiled by his notetaker, Mathesius, contains the story of a 12-year-old boy who may have been severely autistic. Luther reportedly thought the boy was a soulless mass of flesh possessed by the devil, and suggested that he be suffocated, although a later critic has cast doubt on the veracity of this report. The earliest well-documented case of autism is that of Hugh Blair of Borgue, as detailed in a 1747 court case in which his brother successfully petitioned to annul Blair's marriage to gain Blair's inheritance. The Wild Boy of Aveyron, a feral child caught in 1798, showed several signs of autism; the medical student Jean Itard treated him with a behavioral program designed to help him form social attachments and to induce speech via imitation.
|
||||
|
||||
The New Latin word "autismus" (English translation "autism") was coined by the Swiss psychiatrist Eugen Bleuler in 1910 as he was defining symptoms of schizophrenia. He derived it from the Greek word "autós" (αὐτός, meaning "self"), and used it to mean morbid self-admiration, referring to "autistic withdrawal of the patient to his fantasies, against which any influence from outside becomes an intolerable disturbance". A Soviet child psychiatrist, Grunya Sukhareva, described a similar syndrome that was published in Russian in 1925, and in German in 1926.
|
||||
|
||||
The word "autism" first took its modern sense in 1938 when Hans Asperger of the Vienna University Hospital adopted Bleuler's terminology "autistic psychopaths" in a lecture in German about child psychology. Asperger was investigating an ASD now known as Asperger syndrome, though for various reasons it was not widely recognized as a separate diagnosis until 1981. Leo Kanner of the Johns Hopkins Hospital first used "autism" in its modern sense in English when he introduced the label "early infantile autism" in a 1943 report of 11 children with striking behavioral similarities. Almost all the characteristics described in Kanner's first paper on the subject, notably "autistic aloneness" and "insistence on sameness", are still regarded as typical of the autistic spectrum of disorders. It is not known whether Kanner derived the term independently of Asperger.
|
||||
|
||||
Donald Triplett was the first person diagnosed with autism. He was diagnosed by Kanner after being first examined in 1938, and was labeled as "case 1". Triplett was noted for his savant abilities, particularly being able to name musical notes played on a piano and to mentally multiply numbers. His father, Oliver, described him as socially withdrawn but interested in number patterns, music notes, letters of the alphabet, and U.S. president pictures. By the age of 2, he had the ability to recite the 23rd Psalm and memorized 25 questions and answers from the Presbyterian catechism. He was also interested in creating musical chords.
|
||||
|
||||
Kanner's reuse of "autism" led to decades of confused terminology like "infantile schizophrenia", and child psychiatry's focus on maternal deprivation led to misconceptions of autism as an infant's response to "refrigerator mothers". Starting in the late 1960s autism was established as a separate syndrome.
|
||||
|
||||
As late as the mid-1970s there was little evidence of a genetic role in autism; while in 2007 it was believed to be one of the most heritable psychiatric conditions. Although the rise of parent organizations and the destigmatization of childhood ASD have affected how ASD is viewed, parents continue to feel social stigma in situations where their child's autistic behavior is perceived negatively, and many primary care physicians and medical specialists express some beliefs consistent with outdated autism research.
|
||||
|
||||
It took until 1980 for the DSM-III to differentiate autism from childhood schizophrenia. In 1987, the DSM-III-R provided a checklist for diagnosing autism. In May 2013, the DSM-5 was released, updating the classification for pervasive developmental disorders. The grouping of disorders, including PDD-NOS, autism, Asperger syndrome, Rett syndrome, and CDD, has been removed and replaced with the general term of Autism Spectrum Disorders. The two categories that exist are impaired social communication and/or interaction, and restricted and/or repetitive behaviors.
|
||||
|
||||
The Internet has helped autistic individuals bypass nonverbal cues and emotional sharing that they find difficult to deal with, and has given them a way to form online communities and work remotely. Societal and cultural aspects of autism have developed: some in the community seek a cure, while others believe that autism is simply another way of being.
|
||||
|
||||
An autistic culture has emerged, accompanied by the autistic rights and neurodiversity movements. Events include World Autism Awareness Day, Autism Sunday, Autistic Pride Day, Autreat, and others. Organizations dedicated to promoting awareness of autism include Autistic Self Advocacy Network, Aspies For Freedom, Autism National Committee, and Autism Society of America. At the same time, some organizations, including Autism Speaks, have been condemned by disability rights organizations for failing to support autistic people. Social-science scholars study those with autism in hopes to learn more about "autism as a culture, transcultural comparisons... and research on social movements." While most autistic individuals do not have savant skills, many have been successful in their fields.
|
||||
|
||||
The autism rights movement is a social movement within the context of disability rights that emphasizes the concept of neurodiversity, viewing the autism spectrum as a result of natural variations in the human brain rather than a disorder to be cured. The autism rights movement advocates for including greater acceptance of autistic behaviors; therapies that focus on coping skills rather than on imitating the behaviors of those without autism, and the recognition of the autistic community as a minority group. Autism rights or neurodiversity advocates believe that the autism spectrum is genetic and should be accepted as a natural expression of the human genome. This perspective is distinct from two other likewise distinct views: the medical perspective, that autism is caused by a genetic defect and should be addressed by targeting the autism gene(s), and fringe theories that autism is caused by environmental factors such as vaccines. A common criticism against autistic activists is that the majority of them are "high-functioning" or have Asperger syndrome and do not represent the views of "low-functioning" autistic people.
|
||||
|
||||
About half of autistics are unemployed, and one third of those with graduate degrees may be unemployed. Among autistics who find work, most are employed in sheltered settings working for wages below the national minimum. While employers state hiring concerns about productivity and supervision, experienced employers of autistics give positive reports of above average memory and detail orientation as well as a high regard for rules and procedure in autistic employees. A majority of the economic burden of autism is caused by decreased earnings in the job market. Some studies also find decreased earning among parents who care for autistic children.
|
||||
|
||||
|
||||
</doc>
|
||||
@@ -11,9 +11,11 @@ if is_torch_available():
|
||||
DataCollatorForLanguageModeling,
|
||||
DataCollatorForNextSentencePrediction,
|
||||
DataCollatorForPermutationLanguageModeling,
|
||||
DataCollatorForSOP,
|
||||
GlueDataset,
|
||||
GlueDataTrainingArguments,
|
||||
LineByLineTextDataset,
|
||||
LineByLineWithSOPTextDataset,
|
||||
TextDataset,
|
||||
TextDatasetForNextSentencePrediction,
|
||||
default_data_collator,
|
||||
@@ -21,6 +23,7 @@ if is_torch_available():
|
||||
|
||||
|
||||
PATH_SAMPLE_TEXT = "./tests/fixtures/sample_text.txt"
|
||||
PATH_SAMPLE_TEXT_DIR = "./tests/fixtures/tests_samples/wiki_text"
|
||||
|
||||
|
||||
@require_torch
|
||||
@@ -168,3 +171,19 @@ class DataCollatorIntegrationTest(unittest.TestCase):
|
||||
self.assertEqual(batch["token_type_ids"].shape, torch.Size((total_samples, 512)))
|
||||
self.assertEqual(batch["masked_lm_labels"].shape, torch.Size((total_samples, 512)))
|
||||
self.assertEqual(batch["next_sentence_label"].shape, torch.Size((total_samples,)))
|
||||
|
||||
def test_sop(self):
|
||||
tokenizer = AutoTokenizer.from_pretrained("albert-base-v2")
|
||||
data_collator = DataCollatorForSOP(tokenizer)
|
||||
|
||||
dataset = LineByLineWithSOPTextDataset(tokenizer, file_dir=PATH_SAMPLE_TEXT_DIR, block_size=512)
|
||||
examples = [dataset[i] for i in range(len(dataset))]
|
||||
batch = data_collator(examples)
|
||||
self.assertIsInstance(batch, dict)
|
||||
|
||||
# Since there are randomly generated false samples, the total number of samples is not fixed.
|
||||
total_samples = batch["input_ids"].shape[0]
|
||||
self.assertEqual(batch["input_ids"].shape, torch.Size((total_samples, 512)))
|
||||
self.assertEqual(batch["token_type_ids"].shape, torch.Size((total_samples, 512)))
|
||||
self.assertEqual(batch["labels"].shape, torch.Size((total_samples, 512)))
|
||||
self.assertEqual(batch["sentence_order_label"].shape, torch.Size((total_samples,)))
|
||||
@@ -0,0 +1,103 @@
|
||||
# coding=utf-8
|
||||
# Copyright 2020 The Hugging Face 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 unittest
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
from transformers.file_utils import ModelOutput
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModelOutputTest(ModelOutput):
|
||||
a: float
|
||||
b: Optional[float] = None
|
||||
c: Optional[float] = None
|
||||
|
||||
|
||||
class ModelOutputTester(unittest.TestCase):
|
||||
def test_get_attributes(self):
|
||||
x = ModelOutputTest(a=30)
|
||||
self.assertEqual(x.a, 30)
|
||||
self.assertIsNone(x.b)
|
||||
self.assertIsNone(x.c)
|
||||
with self.assertRaises(AttributeError):
|
||||
_ = x.d
|
||||
|
||||
def test_index_with_ints_and_slices(self):
|
||||
x = ModelOutputTest(a=30, b=10)
|
||||
self.assertEqual(x[0], 30)
|
||||
self.assertEqual(x[1], 10)
|
||||
self.assertEqual(x[:2], (30, 10))
|
||||
self.assertEqual(x[:], (30, 10))
|
||||
|
||||
x = ModelOutputTest(a=30, c=10)
|
||||
self.assertEqual(x[0], 30)
|
||||
self.assertEqual(x[1], 10)
|
||||
self.assertEqual(x[:2], (30, 10))
|
||||
self.assertEqual(x[:], (30, 10))
|
||||
|
||||
def test_index_with_strings(self):
|
||||
x = ModelOutputTest(a=30, b=10)
|
||||
self.assertEqual(x["a"], 30)
|
||||
self.assertEqual(x["b"], 10)
|
||||
with self.assertRaises(KeyError):
|
||||
_ = x["c"]
|
||||
|
||||
x = ModelOutputTest(a=30, c=10)
|
||||
self.assertEqual(x["a"], 30)
|
||||
self.assertEqual(x["c"], 10)
|
||||
with self.assertRaises(KeyError):
|
||||
_ = x["b"]
|
||||
|
||||
def test_dict_like_properties(self):
|
||||
x = ModelOutputTest(a=30)
|
||||
self.assertEqual(list(x.keys()), ["a"])
|
||||
self.assertEqual(list(x.values()), [30])
|
||||
self.assertEqual(list(x.items()), [("a", 30)])
|
||||
self.assertEqual(list(x), ["a"])
|
||||
|
||||
x = ModelOutputTest(a=30, b=10)
|
||||
self.assertEqual(list(x.keys()), ["a", "b"])
|
||||
self.assertEqual(list(x.values()), [30, 10])
|
||||
self.assertEqual(list(x.items()), [("a", 30), ("b", 10)])
|
||||
self.assertEqual(list(x), ["a", "b"])
|
||||
|
||||
x = ModelOutputTest(a=30, c=10)
|
||||
self.assertEqual(list(x.keys()), ["a", "c"])
|
||||
self.assertEqual(list(x.values()), [30, 10])
|
||||
self.assertEqual(list(x.items()), [("a", 30), ("c", 10)])
|
||||
self.assertEqual(list(x), ["a", "c"])
|
||||
|
||||
with self.assertRaises(Exception):
|
||||
x = x.update({"d": 20})
|
||||
with self.assertRaises(Exception):
|
||||
del x["a"]
|
||||
with self.assertRaises(Exception):
|
||||
_ = x.pop("a")
|
||||
with self.assertRaises(Exception):
|
||||
_ = x.setdefault("d", 32)
|
||||
|
||||
def test_set_attributes(self):
|
||||
x = ModelOutputTest(a=30)
|
||||
x.a = 10
|
||||
self.assertEqual(x.a, 10)
|
||||
self.assertEqual(x["a"], 10)
|
||||
|
||||
def test_set_keys(self):
|
||||
x = ModelOutputTest(a=30)
|
||||
x["a"] = 10
|
||||
self.assertEqual(x.a, 10)
|
||||
self.assertEqual(x["a"], 10)
|
||||
@@ -453,83 +453,6 @@ class BertModelTest(ModelTesterMixin, unittest.TestCase):
|
||||
config_and_inputs = self.model_tester.prepare_config_and_inputs()
|
||||
self.model_tester.create_and_check_for_token_classification(*config_and_inputs)
|
||||
|
||||
# Copied from test_modeling_common.test_torchscript, but using jit.script, not jit.trace
|
||||
def test_full_torchscript(self):
|
||||
import copy
|
||||
import os
|
||||
import tempfile
|
||||
|
||||
if is_torch_available():
|
||||
import torch
|
||||
|
||||
from transformers.modeling_bert import (
|
||||
BertScriptableForMultipleChoice,
|
||||
BertScriptableForNextSentencePrediction,
|
||||
BertScriptableForPreTraining,
|
||||
BertScriptableForQuestionAnswering,
|
||||
BertScriptableForSequenceClassification,
|
||||
BertScriptableForTokenClassification,
|
||||
BertScriptableModel,
|
||||
)
|
||||
|
||||
config, unused_inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()
|
||||
config.return_dict = False
|
||||
configs_no_init = copy.deepcopy(config)
|
||||
for key in configs_no_init.__dict__.keys():
|
||||
if "_range" in key or "_std" in key or "initializer_factor" in key:
|
||||
setattr(configs_no_init, key, 1e-10)
|
||||
|
||||
scriptable_model_classes = (
|
||||
BertScriptableModel,
|
||||
BertScriptableForMultipleChoice,
|
||||
BertScriptableForNextSentencePrediction,
|
||||
BertScriptableForPreTraining,
|
||||
BertScriptableForQuestionAnswering,
|
||||
BertScriptableForSequenceClassification,
|
||||
BertScriptableForTokenClassification,
|
||||
)
|
||||
for model_class in scriptable_model_classes:
|
||||
model = model_class(config=configs_no_init)
|
||||
model.to(torch_device)
|
||||
model.eval()
|
||||
|
||||
try:
|
||||
scripted = torch.jit.script(model)
|
||||
except RuntimeError:
|
||||
self.fail("Couldn't script module.")
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmp_dir_name:
|
||||
pt_file_name = os.path.join(tmp_dir_name, "scripted_model.pt")
|
||||
|
||||
try:
|
||||
torch.jit.save(scripted, pt_file_name)
|
||||
except Exception:
|
||||
self.fail("Couldn't save scripted module.")
|
||||
|
||||
try:
|
||||
loaded_model = torch.jit.load(pt_file_name)
|
||||
except Exception:
|
||||
self.fail("Couldn't load scripted module.")
|
||||
|
||||
model.to(torch_device)
|
||||
model.eval()
|
||||
|
||||
loaded_model.to(torch_device)
|
||||
loaded_model.eval()
|
||||
|
||||
model_state_dict = model.state_dict()
|
||||
loaded_model_state_dict = loaded_model.state_dict()
|
||||
|
||||
self.assertEqual(set(model_state_dict.keys()), set(loaded_model_state_dict.keys()))
|
||||
|
||||
models_equal = True
|
||||
for layer_name, p1 in model_state_dict.items():
|
||||
p2 = loaded_model_state_dict[layer_name]
|
||||
if p1.data.ne(p2.data).sum() > 0:
|
||||
models_equal = False
|
||||
|
||||
self.assertTrue(models_equal)
|
||||
|
||||
@slow
|
||||
def test_model_from_pretrained(self):
|
||||
for model_name in BERT_PRETRAINED_MODEL_ARCHIVE_LIST[:1]:
|
||||
|
||||
Executable
+234
@@ -0,0 +1,234 @@
|
||||
# coding=utf-8
|
||||
# Copyright 2018 The Google AI Language Team Authors.
|
||||
#
|
||||
# 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 unittest
|
||||
|
||||
from transformers import is_torch_available
|
||||
from transformers.testing_utils import require_torch, slow, torch_device
|
||||
|
||||
from .test_configuration_common import ConfigTester
|
||||
from .test_modeling_common import ModelTesterMixin, floats_tensor, ids_tensor, random_attention_mask
|
||||
|
||||
|
||||
if is_torch_available():
|
||||
from transformers import BertGenerationConfig, BertGenerationDecoder, BertGenerationEncoder
|
||||
|
||||
|
||||
class BertGenerationEncoderTester:
|
||||
def __init__(
|
||||
self,
|
||||
parent,
|
||||
batch_size=13,
|
||||
seq_length=7,
|
||||
is_training=True,
|
||||
use_input_mask=True,
|
||||
vocab_size=99,
|
||||
hidden_size=32,
|
||||
num_hidden_layers=5,
|
||||
num_attention_heads=4,
|
||||
intermediate_size=37,
|
||||
hidden_act="gelu",
|
||||
hidden_dropout_prob=0.1,
|
||||
attention_probs_dropout_prob=0.1,
|
||||
max_position_embeddings=50,
|
||||
initializer_range=0.02,
|
||||
use_labels=True,
|
||||
scope=None,
|
||||
):
|
||||
self.parent = parent
|
||||
self.batch_size = batch_size
|
||||
self.seq_length = seq_length
|
||||
self.is_training = is_training
|
||||
self.use_input_mask = use_input_mask
|
||||
self.vocab_size = vocab_size
|
||||
self.hidden_size = hidden_size
|
||||
self.num_hidden_layers = num_hidden_layers
|
||||
self.num_attention_heads = num_attention_heads
|
||||
self.intermediate_size = intermediate_size
|
||||
self.hidden_act = hidden_act
|
||||
self.hidden_dropout_prob = hidden_dropout_prob
|
||||
self.attention_probs_dropout_prob = attention_probs_dropout_prob
|
||||
self.max_position_embeddings = max_position_embeddings
|
||||
self.initializer_range = initializer_range
|
||||
self.use_labels = use_labels
|
||||
self.scope = scope
|
||||
|
||||
def prepare_config_and_inputs(self):
|
||||
input_ids = ids_tensor([self.batch_size, self.seq_length], self.vocab_size)
|
||||
|
||||
input_mask = None
|
||||
if self.use_input_mask:
|
||||
input_mask = random_attention_mask([self.batch_size, self.seq_length])
|
||||
|
||||
if self.use_labels:
|
||||
token_labels = ids_tensor([self.batch_size, self.seq_length], self.vocab_size)
|
||||
|
||||
config = BertGenerationConfig(
|
||||
vocab_size=self.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,
|
||||
is_decoder=False,
|
||||
initializer_range=self.initializer_range,
|
||||
return_dict=True,
|
||||
)
|
||||
|
||||
return config, input_ids, input_mask, token_labels
|
||||
|
||||
def prepare_config_and_inputs_for_decoder(self):
|
||||
(
|
||||
config,
|
||||
input_ids,
|
||||
input_mask,
|
||||
token_labels,
|
||||
) = self.prepare_config_and_inputs()
|
||||
|
||||
config.is_decoder = True
|
||||
encoder_hidden_states = floats_tensor([self.batch_size, self.seq_length, self.hidden_size])
|
||||
encoder_attention_mask = ids_tensor([self.batch_size, self.seq_length], vocab_size=2)
|
||||
|
||||
return (
|
||||
config,
|
||||
input_ids,
|
||||
input_mask,
|
||||
token_labels,
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
)
|
||||
|
||||
def create_and_check_model(
|
||||
self,
|
||||
config,
|
||||
input_ids,
|
||||
input_mask,
|
||||
token_labels,
|
||||
**kwargs,
|
||||
):
|
||||
model = BertGenerationEncoder(config=config)
|
||||
model.to(torch_device)
|
||||
model.eval()
|
||||
result = model(input_ids, attention_mask=input_mask)
|
||||
result = model(input_ids)
|
||||
self.parent.assertEqual(result.last_hidden_state.shape, (self.batch_size, self.seq_length, self.hidden_size))
|
||||
|
||||
def create_and_check_model_as_decoder(
|
||||
self,
|
||||
config,
|
||||
input_ids,
|
||||
input_mask,
|
||||
token_labels,
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
**kwargs,
|
||||
):
|
||||
config.add_cross_attention = True
|
||||
model = BertGenerationEncoder(config=config)
|
||||
model.to(torch_device)
|
||||
model.eval()
|
||||
result = model(
|
||||
input_ids,
|
||||
attention_mask=input_mask,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
encoder_attention_mask=encoder_attention_mask,
|
||||
)
|
||||
result = model(
|
||||
input_ids,
|
||||
attention_mask=input_mask,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
)
|
||||
self.parent.assertEqual(result.last_hidden_state.shape, (self.batch_size, self.seq_length, self.hidden_size))
|
||||
|
||||
def create_and_check_for_causal_lm(
|
||||
self,
|
||||
config,
|
||||
input_ids,
|
||||
input_mask,
|
||||
token_labels,
|
||||
*args,
|
||||
):
|
||||
model = BertGenerationDecoder(config)
|
||||
model.to(torch_device)
|
||||
model.eval()
|
||||
result = model(input_ids, attention_mask=input_mask, labels=token_labels)
|
||||
self.parent.assertEqual(result.logits.shape, (self.batch_size, self.seq_length, self.vocab_size))
|
||||
|
||||
def prepare_config_and_inputs_for_common(self):
|
||||
config_and_inputs = self.prepare_config_and_inputs()
|
||||
(
|
||||
config,
|
||||
input_ids,
|
||||
input_mask,
|
||||
token_labels,
|
||||
) = config_and_inputs
|
||||
inputs_dict = {"input_ids": input_ids, "attention_mask": input_mask}
|
||||
return config, inputs_dict
|
||||
|
||||
|
||||
@require_torch
|
||||
class BertGenerationEncoderTest(ModelTesterMixin, unittest.TestCase):
|
||||
|
||||
all_model_classes = (BertGenerationEncoder, BertGenerationDecoder) if is_torch_available() else ()
|
||||
|
||||
def setUp(self):
|
||||
self.model_tester = BertGenerationEncoderTester(self)
|
||||
self.config_tester = ConfigTester(self, config_class=BertGenerationConfig, hidden_size=37)
|
||||
|
||||
def test_config(self):
|
||||
self.config_tester.run_common_tests()
|
||||
|
||||
def test_model(self):
|
||||
config_and_inputs = self.model_tester.prepare_config_and_inputs()
|
||||
self.model_tester.create_and_check_model(*config_and_inputs)
|
||||
|
||||
def test_model_as_decoder(self):
|
||||
config_and_inputs = self.model_tester.prepare_config_and_inputs_for_decoder()
|
||||
self.model_tester.create_and_check_model_as_decoder(*config_and_inputs)
|
||||
|
||||
def test_model_as_decoder_with_default_input_mask(self):
|
||||
# This regression test was failing with PyTorch < 1.3
|
||||
(
|
||||
config,
|
||||
input_ids,
|
||||
input_mask,
|
||||
token_labels,
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
) = self.model_tester.prepare_config_and_inputs_for_decoder()
|
||||
|
||||
input_mask = None
|
||||
|
||||
self.model_tester.create_and_check_model_as_decoder(
|
||||
config,
|
||||
input_ids,
|
||||
input_mask,
|
||||
token_labels,
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
)
|
||||
|
||||
def test_for_causal_lm(self):
|
||||
config_and_inputs = self.model_tester.prepare_config_and_inputs_for_decoder()
|
||||
self.model_tester.create_and_check_for_causal_lm(*config_and_inputs)
|
||||
|
||||
@slow
|
||||
def test_model_from_pretrained(self):
|
||||
model = BertGenerationEncoder.from_pretrained("google/bert_for_seq_generation_L-24_bbc_encoder")
|
||||
self.assertIsNotNone(model)
|
||||
@@ -21,6 +21,7 @@ from transformers import is_torch_available
|
||||
from transformers.testing_utils import require_torch, slow, torch_device
|
||||
|
||||
from .test_modeling_bert import BertModelTester
|
||||
from .test_modeling_bert_generation import BertGenerationEncoderTester
|
||||
from .test_modeling_common import ids_tensor
|
||||
from .test_modeling_gpt2 import GPT2ModelTester
|
||||
from .test_modeling_roberta import RobertaModelTester
|
||||
@@ -31,6 +32,9 @@ if is_torch_available():
|
||||
import torch
|
||||
|
||||
from transformers import (
|
||||
AutoTokenizer,
|
||||
BertGenerationDecoder,
|
||||
BertGenerationEncoder,
|
||||
BertLMHeadModel,
|
||||
BertModel,
|
||||
BertTokenizer,
|
||||
@@ -489,6 +493,67 @@ class BertEncoderDecoderModelTest(EncoderDecoderMixin, unittest.TestCase):
|
||||
self.assertEqual(summary, EXPECTED_SUMMARY)
|
||||
|
||||
|
||||
class BertGenerationEncoderDecoderModelTest(EncoderDecoderMixin, unittest.TestCase):
|
||||
def get_pretrained_model(self):
|
||||
return EncoderDecoderModel.from_encoder_decoder_pretrained(
|
||||
"google/bert_for_seq_generation_L-24_bbc_encoder", "google/bert_for_seq_generation_L-24_bbc_encoder"
|
||||
)
|
||||
|
||||
def get_encoder_decoder_model(self, config, decoder_config):
|
||||
encoder_model = BertGenerationEncoder(config)
|
||||
decoder_model = BertGenerationDecoder(decoder_config)
|
||||
return encoder_model, decoder_model
|
||||
|
||||
def prepare_config_and_inputs(self):
|
||||
model_tester = BertGenerationEncoderTester(self)
|
||||
encoder_config_and_inputs = model_tester.prepare_config_and_inputs()
|
||||
decoder_config_and_inputs = model_tester.prepare_config_and_inputs_for_decoder()
|
||||
(
|
||||
config,
|
||||
input_ids,
|
||||
input_mask,
|
||||
token_labels,
|
||||
) = encoder_config_and_inputs
|
||||
(
|
||||
decoder_config,
|
||||
decoder_input_ids,
|
||||
decoder_input_mask,
|
||||
decoder_token_labels,
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
) = decoder_config_and_inputs
|
||||
|
||||
# make sure that cross attention layers are added
|
||||
decoder_config.add_cross_attention = True
|
||||
return {
|
||||
"config": config,
|
||||
"input_ids": input_ids,
|
||||
"attention_mask": input_mask,
|
||||
"decoder_config": decoder_config,
|
||||
"decoder_input_ids": decoder_input_ids,
|
||||
"decoder_attention_mask": decoder_input_mask,
|
||||
"decoder_token_labels": decoder_token_labels,
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"labels": decoder_token_labels,
|
||||
}
|
||||
|
||||
@slow
|
||||
def test_roberta2roberta_summarization(self):
|
||||
model = EncoderDecoderModel.from_pretrained("google/roberta2roberta_L-24_bbc")
|
||||
model.to(torch_device)
|
||||
tokenizer = AutoTokenizer.from_pretrained("google/roberta2roberta_L-24_bbc")
|
||||
|
||||
ARTICLE = """The problem is affecting people using the older versions of the PlayStation 3, called the "Fat" model.The problem isn't affecting the newer PS3 Slim systems that have been on sale since September last year.Sony have also said they are aiming to have the problem fixed shortly but is advising some users to avoid using their console for the time being."We hope to resolve this problem within the next 24 hours," a statement reads. "In the meantime, if you have a model other than the new slim PS3, we advise that you do not use your PS3 system, as doing so may result in errors in some functionality, such as recording obtained trophies, and not being able to restore certain data."We believe we have identified that this problem is being caused by a bug in the clock functionality incorporated in the system."The PlayStation Network is used by millions of people around the world.It allows users to play their friends at games like Fifa over the internet and also do things like download software or visit online stores."""
|
||||
|
||||
EXPECTED_SUMMARY = """Sony has said that a bug in its PlayStation 3 console is preventing them from using the machine as a computer."""
|
||||
|
||||
input_ids = tokenizer(ARTICLE, return_tensors="pt").input_ids.to(torch_device)
|
||||
output_ids = model.generate(input_ids)
|
||||
summary = tokenizer.decode(output_ids[0], skip_special_tokens=True)
|
||||
|
||||
self.assertEqual(summary, EXPECTED_SUMMARY)
|
||||
|
||||
|
||||
class RoBertaEncoderDecoderModelTest(EncoderDecoderMixin, unittest.TestCase):
|
||||
def get_encoder_decoder_model(self, config, decoder_config):
|
||||
encoder_model = RobertaModel(config)
|
||||
|
||||
@@ -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)
|
||||
@@ -487,7 +487,10 @@ class TFModelTesterMixin:
|
||||
model = model_class(config)
|
||||
outputs = model(self._prepare_for_class(inputs_dict, model_class))
|
||||
hidden_states = [t.numpy() for t in outputs[-1]]
|
||||
self.assertEqual(len(hidden_states), self.model_tester.num_hidden_layers + 1)
|
||||
expected_num_layers = getattr(
|
||||
self.model_tester, "expected_num_hidden_layers", self.model_tester.num_hidden_layers + 1
|
||||
)
|
||||
self.assertEqual(len(hidden_states), expected_num_layers)
|
||||
self.assertListEqual(
|
||||
list(hidden_states[0].shape[-2:]),
|
||||
[self.model_tester.seq_length, self.model_tester.hidden_size],
|
||||
|
||||
@@ -0,0 +1,394 @@
|
||||
# coding=utf-8
|
||||
# Copyright 2020 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 unittest
|
||||
|
||||
from transformers import FunnelConfig, is_tf_available
|
||||
from transformers.testing_utils import require_tf
|
||||
|
||||
from .test_configuration_common import ConfigTester
|
||||
from .test_modeling_tf_common import TFModelTesterMixin, ids_tensor
|
||||
|
||||
|
||||
if is_tf_available():
|
||||
import tensorflow as tf
|
||||
|
||||
from transformers.modeling_tf_funnel import (
|
||||
TFFunnelBaseModel,
|
||||
TFFunnelForMaskedLM,
|
||||
TFFunnelForMultipleChoice,
|
||||
TFFunnelForPreTraining,
|
||||
TFFunnelForQuestionAnswering,
|
||||
TFFunnelForSequenceClassification,
|
||||
TFFunnelForTokenClassification,
|
||||
TFFunnelModel,
|
||||
)
|
||||
|
||||
|
||||
class TFFunnelModelTester:
|
||||
"""You can also import this e.g, from .test_modeling_funnel import FunnelModelTester """
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
parent,
|
||||
batch_size=13,
|
||||
seq_length=7,
|
||||
is_training=True,
|
||||
use_input_mask=True,
|
||||
use_token_type_ids=True,
|
||||
use_labels=True,
|
||||
vocab_size=99,
|
||||
block_sizes=[1, 1, 2],
|
||||
num_decoder_layers=1,
|
||||
d_model=32,
|
||||
n_head=4,
|
||||
d_head=8,
|
||||
d_inner=37,
|
||||
hidden_act="gelu_new",
|
||||
hidden_dropout=0.1,
|
||||
attention_dropout=0.1,
|
||||
activation_dropout=0.0,
|
||||
max_position_embeddings=512,
|
||||
type_vocab_size=3,
|
||||
num_labels=3,
|
||||
num_choices=4,
|
||||
scope=None,
|
||||
base=False,
|
||||
):
|
||||
self.parent = parent
|
||||
self.batch_size = batch_size
|
||||
self.seq_length = seq_length
|
||||
self.is_training = is_training
|
||||
self.use_input_mask = use_input_mask
|
||||
self.use_token_type_ids = use_token_type_ids
|
||||
self.use_labels = use_labels
|
||||
self.vocab_size = vocab_size
|
||||
self.block_sizes = block_sizes
|
||||
self.num_decoder_layers = num_decoder_layers
|
||||
self.d_model = d_model
|
||||
self.n_head = n_head
|
||||
self.d_head = d_head
|
||||
self.d_inner = d_inner
|
||||
self.hidden_act = hidden_act
|
||||
self.hidden_dropout = hidden_dropout
|
||||
self.attention_dropout = attention_dropout
|
||||
self.activation_dropout = activation_dropout
|
||||
self.max_position_embeddings = max_position_embeddings
|
||||
self.type_vocab_size = type_vocab_size
|
||||
self.type_sequence_label_size = 2
|
||||
self.num_labels = num_labels
|
||||
self.num_choices = num_choices
|
||||
self.scope = scope
|
||||
|
||||
# Used in the tests to check the size of the first attention layer
|
||||
self.num_attention_heads = n_head
|
||||
# Used in the tests to check the size of the first hidden state
|
||||
self.hidden_size = self.d_model
|
||||
# Used in the tests to check the number of output hidden states/attentions
|
||||
self.num_hidden_layers = sum(self.block_sizes) + (0 if base else self.num_decoder_layers)
|
||||
# FunnelModel adds two hidden layers: input embeddings and the sum of the upsampled encoder hidden state with
|
||||
# the last hidden state of the first block (which is the first hidden state of the decoder).
|
||||
if not base:
|
||||
self.expected_num_hidden_layers = self.num_hidden_layers + 2
|
||||
|
||||
def prepare_config_and_inputs(self):
|
||||
input_ids = ids_tensor([self.batch_size, self.seq_length], self.vocab_size)
|
||||
|
||||
input_mask = None
|
||||
if self.use_input_mask:
|
||||
input_mask = ids_tensor([self.batch_size, self.seq_length], vocab_size=2)
|
||||
|
||||
token_type_ids = None
|
||||
if self.use_token_type_ids:
|
||||
token_type_ids = ids_tensor([self.batch_size, self.seq_length], self.type_vocab_size)
|
||||
|
||||
sequence_labels = None
|
||||
token_labels = None
|
||||
choice_labels = None
|
||||
if self.use_labels:
|
||||
sequence_labels = ids_tensor([self.batch_size], self.type_sequence_label_size)
|
||||
token_labels = ids_tensor([self.batch_size, self.seq_length], self.num_labels)
|
||||
choice_labels = ids_tensor([self.batch_size], self.num_choices)
|
||||
|
||||
config = FunnelConfig(
|
||||
vocab_size=self.vocab_size,
|
||||
block_sizes=self.block_sizes,
|
||||
num_decoder_layers=self.num_decoder_layers,
|
||||
d_model=self.d_model,
|
||||
n_head=self.n_head,
|
||||
d_head=self.d_head,
|
||||
d_inner=self.d_inner,
|
||||
hidden_act=self.hidden_act,
|
||||
hidden_dropout=self.hidden_dropout,
|
||||
attention_dropout=self.attention_dropout,
|
||||
activation_dropout=self.activation_dropout,
|
||||
max_position_embeddings=self.max_position_embeddings,
|
||||
type_vocab_size=self.type_vocab_size,
|
||||
return_dict=True,
|
||||
)
|
||||
|
||||
return (
|
||||
config,
|
||||
input_ids,
|
||||
token_type_ids,
|
||||
input_mask,
|
||||
sequence_labels,
|
||||
token_labels,
|
||||
choice_labels,
|
||||
)
|
||||
|
||||
def create_and_check_model(
|
||||
self,
|
||||
config,
|
||||
input_ids,
|
||||
token_type_ids,
|
||||
input_mask,
|
||||
sequence_labels,
|
||||
token_labels,
|
||||
choice_labels,
|
||||
):
|
||||
model = TFFunnelModel(config=config)
|
||||
inputs = {"input_ids": input_ids, "attention_mask": input_mask, "token_type_ids": token_type_ids}
|
||||
result = model(inputs)
|
||||
|
||||
inputs = [input_ids, input_mask]
|
||||
result = model(inputs)
|
||||
|
||||
result = model(input_ids)
|
||||
self.parent.assertEqual(result.last_hidden_state.shape, (self.batch_size, self.seq_length, self.d_model))
|
||||
|
||||
config.truncate_seq = False
|
||||
model = TFFunnelModel(config=config)
|
||||
result = model(input_ids)
|
||||
self.parent.assertEqual(result.last_hidden_state.shape, (self.batch_size, self.seq_length, self.d_model))
|
||||
|
||||
config.separate_cls = False
|
||||
model = TFFunnelModel(config=config)
|
||||
result = model(input_ids)
|
||||
self.parent.assertEqual(result.last_hidden_state.shape, (self.batch_size, self.seq_length, self.d_model))
|
||||
|
||||
def create_and_check_base_model(
|
||||
self,
|
||||
config,
|
||||
input_ids,
|
||||
token_type_ids,
|
||||
input_mask,
|
||||
sequence_labels,
|
||||
token_labels,
|
||||
choice_labels,
|
||||
):
|
||||
model = TFFunnelBaseModel(config=config)
|
||||
inputs = {"input_ids": input_ids, "attention_mask": input_mask, "token_type_ids": token_type_ids}
|
||||
result = model(inputs)
|
||||
|
||||
inputs = [input_ids, input_mask]
|
||||
result = model(inputs)
|
||||
|
||||
result = model(input_ids)
|
||||
self.parent.assertEqual(result.last_hidden_state.shape, (self.batch_size, 2, self.d_model))
|
||||
|
||||
config.truncate_seq = False
|
||||
model = TFFunnelBaseModel(config=config)
|
||||
result = model(input_ids)
|
||||
self.parent.assertEqual(result.last_hidden_state.shape, (self.batch_size, 3, self.d_model))
|
||||
|
||||
config.separate_cls = False
|
||||
model = TFFunnelBaseModel(config=config)
|
||||
result = model(input_ids)
|
||||
self.parent.assertEqual(result.last_hidden_state.shape, (self.batch_size, 2, self.d_model))
|
||||
|
||||
def create_and_check_for_pretraining(
|
||||
self,
|
||||
config,
|
||||
input_ids,
|
||||
token_type_ids,
|
||||
input_mask,
|
||||
sequence_labels,
|
||||
token_labels,
|
||||
choice_labels,
|
||||
):
|
||||
model = TFFunnelForPreTraining(config=config)
|
||||
inputs = {"input_ids": input_ids, "attention_mask": input_mask, "token_type_ids": token_type_ids}
|
||||
result = model(inputs)
|
||||
self.parent.assertEqual(result.logits.shape, (self.batch_size, self.seq_length))
|
||||
|
||||
def create_and_check_for_masked_lm(
|
||||
self,
|
||||
config,
|
||||
input_ids,
|
||||
token_type_ids,
|
||||
input_mask,
|
||||
sequence_labels,
|
||||
token_labels,
|
||||
choice_labels,
|
||||
):
|
||||
model = TFFunnelForMaskedLM(config=config)
|
||||
inputs = {"input_ids": input_ids, "attention_mask": input_mask, "token_type_ids": token_type_ids}
|
||||
result = model(inputs)
|
||||
self.parent.assertEqual(result.logits.shape, (self.batch_size, self.seq_length, self.vocab_size))
|
||||
|
||||
def create_and_check_for_sequence_classification(
|
||||
self,
|
||||
config,
|
||||
input_ids,
|
||||
token_type_ids,
|
||||
input_mask,
|
||||
sequence_labels,
|
||||
token_labels,
|
||||
choice_labels,
|
||||
):
|
||||
config.num_labels = self.num_labels
|
||||
model = TFFunnelForSequenceClassification(config=config)
|
||||
inputs = {"input_ids": input_ids, "attention_mask": input_mask, "token_type_ids": token_type_ids}
|
||||
result = model(inputs)
|
||||
self.parent.assertEqual(result.logits.shape, (self.batch_size, self.num_labels))
|
||||
|
||||
def create_and_check_for_multiple_choice(
|
||||
self,
|
||||
config,
|
||||
input_ids,
|
||||
token_type_ids,
|
||||
input_mask,
|
||||
sequence_labels,
|
||||
token_labels,
|
||||
choice_labels,
|
||||
):
|
||||
config.num_choices = self.num_choices
|
||||
model = TFFunnelForMultipleChoice(config=config)
|
||||
multiple_choice_inputs_ids = tf.tile(tf.expand_dims(input_ids, 1), (1, self.num_choices, 1))
|
||||
multiple_choice_input_mask = tf.tile(tf.expand_dims(input_mask, 1), (1, self.num_choices, 1))
|
||||
multiple_choice_token_type_ids = tf.tile(tf.expand_dims(token_type_ids, 1), (1, self.num_choices, 1))
|
||||
inputs = {
|
||||
"input_ids": multiple_choice_inputs_ids,
|
||||
"attention_mask": multiple_choice_input_mask,
|
||||
"token_type_ids": multiple_choice_token_type_ids,
|
||||
}
|
||||
result = model(inputs)
|
||||
self.parent.assertEqual(result.logits.shape, (self.batch_size, self.num_choices))
|
||||
|
||||
def create_and_check_for_token_classification(
|
||||
self,
|
||||
config,
|
||||
input_ids,
|
||||
token_type_ids,
|
||||
input_mask,
|
||||
sequence_labels,
|
||||
token_labels,
|
||||
choice_labels,
|
||||
):
|
||||
config.num_labels = self.num_labels
|
||||
model = TFFunnelForTokenClassification(config=config)
|
||||
inputs = {"input_ids": input_ids, "attention_mask": input_mask, "token_type_ids": token_type_ids}
|
||||
result = model(inputs)
|
||||
self.parent.assertEqual(result.logits.shape, (self.batch_size, self.seq_length, self.num_labels))
|
||||
|
||||
def create_and_check_for_question_answering(
|
||||
self,
|
||||
config,
|
||||
input_ids,
|
||||
token_type_ids,
|
||||
input_mask,
|
||||
sequence_labels,
|
||||
token_labels,
|
||||
choice_labels,
|
||||
):
|
||||
model = TFFunnelForQuestionAnswering(config=config)
|
||||
inputs = {"input_ids": input_ids, "attention_mask": input_mask, "token_type_ids": token_type_ids}
|
||||
result = model(inputs)
|
||||
self.parent.assertEqual(result.start_logits.shape, (self.batch_size, self.seq_length))
|
||||
self.parent.assertEqual(result.end_logits.shape, (self.batch_size, self.seq_length))
|
||||
|
||||
def prepare_config_and_inputs_for_common(self):
|
||||
config_and_inputs = self.prepare_config_and_inputs()
|
||||
(
|
||||
config,
|
||||
input_ids,
|
||||
token_type_ids,
|
||||
input_mask,
|
||||
sequence_labels,
|
||||
token_labels,
|
||||
choice_labels,
|
||||
) = config_and_inputs
|
||||
inputs_dict = {"input_ids": input_ids, "token_type_ids": token_type_ids, "attention_mask": input_mask}
|
||||
return config, inputs_dict
|
||||
|
||||
|
||||
@require_tf
|
||||
class FunnelModelTest(TFModelTesterMixin, unittest.TestCase):
|
||||
all_model_classes = (
|
||||
(
|
||||
TFFunnelModel,
|
||||
TFFunnelForMaskedLM,
|
||||
TFFunnelForPreTraining,
|
||||
TFFunnelForQuestionAnswering,
|
||||
TFFunnelForTokenClassification,
|
||||
)
|
||||
if is_tf_available()
|
||||
else ()
|
||||
)
|
||||
|
||||
def setUp(self):
|
||||
self.model_tester = TFFunnelModelTester(self)
|
||||
self.config_tester = ConfigTester(self, config_class=FunnelConfig)
|
||||
|
||||
def test_config(self):
|
||||
self.config_tester.run_common_tests()
|
||||
|
||||
def test_model(self):
|
||||
config_and_inputs = self.model_tester.prepare_config_and_inputs()
|
||||
self.model_tester.create_and_check_model(*config_and_inputs)
|
||||
|
||||
def test_for_pretraining(self):
|
||||
config_and_inputs = self.model_tester.prepare_config_and_inputs()
|
||||
self.model_tester.create_and_check_for_pretraining(*config_and_inputs)
|
||||
|
||||
def test_for_masked_lm(self):
|
||||
config_and_inputs = self.model_tester.prepare_config_and_inputs()
|
||||
self.model_tester.create_and_check_for_masked_lm(*config_and_inputs)
|
||||
|
||||
def test_for_token_classification(self):
|
||||
config_and_inputs = self.model_tester.prepare_config_and_inputs()
|
||||
self.model_tester.create_and_check_for_token_classification(*config_and_inputs)
|
||||
|
||||
def test_for_question_answering(self):
|
||||
config_and_inputs = self.model_tester.prepare_config_and_inputs()
|
||||
self.model_tester.create_and_check_for_question_answering(*config_and_inputs)
|
||||
|
||||
|
||||
@require_tf
|
||||
class TFFunnelBaseModelTest(TFModelTesterMixin, unittest.TestCase):
|
||||
all_model_classes = (
|
||||
(TFFunnelBaseModel, TFFunnelForMultipleChoice, TFFunnelForSequenceClassification) if is_tf_available() else ()
|
||||
)
|
||||
|
||||
def setUp(self):
|
||||
self.model_tester = TFFunnelModelTester(self, base=True)
|
||||
self.config_tester = ConfigTester(self, config_class=FunnelConfig)
|
||||
|
||||
def test_config(self):
|
||||
self.config_tester.run_common_tests()
|
||||
|
||||
def test_base_model(self):
|
||||
config_and_inputs = self.model_tester.prepare_config_and_inputs()
|
||||
self.model_tester.create_and_check_base_model(*config_and_inputs)
|
||||
|
||||
def test_for_sequence_classification(self):
|
||||
config_and_inputs = self.model_tester.prepare_config_and_inputs()
|
||||
self.model_tester.create_and_check_for_sequence_classification(*config_and_inputs)
|
||||
|
||||
def test_for_multiple_choice(self):
|
||||
config_and_inputs = self.model_tester.prepare_config_and_inputs()
|
||||
self.model_tester.create_and_check_for_multiple_choice(*config_and_inputs)
|
||||
@@ -0,0 +1,210 @@
|
||||
# coding=utf-8
|
||||
# Copyright 2020 The Google AI Language Team Authors.
|
||||
#
|
||||
# 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 os
|
||||
import unittest
|
||||
|
||||
from transformers.file_utils import cached_property
|
||||
from transformers.testing_utils import require_torch, slow
|
||||
from transformers.tokenization_bert_generation import BertGenerationTokenizer
|
||||
|
||||
from .test_tokenization_common import TokenizerTesterMixin
|
||||
|
||||
|
||||
SPIECE_UNDERLINE = "▁"
|
||||
|
||||
SAMPLE_VOCAB = os.path.join(os.path.dirname(os.path.abspath(__file__)), "fixtures/test_sentencepiece.model")
|
||||
|
||||
|
||||
class BertGenerationTokenizationTest(TokenizerTesterMixin, unittest.TestCase):
|
||||
|
||||
tokenizer_class = BertGenerationTokenizer
|
||||
|
||||
def setUp(self):
|
||||
super().setUp()
|
||||
|
||||
tokenizer = BertGenerationTokenizer(SAMPLE_VOCAB, keep_accents=True)
|
||||
tokenizer.save_pretrained(self.tmpdirname)
|
||||
|
||||
def test_full_tokenizer(self):
|
||||
tokenizer = BertGenerationTokenizer(SAMPLE_VOCAB, keep_accents=True)
|
||||
|
||||
tokens = tokenizer.tokenize("This is a test")
|
||||
self.assertListEqual(tokens, ["▁This", "▁is", "▁a", "▁t", "est"])
|
||||
|
||||
self.assertListEqual(
|
||||
tokenizer.convert_tokens_to_ids(tokens),
|
||||
[285, 46, 10, 170, 382],
|
||||
)
|
||||
|
||||
tokens = tokenizer.tokenize("I was born in 92000, and this is falsé.")
|
||||
self.assertListEqual(
|
||||
tokens,
|
||||
[
|
||||
SPIECE_UNDERLINE + "I",
|
||||
SPIECE_UNDERLINE + "was",
|
||||
SPIECE_UNDERLINE + "b",
|
||||
"or",
|
||||
"n",
|
||||
SPIECE_UNDERLINE + "in",
|
||||
SPIECE_UNDERLINE + "",
|
||||
"9",
|
||||
"2",
|
||||
"0",
|
||||
"0",
|
||||
"0",
|
||||
",",
|
||||
SPIECE_UNDERLINE + "and",
|
||||
SPIECE_UNDERLINE + "this",
|
||||
SPIECE_UNDERLINE + "is",
|
||||
SPIECE_UNDERLINE + "f",
|
||||
"al",
|
||||
"s",
|
||||
"é",
|
||||
".",
|
||||
],
|
||||
)
|
||||
ids = tokenizer.convert_tokens_to_ids(tokens)
|
||||
self.assertListEqual(
|
||||
ids,
|
||||
[8, 21, 84, 55, 24, 19, 7, 0, 602, 347, 347, 347, 3, 12, 66, 46, 72, 80, 6, 0, 4],
|
||||
)
|
||||
|
||||
back_tokens = tokenizer.convert_ids_to_tokens(ids)
|
||||
self.assertListEqual(
|
||||
back_tokens,
|
||||
[
|
||||
SPIECE_UNDERLINE + "I",
|
||||
SPIECE_UNDERLINE + "was",
|
||||
SPIECE_UNDERLINE + "b",
|
||||
"or",
|
||||
"n",
|
||||
SPIECE_UNDERLINE + "in",
|
||||
SPIECE_UNDERLINE + "",
|
||||
"<unk>",
|
||||
"2",
|
||||
"0",
|
||||
"0",
|
||||
"0",
|
||||
",",
|
||||
SPIECE_UNDERLINE + "and",
|
||||
SPIECE_UNDERLINE + "this",
|
||||
SPIECE_UNDERLINE + "is",
|
||||
SPIECE_UNDERLINE + "f",
|
||||
"al",
|
||||
"s",
|
||||
"<unk>",
|
||||
".",
|
||||
],
|
||||
)
|
||||
|
||||
@cached_property
|
||||
def big_tokenizer(self):
|
||||
return BertGenerationTokenizer.from_pretrained("google/bert_for_seq_generation_L-24_bbc_encoder")
|
||||
|
||||
@slow
|
||||
def test_tokenization_base_easy_symbols(self):
|
||||
symbols = "Hello World!"
|
||||
original_tokenizer_encodings = [18536, 2260, 101]
|
||||
|
||||
self.assertListEqual(original_tokenizer_encodings, self.big_tokenizer.encode(symbols))
|
||||
|
||||
@slow
|
||||
def test_tokenization_base_hard_symbols(self):
|
||||
symbols = 'This is a very long text with a lot of weird characters, such as: . , ~ ? ( ) " [ ] ! : - . Also we will add words that should not exsist and be tokenized to <unk>, such as saoneuhaoesuth'
|
||||
original_tokenizer_encodings = [
|
||||
871,
|
||||
419,
|
||||
358,
|
||||
946,
|
||||
991,
|
||||
2521,
|
||||
452,
|
||||
358,
|
||||
1357,
|
||||
387,
|
||||
7751,
|
||||
3536,
|
||||
112,
|
||||
985,
|
||||
456,
|
||||
126,
|
||||
865,
|
||||
938,
|
||||
5400,
|
||||
5734,
|
||||
458,
|
||||
1368,
|
||||
467,
|
||||
786,
|
||||
2462,
|
||||
5246,
|
||||
1159,
|
||||
633,
|
||||
865,
|
||||
4519,
|
||||
457,
|
||||
582,
|
||||
852,
|
||||
2557,
|
||||
427,
|
||||
916,
|
||||
508,
|
||||
405,
|
||||
34324,
|
||||
497,
|
||||
391,
|
||||
408,
|
||||
11342,
|
||||
1244,
|
||||
385,
|
||||
100,
|
||||
938,
|
||||
985,
|
||||
456,
|
||||
574,
|
||||
362,
|
||||
12597,
|
||||
3200,
|
||||
3129,
|
||||
1172,
|
||||
]
|
||||
|
||||
self.assertListEqual(original_tokenizer_encodings, self.big_tokenizer.encode(symbols))
|
||||
|
||||
@slow
|
||||
@require_torch
|
||||
def test_torch_encode_plus_sent_to_model(self):
|
||||
import torch
|
||||
|
||||
from transformers import BertGenerationConfig, BertGenerationEncoder
|
||||
|
||||
# Build sequence
|
||||
first_ten_tokens = list(self.big_tokenizer.get_vocab().keys())[:10]
|
||||
sequence = " ".join(first_ten_tokens)
|
||||
encoded_sequence = self.big_tokenizer.encode_plus(sequence, return_tensors="pt", return_token_type_ids=False)
|
||||
batch_encoded_sequence = self.big_tokenizer.batch_encode_plus(
|
||||
[sequence + " " + sequence], return_tensors="pt", return_token_type_ids=False
|
||||
)
|
||||
|
||||
config = BertGenerationConfig()
|
||||
model = BertGenerationEncoder(config)
|
||||
|
||||
assert model.get_input_embeddings().weight.shape[0] >= self.big_tokenizer.vocab_size
|
||||
|
||||
with torch.no_grad():
|
||||
model(**encoded_sequence)
|
||||
model(**batch_encoded_sequence)
|
||||
@@ -21,11 +21,12 @@ from transformers import BatchEncoding
|
||||
from transformers.file_utils import cached_property
|
||||
from transformers.testing_utils import _torch_available
|
||||
from transformers.tokenization_t5 import T5Tokenizer
|
||||
from transformers.tokenization_xlnet import SPIECE_UNDERLINE
|
||||
|
||||
from .test_tokenization_common import TokenizerTesterMixin
|
||||
|
||||
|
||||
SPIECE_UNDERLINE = "▁"
|
||||
|
||||
SAMPLE_VOCAB = os.path.join(os.path.dirname(os.path.abspath(__file__)), "fixtures/test_sentencepiece.model")
|
||||
|
||||
FRAMEWORK = "pt" if _torch_available else "tf"
|
||||
@@ -138,9 +139,6 @@ class T5TokenizationTest(TokenizerTesterMixin, unittest.TestCase):
|
||||
self.assertEqual((2, 9), batch.input_ids.shape)
|
||||
self.assertEqual((2, 9), batch.attention_mask.shape)
|
||||
|
||||
# Test that special tokens are reset
|
||||
self.assertEqual(tokenizer.prefix_tokens, [])
|
||||
|
||||
def test_empty_target_text(self):
|
||||
tokenizer = self.t5_base_tokenizer
|
||||
src_text = ["A long paragraph for summarization.", "Another paragraph for summarization."]
|
||||
@@ -183,7 +181,7 @@ class T5TokenizationTest(TokenizerTesterMixin, unittest.TestCase):
|
||||
src_text = ["A long paragraph for summarization. </s>"]
|
||||
tgt_text = ["Summary of the text. </s>"]
|
||||
expected_src_tokens = [71, 307, 8986, 21, 4505, 1635, 1707, 5, 1]
|
||||
expected_tgt_tokens = [0, 20698, 13, 8, 1499, 5, 1]
|
||||
expected_tgt_tokens = [20698, 13, 8, 1499, 5, 1]
|
||||
|
||||
batch = tokenizer.prepare_seq2seq_batch(src_text, tgt_texts=tgt_text, return_tensors=FRAMEWORK)
|
||||
|
||||
|
||||
+11
-5
@@ -1,10 +1,10 @@
|
||||
import unittest
|
||||
|
||||
import nlp
|
||||
import datasets
|
||||
import numpy as np
|
||||
|
||||
from transformers import AutoTokenizer, TrainingArguments, is_torch_available
|
||||
from transformers.testing_utils import get_tests_dir, require_torch
|
||||
from transformers.testing_utils import get_tests_dir, require_non_multigpu, require_torch
|
||||
|
||||
|
||||
if is_torch_available():
|
||||
@@ -111,6 +111,7 @@ class TrainerIntegrationTest(unittest.TestCase):
|
||||
self.n_epochs = args.num_train_epochs
|
||||
self.batch_size = args.per_device_train_batch_size
|
||||
|
||||
@require_non_multigpu
|
||||
def test_reproducible_training(self):
|
||||
# Checks that training worked, model trained and seed made a reproducible training.
|
||||
trainer = get_regression_trainer(learning_rate=0.1)
|
||||
@@ -122,6 +123,7 @@ class TrainerIntegrationTest(unittest.TestCase):
|
||||
trainer.train()
|
||||
self.check_trained_model(trainer.model, alternate_seed=True)
|
||||
|
||||
@require_non_multigpu
|
||||
def test_number_of_steps_in_training(self):
|
||||
# Regular training has n_epochs * len(train_dl) steps
|
||||
trainer = get_regression_trainer(learning_rate=0.1)
|
||||
@@ -138,6 +140,7 @@ class TrainerIntegrationTest(unittest.TestCase):
|
||||
train_output = trainer.train()
|
||||
self.assertEqual(train_output.global_step, 10)
|
||||
|
||||
@require_non_multigpu
|
||||
def test_train_and_eval_dataloaders(self):
|
||||
trainer = get_regression_trainer(learning_rate=0.1, per_device_train_batch_size=16)
|
||||
self.assertEqual(trainer.get_train_dataloader().batch_size, 16)
|
||||
@@ -200,11 +203,12 @@ class TrainerIntegrationTest(unittest.TestCase):
|
||||
x = trainer.eval_dataset.x
|
||||
self.assertTrue(np.allclose(preds, 1.5 * x + 2.5))
|
||||
|
||||
def test_trainer_with_nlp(self):
|
||||
@require_non_multigpu
|
||||
def test_trainer_with_datasets(self):
|
||||
np.random.seed(42)
|
||||
x = np.random.normal(size=(64,)).astype(np.float32)
|
||||
y = 2.0 * x + 3.0 + np.random.normal(scale=0.1, size=(64,))
|
||||
train_dataset = nlp.Dataset.from_dict({"input_x": x, "label": y})
|
||||
train_dataset = datasets.Dataset.from_dict({"input_x": x, "label": y})
|
||||
|
||||
# Base training. Should have the same results as test_reproducible_training
|
||||
model = RegressionModel()
|
||||
@@ -222,12 +226,13 @@ class TrainerIntegrationTest(unittest.TestCase):
|
||||
|
||||
# Adding one column not used by the model should have no impact
|
||||
z = np.random.normal(size=(64,)).astype(np.float32)
|
||||
train_dataset = nlp.Dataset.from_dict({"input_x": x, "label": y, "extra": z})
|
||||
train_dataset = datasets.Dataset.from_dict({"input_x": x, "label": y, "extra": z})
|
||||
model = RegressionModel()
|
||||
trainer = Trainer(model, args, train_dataset=train_dataset)
|
||||
trainer.train()
|
||||
self.check_trained_model(trainer.model)
|
||||
|
||||
@require_non_multigpu
|
||||
def test_custom_optimizer(self):
|
||||
train_dataset = RegressionDataset()
|
||||
args = TrainingArguments("./regression")
|
||||
@@ -241,6 +246,7 @@ class TrainerIntegrationTest(unittest.TestCase):
|
||||
self.assertTrue(torch.abs(trainer.model.b - 2.5656) < 1e-4)
|
||||
self.assertEqual(trainer.optimizer.state_dict()["param_groups"][0]["lr"], 1.0)
|
||||
|
||||
@require_non_multigpu
|
||||
def test_model_init(self):
|
||||
train_dataset = RegressionDataset()
|
||||
args = TrainingArguments("./regression", learning_rate=0.1)
|
||||
|
||||
@@ -47,6 +47,7 @@ MODEL_NAME_TO_DOC_FILE = {
|
||||
"openai": "gpt.rst",
|
||||
"transfo_xl": "transformerxl.rst",
|
||||
"xlm_roberta": "xlmroberta.rst",
|
||||
"bert_generation": "bertgeneration.rst",
|
||||
}
|
||||
|
||||
# This is to make sure the transformers module imported is the one in the repo.
|
||||
@@ -230,6 +231,9 @@ def _get_model_name(module):
|
||||
# Secial case for xlm_roberta
|
||||
if splits[-1] == "roberta" and splits[-2] == "xlm":
|
||||
return "_".join(splits[-2:])
|
||||
# Special case for bert_generation
|
||||
if splits[-1] == "generation" and splits[-2] == "bert":
|
||||
return "_".join(splits[-2:])
|
||||
return splits[-1]
|
||||
|
||||
|
||||
|
||||
Reference in new issue
Block a user