Compare commits
102
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a2f830d5a1 | ||
|
|
8f5fd79b8f | ||
|
|
c1be41f452 | ||
|
|
135689bba3 | ||
|
|
64141bab07 | ||
|
|
3cd4a574c2 | ||
|
|
237f27f724 | ||
|
|
4274e9c223 | ||
|
|
47b137e175 | ||
|
|
82afc4b93e | ||
|
|
59ce19cde4 | ||
|
|
3abfc19ae0 | ||
|
|
5b47f0bc3b | ||
|
|
897101cfce | ||
|
|
60c8defa01 | ||
|
|
d7e169b3d9 | ||
|
|
1cf7dcbe71 | ||
|
|
82a22c20a8 | ||
|
|
d8fb4c836b | ||
|
|
6bfa18e1a7 | ||
|
|
aca8e30ddf | ||
|
|
349f85f241 | ||
|
|
3860f3144a | ||
|
|
f69a9d32fa | ||
|
|
c8c5ce0fd3 | ||
|
|
6ab7a4584b | ||
|
|
e210739bef | ||
|
|
8977533c8d | ||
|
|
cf9561a4bb | ||
|
|
f64b6c1dc8 | ||
|
|
b378005edf | ||
|
|
eaf68afffe | ||
|
|
4f546ad160 | ||
|
|
2884cd7bdb | ||
|
|
0f32ad8319 | ||
|
|
2c83e1bfd6 | ||
|
|
15af641996 | ||
|
|
95cd16275c | ||
|
|
044fa94285 | ||
|
|
f780b9f415 | ||
|
|
2f1211bbb5 | ||
|
|
00a1fc9ae4 | ||
|
|
2faaa4ad3c | ||
|
|
c1bc9fe05d | ||
|
|
6e9f30748f | ||
|
|
bc440f3e7c | ||
|
|
2094d37888 | ||
|
|
b3f9e986d9 | ||
|
|
b4a094dd97 | ||
|
|
3688823d19 | ||
|
|
e0a37450e1 | ||
|
|
720ee41342 | ||
|
|
17a8621661 | ||
|
|
64796004dc | ||
|
|
df5eec9f14 | ||
|
|
1497120d89 | ||
|
|
593765088e | ||
|
|
4ad9cb1bd9 | ||
|
|
6215f3b5e0 | ||
|
|
9795dc3464 | ||
|
|
a4c7c25cd2 | ||
|
|
945d56995e | ||
|
|
e82aca09a5 | ||
|
|
845c18d9af | ||
|
|
0f3dc78c0b | ||
|
|
21e8c67bc7 | ||
|
|
972d240ae6 | ||
|
|
2b8ab2eef3 | ||
|
|
706a7c064d | ||
|
|
d37b95d39f | ||
|
|
cbf479df13 | ||
|
|
0b9b2840c3 | ||
|
|
ce3684bbcb | ||
|
|
b8a2ba8624 | ||
|
|
cc4ba034d6 | ||
|
|
84d8061dee | ||
|
|
52bbf8dfe3 | ||
|
|
930bdfaaa8 | ||
|
|
1e1a671614 | ||
|
|
edbfab74d9 | ||
|
|
3fa122af11 | ||
|
|
85929c02dd | ||
|
|
2e4dac6802 | ||
|
|
7184c2b2c2 | ||
|
|
5ed8b6460a | ||
|
|
5774e511b0 | ||
|
|
71d2aa59e0 | ||
|
|
0a6be7b00f | ||
|
|
3cb50b0c31 | ||
|
|
3bf6f0b8f1 | ||
|
|
9ea0a9e825 | ||
|
|
3f30195fd0 | ||
|
|
90bb4188e4 | ||
|
|
6c7d856b86 | ||
|
|
17007b7b66 | ||
|
|
ae3079a9c4 | ||
|
|
97c9900124 | ||
|
|
cbcfc0e948 | ||
|
|
5fd9e5d6ab | ||
|
|
20c8a9aa36 | ||
|
|
f8b1e89f8b | ||
|
|
23eaa7afcf |
@@ -245,7 +245,7 @@ jobs:
|
||||
paths:
|
||||
- '~/.cache/pip'
|
||||
- run: black --check --line-length 119 --target-version py35 examples templates tests src utils
|
||||
- run: isort --check-only examples templates tests src utils
|
||||
- run: isort --check-only --recursive examples templates tests src utils
|
||||
- run: flake8 examples templates tests src utils
|
||||
- run: python utils/check_repo.py
|
||||
check_repository_consistency:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -22,125 +22,59 @@
|
||||
<p>State-of-the-art Natural Language Processing for PyTorch and TensorFlow 2.0
|
||||
</h3>
|
||||
|
||||
🤗 Transformers provides thousands of pretrained models to perform tasks on texts such as classification, information extraction, question answering, summarization, translation, text generation, etc in 100+ languages. Its aim is to make cutting-edge NLP easier to use for everyone.
|
||||
|
||||
🤗 Transformers provides APIs to quickly download and use those pretrained models on a given text, fine-tune them on your own datasets then share them with the community on our [model hub](https://huggingface.co/models). At the same time, each python module defining an architecture can be used as a standalone and modified to enable quick research experiments.
|
||||
|
||||
🤗 Transformers is backed by the two most popular deep learning libraries, [PyTorch](https://pytorch.org/) and [TensorFlow](https://www.tensorflow.org/), with a seamless integration between them, allowing you to train your models with one then load it for inference with the other.
|
||||
🤗 Transformers (formerly known as `pytorch-transformers` and `pytorch-pretrained-bert`) provides state-of-the-art general-purpose architectures (BERT, GPT-2, RoBERTa, XLM, DistilBert, XLNet, T5, CTRL...) for Natural Language Understanding (NLU) and Natural Language Generation (NLG) with over thousands of pretrained models in 100+ languages and deep interoperability between PyTorch & TensorFlow 2.0.
|
||||
|
||||
### Recent contributors
|
||||
[](https://sourcerer.io/fame/clmnt/huggingface/transformers/links/0)[](https://sourcerer.io/fame/clmnt/huggingface/transformers/links/1)[](https://sourcerer.io/fame/clmnt/huggingface/transformers/links/2)[](https://sourcerer.io/fame/clmnt/huggingface/transformers/links/3)[](https://sourcerer.io/fame/clmnt/huggingface/transformers/links/4)[](https://sourcerer.io/fame/clmnt/huggingface/transformers/links/5)[](https://sourcerer.io/fame/clmnt/huggingface/transformers/links/6)[](https://sourcerer.io/fame/clmnt/huggingface/transformers/links/7)
|
||||
|
||||
## Online demos
|
||||
### Features
|
||||
- High performance on NLU and NLG tasks
|
||||
- Low barrier to entry for educators and practitioners
|
||||
|
||||
You can test most of our models directly on their pages from the [model hub](https://huggingface.co/models). We also offer an [inference API](https://huggingface.co/pricing) to use those models.
|
||||
State-of-the-art NLP for everyone
|
||||
- Deep learning researchers
|
||||
- Hands-on practitioners
|
||||
- AI/ML/NLP teachers and educators
|
||||
|
||||
Here are a few examples:
|
||||
- [Masked word completion with BERT](https://huggingface.co/bert-base-uncased?text=Paris+is+the+%5BMASK%5D+of+France)
|
||||
- [Name Entity Recognition with Electra](https://huggingface.co/dbmdz/electra-large-discriminator-finetuned-conll03-english?text=My+name+is+Sarah+and+I+live+in+London+city)
|
||||
- [Text generation with GPT-2](https://huggingface.co/gpt2?text=A+long+time+ago%2C+)
|
||||
- [Natural Langugage Inference with RoBERTa](https://huggingface.co/roberta-large-mnli?text=The+dog+was+lost.+Nobody+lost+any+animal)
|
||||
- [Summarization with BART](https://huggingface.co/facebook/bart-large-cnn?text=The+tower+is+324+metres+%281%2C063+ft%29+tall%2C+about+the+same+height+as+an+81-storey+building%2C+and+the+tallest+structure+in+Paris.+Its+base+is+square%2C+measuring+125+metres+%28410+ft%29+on+each+side.+During+its+construction%2C+the+Eiffel+Tower+surpassed+the+Washington+Monument+to+become+the+tallest+man-made+structure+in+the+world%2C+a+title+it+held+for+41+years+until+the+Chrysler+Building+in+New+York+City+was+finished+in+1930.+It+was+the+first+structure+to+reach+a+height+of+300+metres.+Due+to+the+addition+of+a+broadcasting+aerial+at+the+top+of+the+tower+in+1957%2C+it+is+now+taller+than+the+Chrysler+Building+by+5.2+metres+%2817+ft%29.+Excluding+transmitters%2C+the+Eiffel+Tower+is+the+second+tallest+free-standing+structure+in+France+after+the+Millau+Viaduct)
|
||||
- [Question answering with DistilBERT](https://huggingface.co/distilbert-base-uncased-distilled-squad?text=Which+name+is+also+used+to+describe+the+Amazon+rainforest+in+English%3F&context=The+Amazon+rainforest+%28Portuguese%3A+Floresta+Amaz%C3%B4nica+or+Amaz%C3%B4nia%3B+Spanish%3A+Selva+Amaz%C3%B3nica%2C+Amazon%C3%ADa+or+usually+Amazonia%3B+French%3A+For%C3%AAt+amazonienne%3B+Dutch%3A+Amazoneregenwoud%29%2C+also+known+in+English+as+Amazonia+or+the+Amazon+Jungle%2C+is+a+moist+broadleaf+forest+that+covers+most+of+the+Amazon+basin+of+South+America.+This+basin+encompasses+7%2C000%2C000+square+kilometres+%282%2C700%2C000+sq+mi%29%2C+of+which+5%2C500%2C000+square+kilometres+%282%2C100%2C000+sq+mi%29+are+covered+by+the+rainforest.+This+region+includes+territory+belonging+to+nine+nations.+The+majority+of+the+forest+is+contained+within+Brazil%2C+with+60%25+of+the+rainforest%2C+followed+by+Peru+with+13%25%2C+Colombia+with+10%25%2C+and+with+minor+amounts+in+Venezuela%2C+Ecuador%2C+Bolivia%2C+Guyana%2C+Suriname+and+French+Guiana.+States+or+departments+in+four+nations+contain+%22Amazonas%22+in+their+names.+The+Amazon+represents+over+half+of+the+planet%27s+remaining+rainforests%2C+and+comprises+the+largest+and+most+biodiverse+tract+of+tropical+rainforest+in+the+world%2C+with+an+estimated+390+billion+individual+trees+divided+into+16%2C000+species)
|
||||
- [Translation with T5](https://huggingface.co/t5-base?text=My+name+is+Wolfgang+and+I+live+in+Berlin)
|
||||
Lower compute costs, smaller carbon footprint
|
||||
- Researchers can share trained models instead of always retraining
|
||||
- Practitioners can reduce compute time and production costs
|
||||
- Dozens of architectures with over 1,000 pretrained models, some in more than 100 languages
|
||||
|
||||
**[Write With Transformer](https://transformer.huggingface.co)**, built by the Hugging Face team, is the official demo of this repo’s text generation capabilities.
|
||||
Choose the right framework for every part of a model's lifetime
|
||||
- Train state-of-the-art models in 3 lines of code
|
||||
- Deep interoperability between TensorFlow 2.0 and PyTorch models
|
||||
- Move a single model between TF2.0/PyTorch frameworks at will
|
||||
- Seamlessly pick the right framework for training, evaluation, production
|
||||
|
||||
## Quick tour
|
||||
|
||||
To immediately use a model on a given text, we provide the `pipeline` API. Pipelines group together a pretrained model with the preprocessing that was used during that model training. Here is how to quickly use a pipeline to classify postivive versus negative texts
|
||||
|
||||
```python
|
||||
>>> from transformers import pipeline
|
||||
|
||||
# Allocate a pipeline for sentiment-analysis
|
||||
>>> classifier = pipeline('sentiment-analysis')
|
||||
>>> classifier('We are very happy to include pipeline into the transformers repository.')
|
||||
[{'label': 'POSITIVE', 'score': 0.9978193640708923}]
|
||||
```
|
||||
|
||||
The second line of code downloads and caches the pretrained model used by the pipeline, the third line evaluates it on the given text. Here the answer is "positive" with a confidence of 99.8%.
|
||||
|
||||
This is another example of pipeline used for that can extract question answers from some context:
|
||||
|
||||
``` python
|
||||
>>> from transformers import pipeline
|
||||
|
||||
# Allocate a pipeline for question-answering
|
||||
>>> question_answerer = pipeline('question-answering')
|
||||
>>> question_answerer({
|
||||
... 'question': 'What is the name of the repository ?',
|
||||
... 'context': 'Pipeline have been included in the huggingface/transformers repository'
|
||||
... })
|
||||
{'score': 0.5135612454720828, 'start': 35, 'end': 59, 'answer': 'huggingface/transformers'}
|
||||
|
||||
```
|
||||
|
||||
On top of the answer, the pretrained model used here returned its confidence score, along with the start position and its end position in the tokenized sentence. You can learn more about the tasks supported by the `pipeline` API in [this tutorial](https://huggingface.co/transformers/task_summary.html).
|
||||
|
||||
To download and use any of the pretrained models on your given task, you just need to use those three lines of codes (PyTorch verison):
|
||||
```python
|
||||
>>> from transformers import AutoTokenizer, AutoModel
|
||||
|
||||
>>> tokenizer = AutoTokenizer.from_pretrained("bert-base-uncased")
|
||||
>>> model = AutoModel.from_pretrained("bert_base_uncased")
|
||||
|
||||
>>> inputs = tokenizer("Hello world!", return_tensors="pt")
|
||||
>>> outputs = model(**inputs)
|
||||
```
|
||||
or for TensorFlow:
|
||||
```python
|
||||
>>> from transformers import AutoTokenizer, TFAutoModel
|
||||
|
||||
>>> tokenizer = AutoTokenizer.from_pretrained("bert-base-uncased")
|
||||
>>> model = TFAutoModel.from_pretrained("bert_base_uncased")
|
||||
|
||||
>>> inputs = tokenizer("Hello world!", return_tensors="tf")
|
||||
>>> outputs = model(**inputs)
|
||||
```
|
||||
|
||||
The tokenizer is responsible for all the preprocessing the pretrained model expects, and can be called directly on one (or list) of texts (as we can see on the fourth line of both code examples). It will output a dictionary you can directly pass to your model (which is done on the fifth line).
|
||||
|
||||
The model itself is a regular [Pytorch `nn.Module`](https://pytorch.org/docs/stable/nn.html#torch.nn.Module) or a [TensorFlow `tf.keras.Model`](https://www.tensorflow.org/api_docs/python/tf/keras/Model) (depending on your backend) which you can use normally. For instance, [this tutorial](https://huggingface.co/transformers/training.html) explains how to integrate such a model in classic PyTorch or TensorFlow training loop, or how to use our `Trainer` API to quickly fine-tune the on a new dataset.
|
||||
|
||||
## Why should I use transformers?
|
||||
|
||||
1. Easy-to-use state-of-the-art models:
|
||||
- High performance on NLU and NLG tasks.
|
||||
- Low barrier to entry for educators and practitioners.
|
||||
- Few user-facing abastractions with just three classes to learn.
|
||||
- A unified API for using all our pretrained models.
|
||||
|
||||
1. Lower compute costs, smaller carbon footprint:
|
||||
- Researchers can share trained models instead of always retraining.
|
||||
- Practitioners can reduce compute time and production costs.
|
||||
- Dozens of architectures with over 2,000 pretrained models, some in more than 100 languages.
|
||||
|
||||
1. Choose the right framework for every part of a model's lifetime:
|
||||
- Train state-of-the-art models in 3 lines of code.
|
||||
- Move a single model between TF2.0/PyTorch frameworks at will.
|
||||
- Seamlessly pick the right framework for training, evaluation, production.
|
||||
|
||||
1. Easily customize a model or an example to your needs:
|
||||
- Examples for each architecture to reproduce the results by the official authors of said architecture.
|
||||
- Expose the models internal as consistently as possible.
|
||||
- Model files can be used independently of the library for quick experiments.
|
||||
|
||||
## Why shouldn't I use transformers?
|
||||
|
||||
- This library is not a modular toolbox of building blocks for neural nets. The code in the model files is not refactored with additional abstractions on purpose, so that researchers can quickly iterate on each of the models without diving in additional abstractions/files.
|
||||
- The training API is not intended to work on any model but is optimized to work with the models provided by the library. For generic machine learning loops, you should use another library.
|
||||
- While we strive to present as many use cases as possible, the scripts in our [examples folder](https://github.com/huggingface/transformers/tree/master/examples) are just that: examples. It is expected that they won't work out-of-the box on your specific problem and that you will be required to change a few lines of code to adapt them to your needs.
|
||||
| Section | Description |
|
||||
|-|-|
|
||||
| [Installation](#installation) | How to install the package |
|
||||
| [Model architectures](#model-architectures) | Architectures (with pretrained weights) |
|
||||
| [Online demo](#online-demo) | Experimenting with this repo’s text generation capabilities |
|
||||
| [Quick tour: Usage](#quick-tour) | Tokenizers & models usage: Bert and GPT-2 |
|
||||
| [Quick tour: TF 2.0 and PyTorch ](#Quick-tour-TF-20-training-and-PyTorch-interoperability) | Train a TF 2.0 model in 10 lines of code, load it in PyTorch |
|
||||
| [Quick tour: pipelines](#quick-tour-of-pipelines) | Using Pipelines: Wrapper around tokenizer and models to use finetuned models |
|
||||
| [Quick tour: Fine-tuning/usage scripts](#quick-tour-of-the-fine-tuningusage-scripts) | Using provided scripts: GLUE, SQuAD and Text generation |
|
||||
| [Quick tour: Share your models ](#Quick-tour-of-model-sharing) | Upload and share your fine-tuned models with the community |
|
||||
| [Migrating from pytorch-transformers to transformers](#Migrating-from-pytorch-transformers-to-transformers) | Migrating your code from pytorch-transformers to transformers |
|
||||
| [Migrating from pytorch-pretrained-bert to pytorch-transformers](#Migrating-from-pytorch-pretrained-bert-to-transformers) | Migrating your code from pytorch-pretrained-bert to transformers |
|
||||
| [Documentation](https://huggingface.co/transformers/) | Full API documentation and more |
|
||||
|
||||
## Installation
|
||||
|
||||
This repository is tested on Python 3.6+, PyTorch 1.0.0+ (PyTorch 1.3.1+ for [examples](https://github.com/huggingface/transformers/tree/master/examples)) and TensorFlow 2.0.
|
||||
This repo is tested on Python 3.6+, PyTorch 1.0.0+ (PyTorch 1.3.1+ for examples) and TensorFlow 2.0.
|
||||
|
||||
You should install 🤗 Transformers in a [virtual environment](https://docs.python.org/3/library/venv.html). If you're unfamiliar with Python virtual environments, check out the [user guide](https://packaging.python.org/guides/installing-using-pip-and-virtual-environments/).
|
||||
|
||||
First, create a virtual environment with the version of Python you're going to use and activate it.
|
||||
Create a virtual environment with the version of Python you're going to use and activate it.
|
||||
|
||||
Then, you will need to install one of, or both, TensorFlow 2.0 and PyTorch.
|
||||
Now, if you want to use 🤗 Transformers, you can install it with pip. If you'd like to play with the examples, you must install it from source.
|
||||
|
||||
### With pip
|
||||
|
||||
First you need to install one of, or both, TensorFlow 2.0 and PyTorch.
|
||||
Please refer to [TensorFlow installation page](https://www.tensorflow.org/install/pip#tensorflow-2.0-rc-is-available) and/or [PyTorch installation page](https://pytorch.org/get-started/locally/#start-locally) regarding the specific install command for your platform.
|
||||
|
||||
When TensorFlow 2.0 and/or PyTorch has been installed, 🤗 Transformers can be installed using pip as follows:
|
||||
@@ -149,11 +83,68 @@ When TensorFlow 2.0 and/or PyTorch has been installed, 🤗 Transformers can be
|
||||
pip install transformers
|
||||
```
|
||||
|
||||
If you'd like to play with the examples, you must [install the library from source](https://huggingface.co/transformers/installation.html#installing-from-source).
|
||||
### From source
|
||||
|
||||
## Models architectures
|
||||
Here also, you first need to install one of, or both, TensorFlow 2.0 and PyTorch.
|
||||
Please refer to [TensorFlow installation page](https://www.tensorflow.org/install/pip#tensorflow-2.0-rc-is-available) and/or [PyTorch installation page](https://pytorch.org/get-started/locally/#start-locally) regarding the specific install command for your platform.
|
||||
|
||||
🤗 Transformers currently provides the following architectures (see [here](https://huggingface.co/transformers/model_summary.html) for a high-level summary of each them):
|
||||
When TensorFlow 2.0 and/or PyTorch has been installed, you can install from source by cloning the repository and running:
|
||||
|
||||
```bash
|
||||
git clone https://github.com/huggingface/transformers
|
||||
cd transformers
|
||||
pip install .
|
||||
```
|
||||
|
||||
When you update the repository, you should upgrade the transformers installation and its dependencies as follows:
|
||||
|
||||
```bash
|
||||
git pull
|
||||
pip install --upgrade .
|
||||
```
|
||||
|
||||
### Run the examples
|
||||
|
||||
Examples are included in the repository but are not shipped with the library.
|
||||
|
||||
Therefore, in order to run the latest versions of the examples, you need to install from source, as described above.
|
||||
|
||||
Look at the [README](https://github.com/huggingface/transformers/blob/master/examples/README.md) for how to run examples.
|
||||
|
||||
### Tests
|
||||
|
||||
A series of tests are included for the library and for some example scripts. Library tests can be found in the [tests folder](https://github.com/huggingface/transformers/tree/master/tests) and examples tests in the [examples folder](https://github.com/huggingface/transformers/tree/master/examples).
|
||||
|
||||
Depending on which framework is installed (TensorFlow 2.0 and/or PyTorch), the irrelevant tests will be skipped. Ensure that both frameworks are installed if you want to execute all tests.
|
||||
|
||||
Here's the easiest way to run tests for the library:
|
||||
|
||||
```bash
|
||||
pip install -e ".[testing]"
|
||||
make test
|
||||
```
|
||||
|
||||
and for the examples:
|
||||
|
||||
```bash
|
||||
pip install -e ".[testing]"
|
||||
pip install -r examples/requirements.txt
|
||||
make test-examples
|
||||
```
|
||||
|
||||
For details, refer to the [contributing guide](https://github.com/huggingface/transformers/blob/master/CONTRIBUTING.md#tests).
|
||||
|
||||
### Do you want to run a Transformer model on a mobile device?
|
||||
|
||||
You should check out our [`swift-coreml-transformers`](https://github.com/huggingface/swift-coreml-transformers) repo.
|
||||
|
||||
It contains a set of tools to convert PyTorch or TensorFlow 2.0 trained Transformer models (currently contains `GPT-2`, `DistilGPT-2`, `BERT`, and `DistilBERT`) to CoreML models that run on iOS devices.
|
||||
|
||||
At some point in the future, you'll be able to seamlessly move from pre-training or fine-tuning models to productizing them in CoreML, or prototype a model or an app in CoreML then research its hyperparameters or architecture from TensorFlow 2.0 and/or PyTorch. Super exciting!
|
||||
|
||||
## Model architectures
|
||||
|
||||
🤗 Transformers currently provides the following NLU/NLG architectures:
|
||||
|
||||
1. **[BERT](https://huggingface.co/transformers/model_doc/bert.html)** (from Google) released with the paper [BERT: Pre-training of Deep Bidirectional Transformers for Language Understanding](https://arxiv.org/abs/1810.04805) by Jacob Devlin, Ming-Wei Chang, Kenton Lee and Kristina Toutanova.
|
||||
2. **[GPT](https://huggingface.co/transformers/model_doc/gpt.html)** (from OpenAI) released with the paper [Improving Language Understanding by Generative Pre-Training](https://blog.openai.com/language-unsupervised/) by Alec Radford, Karthik Narasimhan, Tim Salimans and Ilya Sutskever.
|
||||
@@ -186,20 +177,526 @@ Min, Patrick Lewis, Ledell Wu, Sergey Edunov, Danqi Chen, and Wen-tau Yih.
|
||||
27. **[Other community models](https://huggingface.co/models)**, contributed by the [community](https://huggingface.co/users).
|
||||
28. Want to contribute a new model? We have added a **detailed guide and templates** to guide you in the process of adding a new model. You can find them in the [`templates`](./templates) folder of the repository. Be sure to check the [contributing guidelines](./CONTRIBUTING.md) and contact the maintainers or open an issue to collect feedbacks before starting your PR.
|
||||
|
||||
These implementations have been tested on several datasets (see the example scripts) and should match the performances of the original implementations. You can find more details on the performances in the Examples section of the [documentation](https://huggingface.co/transformers/examples.html).
|
||||
These implementations have been tested on several datasets (see the example scripts) and should match the performances of the original implementations (e.g. ~93 F1 on SQuAD for BERT Whole-Word-Masking, ~88 F1 on RocStories for OpenAI GPT, ~18.3 perplexity on WikiText 103 for Transformer-XL, ~0.916 Pearson R coefficient on STS-B for XLNet). You can find more details on the performances in the Examples section of the [documentation](https://huggingface.co/transformers/examples.html).
|
||||
|
||||
## Online demo
|
||||
|
||||
You can test our inference API on most model pages from the model hub: https://huggingface.co/models
|
||||
|
||||
For example:
|
||||
- [Masked word completion with BERT](https://huggingface.co/bert-base-uncased?text=Paris+is+the+%5BMASK%5D+of+France)
|
||||
- [NER with Electra](https://huggingface.co/dbmdz/electra-large-discriminator-finetuned-conll03-english?text=My+name+is+Sarah+and+I+live+in+London+city)
|
||||
- [Text generation with GPT-2](https://huggingface.co/gpt2?text=A+long+time+ago%2C+)
|
||||
- [NLI with RoBERTa](https://huggingface.co/roberta-large-mnli?text=The+dog+was+lost.+Nobody+lost+any+animal)
|
||||
- [Summarization with BART](https://huggingface.co/facebook/bart-large-cnn?text=The+tower+is+324+metres+%281%2C063+ft%29+tall%2C+about+the+same+height+as+an+81-storey+building%2C+and+the+tallest+structure+in+Paris.+Its+base+is+square%2C+measuring+125+metres+%28410+ft%29+on+each+side.+During+its+construction%2C+the+Eiffel+Tower+surpassed+the+Washington+Monument+to+become+the+tallest+man-made+structure+in+the+world%2C+a+title+it+held+for+41+years+until+the+Chrysler+Building+in+New+York+City+was+finished+in+1930.+It+was+the+first+structure+to+reach+a+height+of+300+metres.+Due+to+the+addition+of+a+broadcasting+aerial+at+the+top+of+the+tower+in+1957%2C+it+is+now+taller+than+the+Chrysler+Building+by+5.2+metres+%2817+ft%29.+Excluding+transmitters%2C+the+Eiffel+Tower+is+the+second+tallest+free-standing+structure+in+France+after+the+Millau+Viaduct)
|
||||
- [Question answering with DistilBERT](https://huggingface.co/distilbert-base-uncased-distilled-squad?text=Which+name+is+also+used+to+describe+the+Amazon+rainforest+in+English%3F&context=The+Amazon+rainforest+%28Portuguese%3A+Floresta+Amaz%C3%B4nica+or+Amaz%C3%B4nia%3B+Spanish%3A+Selva+Amaz%C3%B3nica%2C+Amazon%C3%ADa+or+usually+Amazonia%3B+French%3A+For%C3%AAt+amazonienne%3B+Dutch%3A+Amazoneregenwoud%29%2C+also+known+in+English+as+Amazonia+or+the+Amazon+Jungle%2C+is+a+moist+broadleaf+forest+that+covers+most+of+the+Amazon+basin+of+South+America.+This+basin+encompasses+7%2C000%2C000+square+kilometres+%282%2C700%2C000+sq+mi%29%2C+of+which+5%2C500%2C000+square+kilometres+%282%2C100%2C000+sq+mi%29+are+covered+by+the+rainforest.+This+region+includes+territory+belonging+to+nine+nations.+The+majority+of+the+forest+is+contained+within+Brazil%2C+with+60%25+of+the+rainforest%2C+followed+by+Peru+with+13%25%2C+Colombia+with+10%25%2C+and+with+minor+amounts+in+Venezuela%2C+Ecuador%2C+Bolivia%2C+Guyana%2C+Suriname+and+French+Guiana.+States+or+departments+in+four+nations+contain+%22Amazonas%22+in+their+names.+The+Amazon+represents+over+half+of+the+planet%27s+remaining+rainforests%2C+and+comprises+the+largest+and+most+biodiverse+tract+of+tropical+rainforest+in+the+world%2C+with+an+estimated+390+billion+individual+trees+divided+into+16%2C000+species)
|
||||
- [Translation with T5](https://huggingface.co/t5-base?text=My+name+is+Wolfgang+and+I+live+in+Berlin)
|
||||
|
||||
|
||||
## Learn more
|
||||
**[Write With Transformer](https://transformer.huggingface.co)**, built by the Hugging Face team at transformer.huggingface.co, is the official demo of this repo’s text generation capabilities.
|
||||
|
||||
| Section | Description |
|
||||
|-|-|
|
||||
| [Documentation](https://huggingface.co/transformers/) | Full API documentation and tutorials |
|
||||
| [Task summary](https://huggingface.co/transformers/task_summary.html) | Tasks supported by 🤗 Transformers |
|
||||
| [Preprocessing tutorial](https://huggingface.co/transformers/preprocessing.html) | Using the `Tokenizer` class to prepare data for the models |
|
||||
| [Training and fine-tuning](https://huggingface.co/transformers/training.html) | Using the models provided by 🤗 Transformers in a PyTorch/TensorFlow training loop and the `Trainer` API |
|
||||
| [Quick tour: Fine-tuning/usage scripts](https://github.com/huggingface/transformers/tree/master/examples) | Example scripts for fine-tuning models on a wide range of tasks |
|
||||
| [Model sharing and uploading](https://huggingface.co/transformers/model_sharing.html) | Upload and share your fine-tuned models with the community |
|
||||
| [Migration](https://huggingface.co/transformers/migration.html) | Migrate to 🤗 Transformers from `pytorch-transformers` or `pytorch-pretrained-bert` |
|
||||
## Quick tour
|
||||
|
||||
Let's do a very quick overview of the model architectures in 🤗 Transformers. Detailed examples for each model architecture (Bert, GPT, GPT-2, Transformer-XL, XLNet and XLM) can be found in the [full documentation](https://huggingface.co/transformers/).
|
||||
|
||||
```python
|
||||
import torch
|
||||
from transformers import *
|
||||
|
||||
# Transformers has a unified API
|
||||
# for 10 transformer architectures and 30 pretrained weights.
|
||||
# Model | Tokenizer | Pretrained weights shortcut
|
||||
MODELS = [(BertModel, BertTokenizer, 'bert-base-uncased'),
|
||||
(OpenAIGPTModel, OpenAIGPTTokenizer, 'openai-gpt'),
|
||||
(GPT2Model, GPT2Tokenizer, 'gpt2'),
|
||||
(CTRLModel, CTRLTokenizer, 'ctrl'),
|
||||
(TransfoXLModel, TransfoXLTokenizer, 'transfo-xl-wt103'),
|
||||
(XLNetModel, XLNetTokenizer, 'xlnet-base-cased'),
|
||||
(XLMModel, XLMTokenizer, 'xlm-mlm-enfr-1024'),
|
||||
(DistilBertModel, DistilBertTokenizer, 'distilbert-base-cased'),
|
||||
(RobertaModel, RobertaTokenizer, 'roberta-base'),
|
||||
(XLMRobertaModel, XLMRobertaTokenizer, 'xlm-roberta-base'),
|
||||
]
|
||||
|
||||
# To use TensorFlow 2.0 versions of the models, simply prefix the class names with 'TF', e.g. `TFRobertaModel` is the TF 2.0 counterpart of the PyTorch model `RobertaModel`
|
||||
|
||||
# Let's encode some text in a sequence of hidden-states using each model:
|
||||
for model_class, tokenizer_class, pretrained_weights in MODELS:
|
||||
# Load pretrained model/tokenizer
|
||||
tokenizer = tokenizer_class.from_pretrained(pretrained_weights)
|
||||
model = model_class.from_pretrained(pretrained_weights)
|
||||
|
||||
# Encode text
|
||||
input_ids = torch.tensor([tokenizer.encode("Here is some text to encode", add_special_tokens=True)]) # Add special tokens takes care of adding [CLS], [SEP], <s>... tokens in the right way for each model.
|
||||
with torch.no_grad():
|
||||
last_hidden_states = model(input_ids)[0] # Models outputs are now tuples
|
||||
|
||||
# Each architecture is provided with several class for fine-tuning on down-stream tasks, e.g.
|
||||
BERT_MODEL_CLASSES = [BertModel, BertForPreTraining, BertForMaskedLM, BertForNextSentencePrediction,
|
||||
BertForSequenceClassification, BertForTokenClassification, BertForQuestionAnswering]
|
||||
|
||||
# All the classes for an architecture can be initiated from pretrained weights for this architecture
|
||||
# Note that additional weights added for fine-tuning are only initialized
|
||||
# and need to be trained on the down-stream task
|
||||
pretrained_weights = 'bert-base-uncased'
|
||||
tokenizer = BertTokenizer.from_pretrained(pretrained_weights)
|
||||
for model_class in BERT_MODEL_CLASSES:
|
||||
# Load pretrained model/tokenizer
|
||||
model = model_class.from_pretrained(pretrained_weights)
|
||||
|
||||
# Models can return full list of hidden-states & attentions weights at each layer
|
||||
model = model_class.from_pretrained(pretrained_weights,
|
||||
output_hidden_states=True,
|
||||
output_attentions=True)
|
||||
input_ids = torch.tensor([tokenizer.encode("Let's see all hidden-states and attentions on this text")])
|
||||
all_hidden_states, all_attentions = model(input_ids)[-2:]
|
||||
|
||||
# Models are compatible with Torchscript
|
||||
model = model_class.from_pretrained(pretrained_weights, torchscript=True)
|
||||
traced_model = torch.jit.trace(model, (input_ids,))
|
||||
|
||||
# Simple serialization for models and tokenizers
|
||||
model.save_pretrained('./directory/to/save/') # save
|
||||
model = model_class.from_pretrained('./directory/to/save/') # re-load
|
||||
tokenizer.save_pretrained('./directory/to/save/') # save
|
||||
tokenizer = BertTokenizer.from_pretrained('./directory/to/save/') # re-load
|
||||
|
||||
# SOTA examples for GLUE, SQUAD, text generation...
|
||||
```
|
||||
|
||||
## Quick tour TF 2.0 training and PyTorch interoperability
|
||||
|
||||
Let's do a quick example of how a TensorFlow 2.0 model can be trained in 12 lines of code with 🤗 Transformers and then loaded in PyTorch for fast inspection/tests.
|
||||
|
||||
```python
|
||||
import tensorflow as tf
|
||||
import tensorflow_datasets
|
||||
from transformers import *
|
||||
|
||||
# Load dataset, tokenizer, model from pretrained model/vocabulary
|
||||
tokenizer = BertTokenizer.from_pretrained('bert-base-cased')
|
||||
model = TFBertForSequenceClassification.from_pretrained('bert-base-cased')
|
||||
data = tensorflow_datasets.load('glue/mrpc')
|
||||
|
||||
# Prepare dataset for GLUE as a tf.data.Dataset instance
|
||||
train_dataset = glue_convert_examples_to_features(data['train'], tokenizer, max_length=128, task='mrpc')
|
||||
valid_dataset = glue_convert_examples_to_features(data['validation'], tokenizer, max_length=128, task='mrpc')
|
||||
train_dataset = train_dataset.shuffle(100).batch(32).repeat(2)
|
||||
valid_dataset = valid_dataset.batch(64)
|
||||
|
||||
# Prepare training: Compile tf.keras model with optimizer, loss and learning rate schedule
|
||||
optimizer = tf.keras.optimizers.Adam(learning_rate=3e-5, epsilon=1e-08, clipnorm=1.0)
|
||||
loss = tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True)
|
||||
metric = tf.keras.metrics.SparseCategoricalAccuracy('accuracy')
|
||||
model.compile(optimizer=optimizer, loss=loss, metrics=[metric])
|
||||
|
||||
# Train and evaluate using tf.keras.Model.fit()
|
||||
history = model.fit(train_dataset, epochs=2, steps_per_epoch=115,
|
||||
validation_data=valid_dataset, validation_steps=7)
|
||||
|
||||
# Load the TensorFlow model in PyTorch for inspection
|
||||
model.save_pretrained('./save/')
|
||||
pytorch_model = BertForSequenceClassification.from_pretrained('./save/', from_tf=True)
|
||||
|
||||
# Quickly test a few predictions - MRPC is a paraphrasing task, let's see if our model learned the task
|
||||
sentence_0 = "This research was consistent with his findings."
|
||||
sentence_1 = "His findings were compatible with this research."
|
||||
sentence_2 = "His findings were not compatible with this research."
|
||||
inputs_1 = tokenizer(sentence_0, sentence_1, add_special_tokens=True, return_tensors='pt')
|
||||
inputs_2 = tokenizer(sentence_0, sentence_2, add_special_tokens=True, return_tensors='pt')
|
||||
|
||||
pred_1 = pytorch_model(inputs_1['input_ids'], token_type_ids=inputs_1['token_type_ids'])[0].argmax().item()
|
||||
pred_2 = pytorch_model(inputs_2['input_ids'], token_type_ids=inputs_2['token_type_ids'])[0].argmax().item()
|
||||
|
||||
print("sentence_1 is", "a paraphrase" if pred_1 else "not a paraphrase", "of sentence_0")
|
||||
print("sentence_2 is", "a paraphrase" if pred_2 else "not a paraphrase", "of sentence_0")
|
||||
```
|
||||
|
||||
## Quick tour of the fine-tuning/usage scripts
|
||||
|
||||
**Important**
|
||||
Before running the fine-tuning scripts, please read the
|
||||
[instructions](#run-the-examples) on how to
|
||||
setup your environment to run the examples.
|
||||
|
||||
The library comprises several example scripts with SOTA performances for NLU and NLG tasks:
|
||||
|
||||
- `run_glue.py`: an example fine-tuning sequence classification models on nine different GLUE tasks (*sequence-level classification*)
|
||||
- `run_squad.py`: an example fine-tuning question answering models on the question answering dataset SQuAD 2.0 (*token-level classification*)
|
||||
- `run_ner.py`: an example fine-tuning token classification models on named entity recognition (*token-level classification*)
|
||||
- `run_generation.py`: an example using GPT, GPT-2, CTRL, Transformer-XL and XLNet for conditional language generation
|
||||
- other model-specific examples (see the documentation).
|
||||
|
||||
Here are three quick usage examples for these scripts:
|
||||
|
||||
### `run_glue.py`: Fine-tuning on GLUE tasks for sequence classification
|
||||
|
||||
The [General Language Understanding Evaluation (GLUE) benchmark](https://gluebenchmark.com/) is a collection of nine sentence- or sentence-pair language understanding tasks for evaluating and analyzing natural language understanding systems.
|
||||
|
||||
Before running any of these GLUE tasks you should download the
|
||||
[GLUE data](https://gluebenchmark.com/tasks) by running
|
||||
[this script](https://gist.github.com/W4ngatang/60c2bdb54d156a41194446737ce03e2e)
|
||||
and unpack it to some directory `$GLUE_DIR`.
|
||||
|
||||
You should also install the additional packages required by the examples:
|
||||
|
||||
```shell
|
||||
pip install -r ./examples/requirements.txt
|
||||
```
|
||||
|
||||
```shell
|
||||
export GLUE_DIR=/path/to/glue
|
||||
export TASK_NAME=MRPC
|
||||
|
||||
python ./examples/text-classification/run_glue.py \
|
||||
--model_name_or_path bert-base-uncased \
|
||||
--task_name $TASK_NAME \
|
||||
--do_train \
|
||||
--do_eval \
|
||||
--data_dir $GLUE_DIR/$TASK_NAME \
|
||||
--max_seq_length 128 \
|
||||
--per_device_eval_batch_size=8 \
|
||||
--per_device_train_batch_size=8 \
|
||||
--learning_rate 2e-5 \
|
||||
--num_train_epochs 3.0 \
|
||||
--output_dir /tmp/$TASK_NAME/
|
||||
```
|
||||
|
||||
where task name can be one of CoLA, SST-2, MRPC, STS-B, QQP, MNLI, QNLI, RTE, WNLI.
|
||||
|
||||
The dev set results will be present within the text file 'eval_results.txt' in the specified output_dir. In case of MNLI, since there are two separate dev sets, matched and mismatched, there will be a separate output folder called '/tmp/MNLI-MM/' in addition to '/tmp/MNLI/'.
|
||||
|
||||
#### Fine-tuning XLNet model on the STS-B regression task
|
||||
|
||||
This example code fine-tunes XLNet on the STS-B corpus using parallel training on a server with 4 V100 GPUs.
|
||||
Parallel training is a simple way to use several GPUs (but is slower and less flexible than distributed training, see below).
|
||||
|
||||
```shell
|
||||
export GLUE_DIR=/path/to/glue
|
||||
|
||||
python ./examples/text-classification/run_glue.py \
|
||||
--model_name_or_path xlnet-large-cased \
|
||||
--do_train \
|
||||
--do_eval \
|
||||
--task_name=sts-b \
|
||||
--data_dir=${GLUE_DIR}/STS-B \
|
||||
--output_dir=./proc_data/sts-b-110 \
|
||||
--max_seq_length=128 \
|
||||
--per_device_eval_batch_size=8 \
|
||||
--per_device_train_batch_size=8 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--max_steps=1200 \
|
||||
--model_name=xlnet-large-cased \
|
||||
--overwrite_output_dir \
|
||||
--overwrite_cache \
|
||||
--warmup_steps=120
|
||||
```
|
||||
|
||||
On this machine we thus have a batch size of 32, please increase `gradient_accumulation_steps` to reach the same batch size if you have a smaller machine. These hyper-parameters should result in a Pearson correlation coefficient of `+0.917` on the development set.
|
||||
|
||||
#### Fine-tuning Bert model on the MRPC classification task
|
||||
|
||||
This example code fine-tunes the Bert Whole Word Masking model on the Microsoft Research Paraphrase Corpus (MRPC) corpus using distributed training on 8 V100 GPUs to reach a F1 > 92.
|
||||
|
||||
```bash
|
||||
python -m torch.distributed.launch --nproc_per_node 8 ./examples/text-classification/run_glue.py \
|
||||
--model_name_or_path bert-large-uncased-whole-word-masking \
|
||||
--task_name MRPC \
|
||||
--do_train \
|
||||
--do_eval \
|
||||
--data_dir $GLUE_DIR/MRPC/ \
|
||||
--max_seq_length 128 \
|
||||
--per_device_eval_batch_size=8 \
|
||||
--per_device_train_batch_size=8 \
|
||||
--learning_rate 2e-5 \
|
||||
--num_train_epochs 3.0 \
|
||||
--output_dir /tmp/mrpc_output/ \
|
||||
--overwrite_output_dir \
|
||||
--overwrite_cache \
|
||||
```
|
||||
|
||||
Training with these hyper-parameters gave us the following results:
|
||||
|
||||
```bash
|
||||
acc = 0.8823529411764706
|
||||
acc_and_f1 = 0.901702786377709
|
||||
eval_loss = 0.3418912578906332
|
||||
f1 = 0.9210526315789473
|
||||
global_step = 174
|
||||
loss = 0.07231863956341798
|
||||
```
|
||||
|
||||
### `run_squad.py`: Fine-tuning on SQuAD for question-answering
|
||||
|
||||
This example code fine-tunes BERT on the SQuAD dataset using distributed training on 8 V100 GPUs and Bert Whole Word Masking uncased model to reach a F1 > 93 on SQuAD:
|
||||
|
||||
```bash
|
||||
python -m torch.distributed.launch --nproc_per_node=8 ./examples/question-answering/run_squad.py \
|
||||
--model_type bert \
|
||||
--model_name_or_path bert-large-uncased-whole-word-masking \
|
||||
--do_train \
|
||||
--do_eval \
|
||||
--train_file $SQUAD_DIR/train-v1.1.json \
|
||||
--predict_file $SQUAD_DIR/dev-v1.1.json \
|
||||
--learning_rate 3e-5 \
|
||||
--num_train_epochs 2 \
|
||||
--max_seq_length 384 \
|
||||
--doc_stride 128 \
|
||||
--output_dir ../models/wwm_uncased_finetuned_squad/ \
|
||||
--per_device_eval_batch_size=3 \
|
||||
--per_device_train_batch_size=3 \
|
||||
```
|
||||
|
||||
Training with these hyper-parameters gave us the following results:
|
||||
|
||||
```bash
|
||||
python $SQUAD_DIR/evaluate-v1.1.py $SQUAD_DIR/dev-v1.1.json ../models/wwm_uncased_finetuned_squad/predictions.json
|
||||
{"exact_match": 86.91579943235573, "f1": 93.1532499015869}
|
||||
```
|
||||
|
||||
This is the model provided as `bert-large-uncased-whole-word-masking-finetuned-squad`.
|
||||
|
||||
### `run_generation.py`: Text generation with GPT, GPT-2, CTRL, Transformer-XL and XLNet
|
||||
|
||||
A conditional generation script is also included to generate text from a prompt.
|
||||
The generation script includes the [tricks](https://github.com/rusiaaman/XLNet-gen#methodology) proposed by Aman Rusia to get high-quality generation with memory models like Transformer-XL and XLNet (include a predefined text to make short inputs longer).
|
||||
|
||||
Here is how to run the script with the small version of OpenAI GPT-2 model:
|
||||
|
||||
```shell
|
||||
python ./examples/text-generation/run_generation.py \
|
||||
--model_type=gpt2 \
|
||||
--length=20 \
|
||||
--model_name_or_path=gpt2 \
|
||||
```
|
||||
|
||||
and from the Salesforce CTRL model:
|
||||
```shell
|
||||
python ./examples/text-generation/run_generation.py \
|
||||
--model_type=ctrl \
|
||||
--length=20 \
|
||||
--model_name_or_path=ctrl \
|
||||
--temperature=0 \
|
||||
--repetition_penalty=1.2 \
|
||||
```
|
||||
|
||||
## Quick tour of model sharing
|
||||
|
||||
Starting with `v2.2.2`, you can now upload and share your fine-tuned models with the community, using the <abbr title="Command-line interface">CLI</abbr> that's built-in to the library.
|
||||
|
||||
**First, create an account on [https://huggingface.co/join](https://huggingface.co/join)**. Optionally, join an existing organization or create a new one. Then:
|
||||
|
||||
```shell
|
||||
transformers-cli login
|
||||
# log in using the same credentials as on huggingface.co
|
||||
```
|
||||
Upload your model:
|
||||
```shell
|
||||
transformers-cli upload ./path/to/pretrained_model/
|
||||
|
||||
# ^^ Upload folder containing weights/tokenizer/config
|
||||
# saved via `.save_pretrained()`
|
||||
|
||||
transformers-cli upload ./config.json [--filename folder/foobar.json]
|
||||
|
||||
# ^^ Upload a single file
|
||||
# (you can optionally override its filename, which can be nested inside a folder)
|
||||
```
|
||||
|
||||
If you want your model to be namespaced by your organization name rather than your username, add the following flag to any command:
|
||||
```shell
|
||||
--organization organization_name
|
||||
```
|
||||
|
||||
Your model will then be accessible through its identifier, a concatenation of your username (or organization name) and the folder name above:
|
||||
```python
|
||||
"username/pretrained_model"
|
||||
# or if an org:
|
||||
"organization_name/pretrained_model"
|
||||
```
|
||||
|
||||
**Please add a README.md model card** to the repo under `model_cards/` with: model description, training params (dataset, preprocessing, hardware used, hyperparameters), evaluation results, intended uses & limitations, etc.
|
||||
|
||||
Your model now has a page on huggingface.co/models 🔥
|
||||
|
||||
Anyone can load it from code:
|
||||
```python
|
||||
tokenizer = AutoTokenizer.from_pretrained("namespace/pretrained_model")
|
||||
model = AutoModel.from_pretrained("namespace/pretrained_model")
|
||||
```
|
||||
|
||||
List all your files on S3:
|
||||
```shell
|
||||
transformers-cli s3 ls
|
||||
```
|
||||
|
||||
You can also delete unneeded files:
|
||||
|
||||
```shell
|
||||
transformers-cli s3 rm …
|
||||
```
|
||||
|
||||
## Quick tour of pipelines
|
||||
|
||||
New in version `v2.3`: `Pipeline` are high-level objects which automatically handle tokenization, running your data through a transformers model
|
||||
and outputting the result in a structured object.
|
||||
|
||||
You can create `Pipeline` objects for the following down-stream tasks:
|
||||
|
||||
- `feature-extraction`: Generates a tensor representation for the input sequence
|
||||
- `ner`: Generates named entity mapping for each word in the input sequence.
|
||||
- `sentiment-analysis`: Gives the polarity (positive / negative) of the whole input sequence.
|
||||
- `text-classification`: Initialize a `TextClassificationPipeline` directly, or see `sentiment-analysis` for an example.
|
||||
- `question-answering`: Provided some context and a question refering to the context, it will extract the answer to the question in the context.
|
||||
- `fill-mask`: Takes an input sequence containing a masked token (e.g. `<mask>`) and return list of most probable filled sequences, with their probabilities.
|
||||
- `summarization`
|
||||
- `translation_xx_to_yy`
|
||||
|
||||
```python
|
||||
>>> from transformers import pipeline
|
||||
|
||||
# Allocate a pipeline for sentiment-analysis
|
||||
>>> nlp = pipeline('sentiment-analysis')
|
||||
>>> nlp('We are very happy to include pipeline into the transformers repository.')
|
||||
[{'label': 'POSITIVE', 'score': 0.9978193640708923}]
|
||||
|
||||
# Allocate a pipeline for question-answering
|
||||
>>> nlp = pipeline('question-answering')
|
||||
>>> nlp({
|
||||
... 'question': 'What is the name of the repository ?',
|
||||
... 'context': 'Pipeline have been included in the huggingface/transformers repository'
|
||||
... })
|
||||
{'score': 0.5135612454720828, 'start': 35, 'end': 59, 'answer': 'huggingface/transformers'}
|
||||
|
||||
```
|
||||
|
||||
## Migrating from pytorch-transformers to transformers
|
||||
|
||||
Here is a quick summary of what you should take care of when migrating from `pytorch-transformers` to `transformers`.
|
||||
|
||||
### Positional order of some models' keywords inputs (`attention_mask`, `token_type_ids`...) changed
|
||||
|
||||
To be able to use Torchscript (see #1010, #1204 and #1195) the specific order of some models **keywords inputs** (`attention_mask`, `token_type_ids`...) has been changed.
|
||||
|
||||
If you used to call the models with keyword names for keyword arguments, e.g. `model(inputs_ids, attention_mask=attention_mask, token_type_ids=token_type_ids)`, this should not cause any change.
|
||||
|
||||
If you used to call the models with positional inputs for keyword arguments, e.g. `model(inputs_ids, attention_mask, token_type_ids)`, you may have to double check the exact order of input arguments.
|
||||
|
||||
|
||||
## Migrating from pytorch-pretrained-bert to transformers
|
||||
|
||||
Here is a quick summary of what you should take care of when migrating from `pytorch-pretrained-bert` to `transformers`.
|
||||
|
||||
### Models always output `tuples`
|
||||
|
||||
The main breaking change when migrating from `pytorch-pretrained-bert` to `transformers` is that every model's forward method always outputs a `tuple` with various elements depending on the model and the configuration parameters.
|
||||
|
||||
The exact content of the tuples for each model is detailed in the models' docstrings and the [documentation](https://huggingface.co/transformers/).
|
||||
|
||||
In pretty much every case, you will be fine by taking the first element of the output as the output you previously used in `pytorch-pretrained-bert`.
|
||||
|
||||
Here is a `pytorch-pretrained-bert` to `transformers` conversion example for a `BertForSequenceClassification` classification model:
|
||||
|
||||
```python
|
||||
# Let's load our model
|
||||
model = BertForSequenceClassification.from_pretrained('bert-base-uncased')
|
||||
|
||||
# If you used to have this line in pytorch-pretrained-bert:
|
||||
loss = model(input_ids, labels=labels)
|
||||
|
||||
# Now just use this line in transformers to extract the loss from the output tuple:
|
||||
outputs = model(input_ids, labels=labels)
|
||||
loss = outputs[0]
|
||||
|
||||
# In transformers you can also have access to the logits:
|
||||
loss, logits = outputs[:2]
|
||||
|
||||
# And even the attention weights if you configure the model to output them (and other outputs too, see the docstrings and documentation)
|
||||
model = BertForSequenceClassification.from_pretrained('bert-base-uncased', output_attentions=True)
|
||||
outputs = model(input_ids, labels=labels)
|
||||
loss, logits, attentions = outputs
|
||||
```
|
||||
|
||||
### Using hidden states
|
||||
|
||||
By enabling the configuration option `output_hidden_states`, it was possible to retrieve the last hidden states of the encoder. In `pytorch-transformers` as well as `transformers` the return value has changed slightly: `all_hidden_states` now also includes the hidden state of the embeddings in addition to those of the encoding layers. This allows users to easily access the embeddings final state.
|
||||
|
||||
### Serialization
|
||||
|
||||
Breaking change in the `from_pretrained()` method:
|
||||
|
||||
1. Models are now set in evaluation mode by default when instantiated with the `from_pretrained()` method. To train them, don't forget to set them back in training mode (`model.train()`) to activate the dropout modules.
|
||||
|
||||
2. The additional `*input` and `**kwargs` arguments supplied to the `from_pretrained()` method used to be directly passed to the underlying model's class `__init__()` method. They are now used to update the model configuration attribute instead, which can break derived model classes built based on the previous `BertForSequenceClassification` examples. We are working on a way to mitigate this breaking change in [#866](https://github.com/huggingface/transformers/pull/866) by forwarding the model's `__init__()` method (i) the provided positional arguments and (ii) the keyword arguments which do not match any configuration class attributes.
|
||||
|
||||
Also, while not a breaking change, the serialization methods have been standardized and you probably should switch to the new method `save_pretrained(save_directory)` if you were using any other serialization method before.
|
||||
|
||||
Here is an example:
|
||||
|
||||
```python
|
||||
### Let's load a model and tokenizer
|
||||
model = BertForSequenceClassification.from_pretrained('bert-base-uncased')
|
||||
tokenizer = BertTokenizer.from_pretrained('bert-base-uncased')
|
||||
|
||||
### Do some stuff to our model and tokenizer
|
||||
# Ex: add new tokens to the vocabulary and embeddings of our model
|
||||
tokenizer.add_tokens(['[SPECIAL_TOKEN_1]', '[SPECIAL_TOKEN_2]'])
|
||||
model.resize_token_embeddings(len(tokenizer))
|
||||
# Train our model
|
||||
train(model)
|
||||
|
||||
### Now let's save our model and tokenizer to a directory
|
||||
model.save_pretrained('./my_saved_model_directory/')
|
||||
tokenizer.save_pretrained('./my_saved_model_directory/')
|
||||
|
||||
### Reload the model and the tokenizer
|
||||
model = BertForSequenceClassification.from_pretrained('./my_saved_model_directory/')
|
||||
tokenizer = BertTokenizer.from_pretrained('./my_saved_model_directory/')
|
||||
```
|
||||
|
||||
### Optimizers: BertAdam & OpenAIAdam are now AdamW, schedules are standard PyTorch schedules
|
||||
|
||||
The two optimizers previously included, `BertAdam` and `OpenAIAdam`, have been replaced by a single `AdamW` optimizer which has a few differences:
|
||||
|
||||
- it only implements weights decay correction,
|
||||
- schedules are now externals (see below),
|
||||
- gradient clipping is now also external (see below).
|
||||
|
||||
The new optimizer `AdamW` matches PyTorch `Adam` optimizer API and let you use standard PyTorch or apex methods for the schedule and clipping.
|
||||
|
||||
The schedules are now standard [PyTorch learning rate schedulers](https://pytorch.org/docs/stable/optim.html#how-to-adjust-learning-rate) and not part of the optimizer anymore.
|
||||
|
||||
Here is a conversion examples from `BertAdam` with a linear warmup and decay schedule to `AdamW` and the same schedule:
|
||||
|
||||
```python
|
||||
# Parameters:
|
||||
lr = 1e-3
|
||||
max_grad_norm = 1.0
|
||||
num_training_steps = 1000
|
||||
num_warmup_steps = 100
|
||||
warmup_proportion = float(num_warmup_steps) / float(num_training_steps) # 0.1
|
||||
|
||||
### Previously BertAdam optimizer was instantiated like this:
|
||||
optimizer = BertAdam(model.parameters(), lr=lr, schedule='warmup_linear', warmup=warmup_proportion, t_total=num_training_steps)
|
||||
### and used like this:
|
||||
for batch in train_data:
|
||||
loss = model(batch)
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
|
||||
### In Transformers, optimizer and schedules are splitted and instantiated like this:
|
||||
optimizer = AdamW(model.parameters(), lr=lr, correct_bias=False) # To reproduce BertAdam specific behavior set correct_bias=False
|
||||
scheduler = get_linear_schedule_with_warmup(optimizer, num_warmup_steps=num_warmup_steps, num_training_steps=num_training_steps) # PyTorch scheduler
|
||||
### and used like this:
|
||||
for batch in train_data:
|
||||
model.train()
|
||||
loss = model(batch)
|
||||
loss.backward()
|
||||
torch.nn.utils.clip_grad_norm_(model.parameters(), max_grad_norm) # Gradient clipping is not in AdamW anymore (so you can use amp without issue)
|
||||
optimizer.step()
|
||||
scheduler.step()
|
||||
optimizer.zero_grad()
|
||||
```
|
||||
|
||||
## Citation
|
||||
|
||||
|
||||
@@ -172,7 +172,6 @@ conversion utilities for the following models:
|
||||
converting_tensorflow_models
|
||||
migration
|
||||
contributing
|
||||
testing
|
||||
serialization
|
||||
|
||||
.. toctree::
|
||||
|
||||
@@ -1,131 +1,109 @@
|
||||
AutoClasses
|
||||
AutoModels
|
||||
-----------
|
||||
|
||||
In many cases, the architecture you want to use can be guessed from the name or the path of the pretrained model you
|
||||
are supplying to the :obj:`from_pretrained()` method.
|
||||
are supplying to the ``from_pretrained`` method.
|
||||
|
||||
AutoClasses are here to do this job for you so that you automatically retrieve the relevant model given the name/path
|
||||
to the pretrained weights/config/vocabulary.
|
||||
to the pretrained weights/config/vocabulary:
|
||||
|
||||
Instantiating one of :class:`~transformers.AutoConfig`, :class:`~transformers.AutoModel`, and
|
||||
:class:`~transformers.AutoTokenizer` will directly create a class of the relevant architecture. For instance
|
||||
Instantiating one of ``AutoModel``, ``AutoConfig`` and ``AutoTokenizer`` will directly create a class of the relevant
|
||||
architecture (ex: ``model = AutoModel.from_pretrained('bert-base-cased')`` will create a instance of
|
||||
:class:`~transformers.BertModel`).
|
||||
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
model = AutoModel.from_pretrained('bert-base-cased')
|
||||
|
||||
will create a model that is an instance of :class:`~transformers.BertModel`.
|
||||
|
||||
There is one class of :obj:`AutoModel` for each task, and for each backend (PyTorch or TensorFlow).
|
||||
|
||||
|
||||
AutoConfig
|
||||
~~~~~~~~~~
|
||||
``AutoConfig``
|
||||
~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.AutoConfig
|
||||
:members:
|
||||
|
||||
|
||||
AutoTokenizer
|
||||
~~~~~~~~~~~~~
|
||||
``AutoTokenizer``
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.AutoTokenizer
|
||||
:members:
|
||||
|
||||
|
||||
AutoModel
|
||||
~~~~~~~~~
|
||||
``AutoModel``
|
||||
~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.AutoModel
|
||||
:members:
|
||||
|
||||
|
||||
AutoModelForPreTraining
|
||||
~~~~~~~~~~~~~~~~~~~~~~~
|
||||
``AutoModelForPreTraining``
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.AutoModelForPreTraining
|
||||
:members:
|
||||
|
||||
|
||||
AutoModelWithLMHead
|
||||
~~~~~~~~~~~~~~~~~~~
|
||||
``AutoModelWithLMHead``
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.AutoModelWithLMHead
|
||||
:members:
|
||||
|
||||
|
||||
AutoModelForSequenceClassification
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
``AutoModelForSequenceClassification``
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.AutoModelForSequenceClassification
|
||||
:members:
|
||||
|
||||
|
||||
AutoModelForMultipleChoice
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.AutoModelForMultipleChoice
|
||||
:members:
|
||||
|
||||
|
||||
AutoModelForTokenClassification
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.AutoModelForTokenClassification
|
||||
:members:
|
||||
|
||||
|
||||
AutoModelForQuestionAnswering
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
``AutoModelForQuestionAnswering``
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.AutoModelForQuestionAnswering
|
||||
:members:
|
||||
|
||||
|
||||
TFAutoModel
|
||||
~~~~~~~~~~~
|
||||
``AutoModelForTokenClassification``
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.AutoModelForTokenClassification
|
||||
:members:
|
||||
|
||||
``TFAutoModel``
|
||||
~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.TFAutoModel
|
||||
:members:
|
||||
|
||||
|
||||
TFAutoModelForPreTraining
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
``TFAutoModelForPreTraining``
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.TFAutoModelForPreTraining
|
||||
:members:
|
||||
|
||||
|
||||
TFAutoModelWithLMHead
|
||||
~~~~~~~~~~~~~~~~~~~~~
|
||||
``TFAutoModelWithLMHead``
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.TFAutoModelWithLMHead
|
||||
:members:
|
||||
|
||||
|
||||
TFAutoModelForSequenceClassification
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
``TFAutoModelForSequenceClassification``
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.TFAutoModelForSequenceClassification
|
||||
:members:
|
||||
|
||||
|
||||
TFAutoModelForMultipleChoice
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.TFAutoModelForMultipleChoice
|
||||
:members:
|
||||
|
||||
|
||||
TFAutoModelForTokenClassification
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.TFAutoModelForTokenClassification
|
||||
:members:
|
||||
|
||||
|
||||
TFAutoModelForQuestionAnswering
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
``TFAutoModelForQuestionAnswering``
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.TFAutoModelForQuestionAnswering
|
||||
:members:
|
||||
|
||||
|
||||
``TFAutoModelForTokenClassification``
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
.. autoclass:: transformers.TFAutoModelForTokenClassification
|
||||
:members:
|
||||
|
||||
@@ -7,7 +7,7 @@ The effectiveness of initializing sequence-to-sequence models with pre-trained c
|
||||
|
||||
After such an :class:`~transformers.EncoderDecoderModel` has been trained / fine-tuned, it can be saved / loaded just like any other models (see Examples for more information).
|
||||
|
||||
An application of this architecture could be to leverage two pre-trained :obj:`transformers.BertModel` models as the encoder and decoder for a summarization model as was shown in: `Text Summarization with Pretrained Encoders <https://arxiv.org/abs/1908.08345>`_ by Yang Liu and Mirella Lapata.
|
||||
An application of this architecture could be to leverage two pre-trained :obj:`transformers.BertModel` models as the encoder and decoder for a summarization model as was shown in: `Text Summarization with Pretrained Encoders <https://arxiv.org/abs/1910.13461>`_ by Yang Liu and Mirella Lapata.
|
||||
|
||||
|
||||
``EncoderDecoderConfig``
|
||||
|
||||
@@ -1,936 +0,0 @@
|
||||
Testing
|
||||
==========
|
||||
|
||||
|
||||
Let's take a look at how 🤗 Transformer models are tested and how you can write new tests and improve the existing ones.
|
||||
|
||||
There are 2 test suites in the repository:
|
||||
|
||||
1. ``tests`` -- tests for the general API
|
||||
2. ``examples`` -- tests primarily for various applications that aren't part of the API
|
||||
|
||||
How transformers are tested
|
||||
---------------------------
|
||||
|
||||
1. Once a PR is submitted it gets tested with 9 CircleCi jobs. Every new commit to that PR gets retested. These jobs are defined in this `config file <https://github.com/huggingface/transformers/blob/master/.circleci/config.yml>`__, so that if needed you can reproduce the same environment on your machine.
|
||||
|
||||
These CI jobs don't run ``@slow`` tests.
|
||||
|
||||
2. There are 3 jobs run by `github actions <https://github.com/huggingface/transformers/actions>`__:
|
||||
|
||||
* `torch hub integration <https://github.com/huggingface/transformers/blob/master/.github/workflows/github-torch-hub.yml>`__: checks whether torch hub integration works.
|
||||
|
||||
* `self-hosted (push) <https://github.com/huggingface/transformers/blob/master/.github/workflows/self-push.yml>`__: runs fast tests on GPU only on commits on ``master``. It only runs if a commit on ``master`` has updated the code in one of the following folders: ``src``, ``tests``, ``.github`` (to prevent running on added model cards, notebooks, etc.)
|
||||
|
||||
* `self-hosted runner <https://github.com/huggingface/transformers/blob/master/.github/workflows/self-scheduled.yml>`__: runs slow tests on ``tests`` and ``examples``:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
RUN_SLOW=1 USE_CUDA=1 pytest tests/
|
||||
RUN_SLOW=1 USE_CUDA=1 pytest examples/
|
||||
|
||||
The results can be observed `here <https://github.com/huggingface/transformers/actions>`__.
|
||||
|
||||
|
||||
|
||||
Running tests
|
||||
-------------
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
Choosing which tests to run
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
This document goes into many details of how tests can be run. If after reading everything, you need even more details you will find them `here <https://docs.pytest.org/en/latest/usage.html>`__.
|
||||
|
||||
Here are some most useful ways of running tests.
|
||||
|
||||
Run all:
|
||||
|
||||
.. code-block:: console
|
||||
|
||||
pytest
|
||||
|
||||
or:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
make test
|
||||
|
||||
Note that the latter is defined as:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
python -m pytest -n auto --dist=loadfile -s -v ./tests/
|
||||
|
||||
which tells pytest to:
|
||||
|
||||
* run as many test processes as they are CPU cores (which could be too many if you don't have a ton of RAM!)
|
||||
* ensure that all tests from the same file will be run by the same test process
|
||||
* do not capture output
|
||||
* run in verbose mode
|
||||
|
||||
|
||||
|
||||
Getting the list of all tests
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
All tests of the test suite:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest --collect-only -q
|
||||
|
||||
All tests of a given test file:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest tests/test_optimization.py --collect-only -q
|
||||
|
||||
|
||||
|
||||
Run a specific test module
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
To run an individual test module:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest tests/test_logging.py
|
||||
|
||||
|
||||
Run specific tests
|
||||
~~~~~~~~~~~~~~~~~~
|
||||
|
||||
Since unittest is used inside most of the tests, to run specific subtests you need to know the name of the unittest class containing those tests. For example, it could be:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest tests/test_optimization.py::OptimizationTest::test_adam_w
|
||||
|
||||
Here:
|
||||
|
||||
* ``tests/test_optimization.py`` - the file with tests
|
||||
* ``OptimizationTest`` - the name of the class
|
||||
* ``test_adam_w`` - the name of the specific test function
|
||||
|
||||
If the file contains multiple classes, you can choose to run only tests of a given class. For example:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest tests/test_optimization.py::OptimizationTest
|
||||
|
||||
|
||||
will run all the tests inside that class.
|
||||
|
||||
As mentioned earlier you can see what tests are contained inside the ``OptimizationTest`` class by running:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest tests/test_optimization.py::OptimizationTest --collect-only -q
|
||||
|
||||
|
||||
You can run tests by keyword expressions.
|
||||
|
||||
To run only tests whose name contains ``adam``:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest -k adam tests/test_optimization.py
|
||||
|
||||
To run all tests except those whose name contains ``adam``:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest -k "not adam" tests/test_optimization.py
|
||||
|
||||
And you can combine the two patterns in one:
|
||||
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest -k "ada and not adam" tests/test_optimization.py
|
||||
|
||||
|
||||
|
||||
Run only modified tests
|
||||
~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
You can run the tests related to the unstaged files or the current branch (according to Git) by using `pytest-picked <https://github.com/anapaulagomes/pytest-picked>`__. This is a great way of quickly testing your changes didn't break anything, since it won't run the tests related to files you didn't touch.
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pip install pytest-picked
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest --picked
|
||||
|
||||
All tests will be run from files and folders which are modified, but not
|
||||
yet committed.
|
||||
|
||||
Automatically rerun failed tests on source modification
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
`pytest-xdist <https://github.com/pytest-dev/pytest-xdist>`__ provides a
|
||||
very useful feature of detecting all failed tests, and then waiting for
|
||||
you to modify files and continuously re-rerun those failing tests until
|
||||
they pass while you fix them. So that you don't need to re start pytest
|
||||
after you made the fix. This is repeated until all tests pass after
|
||||
which again a full run is performed.
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pip install pytest-xdist
|
||||
|
||||
To enter the mode: ``pytest -f`` or ``pytest --looponfail``
|
||||
|
||||
File changes are detected by looking at ``looponfailroots`` root
|
||||
directories and all of their contents (recursively). If the default for
|
||||
this value does not work for you, you can change it in your project by
|
||||
setting a configuration option in ``setup.cfg``:
|
||||
|
||||
.. code-block:: ini
|
||||
|
||||
[tool:pytest]
|
||||
looponfailroots = transformers tests
|
||||
|
||||
or ``pytest.ini``/``tox.ini`` files:
|
||||
|
||||
.. code-block:: ini
|
||||
|
||||
[pytest]
|
||||
looponfailroots = transformers tests
|
||||
|
||||
This would lead to only looking for file changes in the respective
|
||||
directories, specified relatively to the ini-file’s directory.
|
||||
|
||||
`pytest-watch <https://github.com/joeyespo/pytest-watch>`__ is an
|
||||
alternative implementation of this functionality.
|
||||
|
||||
|
||||
Skip a test module
|
||||
~~~~~~~~~~~~~~~~~~
|
||||
|
||||
If you want to run all test modules, except a few you can exclude them by giving an explicit list of tests to run. For example, to run all except ``test_modeling_*.py`` tests:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest `ls -1 tests/*py | grep -v test_modeling`
|
||||
|
||||
|
||||
Clearing state
|
||||
~~~~~~~~~~~~~~
|
||||
|
||||
CI builds and when isolation is important (against speed), cache should
|
||||
be cleared:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest --cache-clear tests
|
||||
|
||||
Running tests in parallel
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
As mentioned earlier ``make test`` runs tests in parallel via ``pytest-xdist`` plugin (``-n X`` argument, e.g. ``-n 2`` to run 2 parallel jobs).
|
||||
|
||||
``pytest-xdist``'s ``--dist=`` option allows one to control how the tests are grouped. ``--dist=loadfile`` puts the tests located in one file onto the same process.
|
||||
|
||||
Since the order of executed tests is different and unpredictable, if
|
||||
running the test suite with ``pytest-xdist`` produces failures (meaning
|
||||
we have some undetected coupled tests), use
|
||||
`pytest-replay <https://github.com/ESSS/pytest-replay>`__ to replay the
|
||||
tests in the same order, which should help with then somehow reducing
|
||||
that failing sequence to a minimum.
|
||||
|
||||
Test order and repetition
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
It's good to repeat the tests several times, in sequence, randomly, or
|
||||
in sets, to detect any potential inter-dependency and state-related bugs
|
||||
(tear down). And the straightforward multiple repetition is just good to
|
||||
detect some problems that get uncovered by randomness of DL.
|
||||
|
||||
|
||||
Repeat tests
|
||||
^^^^^^^^^^^^
|
||||
|
||||
* `pytest-flakefinder <https://github.com/dropbox/pytest-flakefinder>`__:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pip install pytest-flakefinder
|
||||
|
||||
And then run every test multiple times (50 by default):
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest --flake-finder --flake-runs=5 tests/test_failing_test.py
|
||||
|
||||
.. note::
|
||||
This plugin doesn't work with ``-n`` flag from ``pytest-xdist``.
|
||||
|
||||
.. note::
|
||||
There is another plugin ``pytest-repeat``, but it doesn't work with ``unittest``.
|
||||
|
||||
|
||||
Run tests in a random order
|
||||
^^^^^^^^^^^^^^^^^^^^^^^^^^^
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pip install pytest-random-order
|
||||
|
||||
Important: the presence of ``pytest-random-order`` will automatically
|
||||
randomize tests, no configuration change or command line options is
|
||||
required.
|
||||
|
||||
As explained earlier this allows detection of coupled tests - where one
|
||||
test's state affects the state of another. When ``pytest-random-order``
|
||||
is installed it will print the random seed it used for that session,
|
||||
e.g:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest tests
|
||||
[...]
|
||||
Using --random-order-bucket=module
|
||||
Using --random-order-seed=573663
|
||||
|
||||
So that if the given particular sequence fails, you can reproduce it by
|
||||
adding that exact seed, e.g.:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest --random-order-seed=573663
|
||||
[...]
|
||||
Using --random-order-bucket=module
|
||||
Using --random-order-seed=573663
|
||||
|
||||
It will only reproduce the exact order if you use the exact same list of
|
||||
tests (or no list at all). Once you start to manually narrowing
|
||||
down the list you can no longer rely on the seed, but have to list them
|
||||
manually in the exact order they failed and tell pytest to not randomize
|
||||
them instead using ``--random-order-bucket=none``, e.g.:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest --random-order-bucket=none tests/test_a.py tests/test_c.py tests/test_b.py
|
||||
|
||||
To disable the shuffling for all tests:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest --random-order-bucket=none
|
||||
|
||||
By default ``--random-order-bucket=module`` is implied, which will
|
||||
shuffle the files on the module levels. It can also shuffle on
|
||||
``class``, ``package``, ``global`` and ``none`` levels. For the complete
|
||||
details please see its `documentation <https://github.com/jbasko/pytest-random-order>`__.
|
||||
|
||||
Another randomization alternative is: ``pytest-randomly`` <https://github.com/pytest-dev/pytest-randomly>`__. This module has a very similar functionality/interface, but it doesn't have the bucket modes available in ``pytest-random-order``. It has the same problem of imposing itself once installed.
|
||||
|
||||
Look and feel variations
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
pytest-sugar
|
||||
^^^^^^^^^^^^
|
||||
|
||||
`pytest-sugar <https://github.com/Frozenball/pytest-sugar>`__ is a
|
||||
plugin that improves the look-n-feel, adds a progressbar, and show tests
|
||||
that fail and the assert instantly. It gets activated automatically upon
|
||||
installation.
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pip install pytest-sugar
|
||||
|
||||
To run tests without it, run:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest -p no:sugar
|
||||
|
||||
or uninstall it.
|
||||
|
||||
|
||||
|
||||
Report each sub-test name and its progress
|
||||
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
|
||||
|
||||
For a single or a group of tests via ``pytest`` (after
|
||||
``pip install pytest-pspec``):
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest --pspec tests/test_optimization.py
|
||||
|
||||
|
||||
|
||||
Instantly shows failed tests
|
||||
^^^^^^^^^^^^^^^^^^^^^^^^^^^^
|
||||
|
||||
`pytest-instafail <https://github.com/pytest-dev/pytest-instafail>`__
|
||||
shows failures and errors instantly instead of waiting until the end of
|
||||
test session.
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pip install pytest-instafail
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest --instafail
|
||||
|
||||
To GPU or not to GPU
|
||||
~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
On a GPU-enabled setup, to test in CPU-only mode add ``CUDA_VISIBLE_DEVICES=""``:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
CUDA_VISIBLE_DEVICES="" pytest tests/test_logging.py
|
||||
|
||||
or if you have multiple gpus, you can tell which one to use in this test session, e.g. to use only the second gpu if you have gpus ``0`` and ``1``, you can run:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
CUDA_VISIBLE_DEVICES="1" pytest tests/test_logging.py
|
||||
|
||||
This is handy when you want to run different tasks on different GPUs.
|
||||
|
||||
And we have these decorators that require the condition described by the marker.
|
||||
|
||||
```
|
||||
@require_torch
|
||||
@require_tf
|
||||
@require_multigpu
|
||||
@require_non_multigpu
|
||||
@require_torch_tpu
|
||||
@require_torch_and_cuda
|
||||
```
|
||||
|
||||
This section will be expanded soon once our work in progress on those decorators is finished.
|
||||
|
||||
Inside tests:
|
||||
|
||||
* How many GPUs are available:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
torch.cuda.device_count()
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
Output capture
|
||||
~~~~~~~~~~~~~~
|
||||
|
||||
During test execution any output sent to ``stdout`` and ``stderr`` is
|
||||
captured. If a test or a setup method fails, its according captured
|
||||
output will usually be shown along with the failure traceback.
|
||||
|
||||
To disable output capturing and to get the ``stdout`` and ``stderr``
|
||||
normally, use ``-s`` or ``--capture=no``:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest -s tests/test_logging.py
|
||||
|
||||
To send test results to JUnit format output:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
py.test tests --junitxml=result.xml
|
||||
|
||||
|
||||
Color control
|
||||
~~~~~~~~~~~~~
|
||||
|
||||
To have no color (e.g., yellow on white background is not readable):
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest --color=no tests/test_logging.py
|
||||
|
||||
|
||||
|
||||
Sending test report to online pastebin service
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
Creating a URL for each test failure:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest --pastebin=failed tests/test_logging.py
|
||||
|
||||
This will submit test run information to a remote Paste service and
|
||||
provide a URL for each failure. You may select tests as usual or add for
|
||||
example -x if you only want to send one particular failure.
|
||||
|
||||
Creating a URL for a whole test session log:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest --pastebin=all tests/test_logging.py
|
||||
|
||||
|
||||
|
||||
Writing tests
|
||||
-------------
|
||||
|
||||
🤗 transformers tests are based on ``unittest``, but run by ``pytest``, so most of the time features from both systems can be used.
|
||||
|
||||
You can read `here <https://docs.pytest.org/en/stable/unittest.html>`__ which features are supported, but the important thing to remember is that most ``pytest`` fixtures don't work. Neither parametrization, but we use the module ``parameterized`` that works in a similar way.
|
||||
|
||||
|
||||
Parametrization
|
||||
~~~~~~~~~~~~~~~
|
||||
|
||||
Often, there is a need to run the same test multiple times, but with different arguments. It could be done from within the test, but then there is no way of running that test for just one set of arguments.
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
# test_this1.py
|
||||
import unittest
|
||||
from parameterized import parameterized
|
||||
class TestMathUnitTest(unittest.TestCase):
|
||||
@parameterized.expand([
|
||||
("negative", -1.5, -2.0),
|
||||
("integer", 1, 1.0),
|
||||
("large fraction", 1.6, 1),
|
||||
])
|
||||
def test_floor(self, name, input, expected):
|
||||
assert_equal(math.floor(input), expected)
|
||||
|
||||
Now, by default this test will be run 3 times, each time with the last 3 arguments of ``test_floor`` being assigned the corresponding arguments in the parameter list.
|
||||
|
||||
and you could run just the ``negative`` and ``integer`` sets of params with:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest -k "negative and integer" tests/test_mytest.py
|
||||
|
||||
or all but ``negative`` sub-tests, with:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest -k "not negative" tests/test_mytest.py
|
||||
|
||||
Besides using the ``-k`` filter that was just mentioned, you can find out the exact name of each sub-test and run any or all of them using their exact names.
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest test_this1.py --collect-only -q
|
||||
|
||||
and it will list:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
test_this1.py::TestMathUnitTest::test_floor_0_negative
|
||||
test_this1.py::TestMathUnitTest::test_floor_1_integer
|
||||
test_this1.py::TestMathUnitTest::test_floor_2_large_fraction
|
||||
|
||||
So now you can run just 2 specific sub-tests:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest test_this1.py::TestMathUnitTest::test_floor_0_negative test_this1.py::TestMathUnitTest::test_floor_1_integer
|
||||
|
||||
The module `parameterized <https://pypi.org/project/parameterized/>`__ which is already in the developer dependencies of ``transformers`` works for both: ``unittests`` and ``pytest`` tests.
|
||||
|
||||
If, however, the test is not a ``unittest``, you may use ``pytest.mark.parametrize`` (or you may see it being used in some existing tests, mostly under ``examples``).
|
||||
|
||||
Here is the same example, this time using ``pytest``'s ``parametrize`` marker:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
# test_this2.py
|
||||
import pytest
|
||||
@pytest.mark.parametrize(
|
||||
"name, input, expected",
|
||||
[
|
||||
("negative", -1.5, -2.0),
|
||||
("integer", 1, 1.0),
|
||||
("large fraction", 1.6, 1),
|
||||
],
|
||||
)
|
||||
def test_floor(name, input, expected):
|
||||
assert_equal(math.floor(input), expected)
|
||||
|
||||
Same as with ``parameterized``, with ``pytest.mark.parametrize`` you can have a fine control over which sub-tests are run, if the ``-k`` filter doesn't do the job. Except, this parametrization function creates a slightly different set of names for the sub-tests. Here is what they look like:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest test_this2.py --collect-only -q
|
||||
|
||||
and it will list:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
test_this2.py::test_floor[integer-1-1.0]
|
||||
test_this2.py::test_floor[negative--1.5--2.0]
|
||||
test_this2.py::test_floor[large fraction-1.6-1]
|
||||
|
||||
So now you can run just the specific test:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest test_this2.py::test_floor[negative--1.5--2.0] test_this2.py::test_floor[integer-1-1.0]
|
||||
|
||||
as in the previous example.
|
||||
|
||||
|
||||
|
||||
Temporary files and directories
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
Using unique temporary files and directories are essential for parallel test running, so that the tests won't overwrite each other's data. Also we want to get the temp files and directories removed at the end of each test that created them. Therefore, using packages like ``tempfile``, which address these needs is essential.
|
||||
|
||||
However, when debugging tests, you need to be able to see what goes into the temp file or directory and you want to know it's exact path and not having it randomized on every test re-run.
|
||||
|
||||
A helper class :obj:`transformers.test_utils.TestCasePlus` is best used for such purposes. It's a sub-class of :obj:`unittest.TestCase`, so we can easily inherit from it in the test modules.
|
||||
|
||||
Here is an example of its usage:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
from transformers.testing_utils import TestCasePlus
|
||||
class ExamplesTests(TestCasePlus):
|
||||
def test_whatever(self):
|
||||
tmp_dir = self.get_auto_remove_tmp_dir()
|
||||
|
||||
This code creates a unique temporary directory, and sets :obj:`tmp_dir` to its location.
|
||||
|
||||
In this and all the following scenarios the temporary directory will be auto-removed at the end of test, unless ``after=False`` is passed to the helper function.
|
||||
|
||||
* Create a temporary directory of my choice and delete it at the end - useful for debugging when you want to monitor a specific directory:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
def test_whatever(self):
|
||||
tmp_dir = self.get_auto_remove_tmp_dir(tmp_dir="./tmp/run/test")
|
||||
|
||||
* Create a temporary directory of my choice and do not delete it at the end---useful for when you want to look at the temp results:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
def test_whatever(self):
|
||||
tmp_dir = self.get_auto_remove_tmp_dir(tmp_dir="./tmp/run/test", after=False)
|
||||
|
||||
* Create a temporary directory of my choice and ensure to delete it right away---useful for when you disabled deletion in the previous test run and want to make sure the that temporary directory is empty before the new test is run:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
def test_whatever(self):
|
||||
tmp_dir = self.get_auto_remove_tmp_dir(tmp_dir="./tmp/run/test", before=True)
|
||||
|
||||
.. note::
|
||||
In order to run the equivalent of ``rm -r`` safely, only subdirs of the project repository checkout are allowed if an explicit obj:`tmp_dir` is used, so that by mistake no ``/tmp`` or similar important part of the filesystem will get nuked. i.e. please always pass paths that start with ``./``.
|
||||
|
||||
.. note::
|
||||
Each test can register multiple temporary directories and they all will get auto-removed, unless requested otherwise.
|
||||
|
||||
|
||||
Skipping tests
|
||||
~~~~~~~~~~~~~~
|
||||
|
||||
This is useful when a bug is found and a new test is written, yet the
|
||||
bug is not fixed yet. In order to be able to commit it to the main
|
||||
repository we need make sure it's skipped during ``make test``.
|
||||
|
||||
Methods:
|
||||
|
||||
- A **skip** means that you expect your test to pass only if some
|
||||
conditions are met, otherwise pytest should skip running the test
|
||||
altogether. Common examples are skipping windows-only tests on
|
||||
non-windows platforms, or skipping tests that depend on an external
|
||||
resource which is not available at the moment (for example a
|
||||
database).
|
||||
|
||||
- A **xfail** means that you expect a test to fail for some reason. A
|
||||
common example is a test for a feature not yet implemented, or a bug
|
||||
not yet fixed. When a test passes despite being expected to fail
|
||||
(marked with pytest.mark.xfail), it’s an xpass and will be reported
|
||||
in the test summary.
|
||||
|
||||
One of the important differences between the two is that ``skip``
|
||||
doesn't run the test, and ``xfail`` does. So if the code that's buggy
|
||||
causes some bad state that will affect other tests, do not use
|
||||
``xfail``.
|
||||
|
||||
Implementation
|
||||
^^^^^^^^^^^^^^
|
||||
|
||||
- Here is how to skip whole test unconditionally:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
@unittest.skip("this bug needs to be fixed")
|
||||
def test_feature_x():
|
||||
|
||||
or via pytest:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
@pytest.mark.skip(reason="this bug needs to be fixed")
|
||||
|
||||
or the ``xfail`` way:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
@pytest.mark.xfail
|
||||
def test_feature_x():
|
||||
|
||||
Here is how to skip a test based on some internal check inside the test:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
def test_feature_x():
|
||||
if not has_something():
|
||||
pytest.skip("unsupported configuration")
|
||||
|
||||
or the whole module:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
import pytest
|
||||
if not pytest.config.getoption("--custom-flag"):
|
||||
pytest.skip("--custom-flag is missing, skipping tests", allow_module_level=True)
|
||||
|
||||
or the ``xfail`` way:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
def test_feature_x():
|
||||
pytest.xfail("expected to fail until bug XYZ is fixed")
|
||||
|
||||
Here is how to skip all tests in a module if some import is missing:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
docutils = pytest.importorskip("docutils", minversion="0.3")
|
||||
|
||||
- Skip a test based on a condition:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
@pytest.mark.skipif(sys.version_info < (3,6), reason="requires python3.6 or higher")
|
||||
def test_feature_x():
|
||||
|
||||
or:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
@unittest.skipIf(torch_device == "cpu", "Can't do half precision")
|
||||
def test_feature_x():
|
||||
|
||||
or skip the whole module:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
@pytest.mark.skipif(sys.platform == 'win32', reason="does not run on windows")
|
||||
class TestClass():
|
||||
def test_feature_x(self):
|
||||
|
||||
More details, example and ways are `here <https://docs.pytest.org/en/latest/skipping.html>`__.
|
||||
|
||||
Custom markers
|
||||
~~~~~~~~~~~~~~
|
||||
|
||||
* Slow tests
|
||||
|
||||
Tests that are too slow (e.g. once downloading huge model files) are marked with:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
from transformers.testing_utils import slow
|
||||
@slow
|
||||
def test_integration_foo():
|
||||
|
||||
To run such tests set ``RUN_SLOW=1`` env var, e.g.:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
RUN_SLOW=1 pytest tests
|
||||
|
||||
|
||||
Testing the stdout/stderr output
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
In order to test functions that write to ``stdout`` and/or ``stderr``,
|
||||
the test can access those streams using the ``pytest``'s `capsys
|
||||
system <https://docs.pytest.org/en/latest/capture.html>`__. Here is how
|
||||
this is accomplished:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
import sys
|
||||
def print_to_stdout(s): print(s)
|
||||
def print_to_stderr(s): sys.stderr.write(s)
|
||||
def test_result_and_stdout(capsys):
|
||||
msg = "Hello"
|
||||
print_to_stdout(msg)
|
||||
print_to_stderr(msg)
|
||||
out, err = capsys.readouterr() # consume the captured output streams
|
||||
# optional: if you want to replay the consumed streams:
|
||||
sys.stdout.write(out)
|
||||
sys.stderr.write(err)
|
||||
# test:
|
||||
assert msg in out
|
||||
assert msg in err
|
||||
|
||||
And, of course, most of the time, ``stderr`` will come as a part of an
|
||||
exception, so try/except has to be used in such a case:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
def raise_exception(msg): raise ValueError(msg)
|
||||
def test_something_exception():
|
||||
msg = "Not a good value"
|
||||
error = ''
|
||||
try:
|
||||
raise_exception(msg)
|
||||
except Exception as e:
|
||||
error = str(e)
|
||||
assert msg in error, f"{msg} is in the exception:\n{error}"
|
||||
|
||||
Another approach to capturing stdout is via ``contextlib.redirect_stdout``:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
from io import StringIO
|
||||
from contextlib import redirect_stdout
|
||||
def print_to_stdout(s): print(s)
|
||||
def test_result_and_stdout():
|
||||
msg = "Hello"
|
||||
buffer = StringIO()
|
||||
with redirect_stdout(buffer):
|
||||
print_to_stdout(msg)
|
||||
out = buffer.getvalue()
|
||||
# optional: if you want to replay the consumed streams:
|
||||
sys.stdout.write(out)
|
||||
# test:
|
||||
assert msg in out
|
||||
|
||||
An important potential issue with capturing stdout is that it may
|
||||
contain ``\r`` characters that in normal ``print`` reset everything that
|
||||
has been printed so far. There is no problem with ``pytest``, but with
|
||||
``pytest -s`` these characters get included in the buffer, so to be able
|
||||
to have the test run with and without ``-s``, you have to make an extra
|
||||
cleanup to the captured output, using ``re.sub(r'~.*\r', '', buf, 0, re.M)``.
|
||||
|
||||
But, then we have a helper context manager wrapper to automatically take
|
||||
care of it all, regardless of whether it has some ``\r``'s in it or
|
||||
not, so it's a simple:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
from transformers.testing_utils import CaptureStdout
|
||||
with CaptureStdout() as cs:
|
||||
function_that_writes_to_stdout()
|
||||
print(cs.out)
|
||||
|
||||
Here is a full test example:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
from transformers.testing_utils import CaptureStdout
|
||||
msg = "Secret message\r"
|
||||
final = "Hello World"
|
||||
with CaptureStdout() as cs:
|
||||
print(msg + final)
|
||||
assert cs.out == final+"\n", f"captured: {cs.out}, expecting {final}"
|
||||
|
||||
If you'd like to capture ``stderr`` use the :obj:`CaptureStderr` class
|
||||
instead:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
from transformers.testing_utils import CaptureStderr
|
||||
with CaptureStderr() as cs:
|
||||
function_that_writes_to_stderr()
|
||||
print(cs.err)
|
||||
|
||||
If you need to capture both streams at once, use the parent
|
||||
:obj:`CaptureStd` class:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
from transformers.testing_utils import CaptureStd
|
||||
with CaptureStd() as cs:
|
||||
function_that_writes_to_stdout_and_stderr()
|
||||
print(cs.err, cs.out)
|
||||
|
||||
|
||||
|
||||
Capturing logger stream
|
||||
~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
If you need to validate the output of a logger, you can use :obj:`CaptureLogger`:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
from transformers import logging
|
||||
from transformers.testing_utils import CaptureLogger
|
||||
|
||||
msg = "Testing 1, 2, 3"
|
||||
logging.set_verbosity_info()
|
||||
logger = logging.get_logger("transformers.tokenization_bart")
|
||||
with CaptureLogger(logger) as cl:
|
||||
logger.info(msg)
|
||||
assert cl.out, msg+"\n"
|
||||
|
||||
|
||||
Testing with environment variables
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
If you want to test the impact of environment variables for a specific test you can use a helper decorator ``transformers.testing_utils.mockenv``
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
from transformers.testing_utils import mockenv
|
||||
class HfArgumentParserTest(unittest.TestCase):
|
||||
@mockenv(TRANSFORMERS_VERBOSITY="error")
|
||||
def test_env_override(self):
|
||||
env_level_str = os.getenv("TRANSFORMERS_VERBOSITY", None)
|
||||
|
||||
|
||||
Getting reproducible results
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
|
||||
In some situations you may want to remove randomness for your tests. To
|
||||
get identical reproducable results set, you will need to fix the seed:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
seed = 42
|
||||
|
||||
# python RNG
|
||||
import random
|
||||
random.seed(seed)
|
||||
|
||||
# pytorch RNGs
|
||||
import torch
|
||||
torch.manual_seed(seed)
|
||||
torch.backends.cudnn.deterministic = True
|
||||
if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
|
||||
|
||||
# numpy RNG
|
||||
import numpy as np
|
||||
np.random.seed(seed)
|
||||
|
||||
# tf RNG
|
||||
tf.random.set_seed(seed)
|
||||
|
||||
Debugging tests
|
||||
~~~~~~~~~~~~~~~
|
||||
|
||||
To start a debugger at the point of the warning, do this:
|
||||
|
||||
.. code-block:: bash
|
||||
|
||||
pytest tests/test_logging.py -W error::UserWarning --pdb
|
||||
@@ -721,7 +721,7 @@ def main():
|
||||
checkpoints = list(
|
||||
os.path.dirname(c) for c in sorted(glob.glob(args.output_dir + "/**/" + WEIGHTS_NAME, recursive=True))
|
||||
)
|
||||
|
||||
logging.getLogger("transformers.modeling_utils").setLevel(logging.WARN) # Reduce logging
|
||||
logger.info("Evaluate the following checkpoints: %s", checkpoints)
|
||||
|
||||
for checkpoint in checkpoints:
|
||||
|
||||
@@ -547,7 +547,7 @@ def main():
|
||||
checkpoints = list(
|
||||
os.path.dirname(c) for c in sorted(glob.glob(args.output_dir + "/**/" + WEIGHTS_NAME, recursive=True))
|
||||
)
|
||||
|
||||
logging.getLogger("transformers.modeling_utils").setLevel(logging.WARN) # Reduce logging
|
||||
logger.info("Evaluate the following checkpoints: %s", checkpoints)
|
||||
for checkpoint in checkpoints:
|
||||
global_step = checkpoint.split("-")[-1] if len(checkpoints) > 1 else ""
|
||||
|
||||
@@ -681,6 +681,7 @@ def main():
|
||||
checkpoints = list(
|
||||
os.path.dirname(c) for c in sorted(glob.glob(args.output_dir + "/**/" + WEIGHTS_NAME, recursive=True))
|
||||
)
|
||||
logging.getLogger("transformers.modeling_utils").setLevel(logging.WARN) # Reduce model loading logs
|
||||
|
||||
logger.info("Evaluate the following checkpoints: %s", checkpoints)
|
||||
|
||||
|
||||
@@ -88,7 +88,7 @@ def main():
|
||||
)
|
||||
)
|
||||
|
||||
model.reset_memory_length(args.mem_len)
|
||||
model.reset_length(args.tgt_len, args.ext_len, args.mem_len)
|
||||
if args.clamp_len > 0:
|
||||
model.clamp_len = args.clamp_len
|
||||
if args.same_length:
|
||||
|
||||
@@ -677,7 +677,7 @@ def main():
|
||||
checkpoints = list(
|
||||
os.path.dirname(c) for c in sorted(glob.glob(args.output_dir + "/**/" + WEIGHTS_NAME, recursive=True))
|
||||
)
|
||||
|
||||
logging.getLogger("transformers.modeling_utils").setLevel(logging.WARN) # Reduce logging
|
||||
logger.info("Evaluate the following checkpoints: %s", checkpoints)
|
||||
for checkpoint in checkpoints:
|
||||
global_step = checkpoint.split("-")[-1] if len(checkpoints) > 1 else ""
|
||||
|
||||
@@ -842,6 +842,7 @@ def main():
|
||||
checkpoints = list(
|
||||
os.path.dirname(c) for c in sorted(glob.glob(args.output_dir + "/**/" + WEIGHTS_NAME, recursive=True))
|
||||
)
|
||||
logging.getLogger("transformers.modeling_utils").setLevel(logging.WARN) # Reduce model loading logs
|
||||
|
||||
logger.info("Evaluate the following checkpoints: %s", checkpoints)
|
||||
|
||||
|
||||
@@ -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 +0,0 @@
|
||||
# LXMERT DEMO
|
||||
|
||||
1. make a virtualenv: ``virtualenv venv`` and activate ``source venv/bin/activate``
|
||||
2. install reqs: ``pip install -r ./requirements.txt``
|
||||
3. usage is as shown in demo.ipynb
|
||||
File diff suppressed because one or more lines are too long
@@ -1,149 +0,0 @@
|
||||
import getopt
|
||||
import json
|
||||
import os
|
||||
|
||||
# import numpy as np
|
||||
import sys
|
||||
from collections import OrderedDict
|
||||
|
||||
import datasets
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from modeling_frcnn import GeneralizedRCNN
|
||||
from processing_image import Preprocess
|
||||
from utils import Config
|
||||
|
||||
|
||||
"""
|
||||
USAGE:
|
||||
``python extracting_data.py -i <img_dir> -o <dataset_file>.datasets <batch_size>``
|
||||
"""
|
||||
|
||||
|
||||
TEST = False
|
||||
CONFIG = Config.from_pretrained("unc-nlp/frcnn-vg-finetuned")
|
||||
DEFAULT_SCHEMA = datasets.Features(
|
||||
OrderedDict(
|
||||
{
|
||||
"attr_ids": datasets.Sequence(length=CONFIG.MAX_DETECTIONS, feature=datasets.Value("float32")),
|
||||
"attr_probs": datasets.Sequence(length=CONFIG.MAX_DETECTIONS, feature=datasets.Value("float32")),
|
||||
"boxes": datasets.Array2D((CONFIG.MAX_DETECTIONS, 4), dtype="float32"),
|
||||
"img_id": datasets.Value("int32"),
|
||||
"obj_ids": datasets.Sequence(length=CONFIG.MAX_DETECTIONS, feature=datasets.Value("float32")),
|
||||
"obj_probs": datasets.Sequence(length=CONFIG.MAX_DETECTIONS, feature=datasets.Value("float32")),
|
||||
"roi_features": datasets.Array2D((CONFIG.MAX_DETECTIONS, 2048), dtype="float32"),
|
||||
"sizes": datasets.Sequence(length=2, feature=datasets.Value("float32")),
|
||||
"preds_per_image": datasets.Value(dtype="int32"),
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
class Extract:
|
||||
def __init__(self, argv=sys.argv[1:]):
|
||||
inputdir = None
|
||||
outputfile = None
|
||||
subset_list = None
|
||||
batch_size = 1
|
||||
opts, args = getopt.getopt(argv, "i:o:b:s", ["inputdir=", "outfile=", "batch_size=", "subset_list="])
|
||||
for opt, arg in opts:
|
||||
if opt in ("-i", "--inputdir"):
|
||||
inputdir = arg
|
||||
elif opt in ("-o", "--outfile"):
|
||||
outputfile = arg
|
||||
elif opt in ("-b", "--batch_size"):
|
||||
batch_size = int(arg)
|
||||
elif opt in ("-s", "--subset_list"):
|
||||
subset_list = arg
|
||||
|
||||
assert inputdir is not None # and os.path.isdir(inputdir), f"{inputdir}"
|
||||
assert outputfile is not None and not os.path.isfile(outputfile), f"{outputfile}"
|
||||
if subset_list is not None:
|
||||
with open(os.path.realpath(subset_list)) as f:
|
||||
self.subset_list = set(map(lambda x: self._vqa_file_split()[0], tryload(f)))
|
||||
else:
|
||||
self.subset_list = None
|
||||
|
||||
self.config = CONFIG
|
||||
if torch.cuda.is_available():
|
||||
self.config.model.device = "cuda"
|
||||
self.inputdir = os.path.realpath(inputdir)
|
||||
self.outputfile = os.path.realpath(outputfile)
|
||||
self.preprocess = Preprocess(self.config)
|
||||
self.model = GeneralizedRCNN.from_pretrained("unc-nlp/frcnn-vg-finetuned", config=self.config)
|
||||
self.batch = batch_size if batch_size != 0 else 1
|
||||
self.schema = DEFAULT_SCHEMA
|
||||
|
||||
def _vqa_file_split(self, file):
|
||||
img_id = int(file.split(".")[0].split("_")[-1])
|
||||
filepath = os.path.join(self.inputdir, file)
|
||||
return (img_id, filepath)
|
||||
|
||||
@property
|
||||
def file_generator(self):
|
||||
batch = []
|
||||
for i, file in enumerate(os.listdir(self.inputdir)):
|
||||
if self.subset_list is not None and i not in self.subset_list:
|
||||
continue
|
||||
batch.append(self._vqa_file_split(file))
|
||||
if len(batch) == self.batch:
|
||||
temp = batch
|
||||
batch = []
|
||||
yield list(map(list, zip(*temp)))
|
||||
|
||||
for i in range(1):
|
||||
yield list(map(list, zip(*batch)))
|
||||
|
||||
def __call__(self):
|
||||
# make writer
|
||||
if not TEST:
|
||||
writer = datasets.ArrowWriter(features=self.schema, path=self.outputfile)
|
||||
# do file generator
|
||||
for i, (img_ids, filepaths) in enumerate(self.file_generator):
|
||||
images, sizes, scales_yx = self.preprocess(filepaths)
|
||||
output_dict = self.model(
|
||||
images,
|
||||
sizes,
|
||||
scales_yx=scales_yx,
|
||||
padding="max_detections",
|
||||
max_detections=self.config.MAX_DETECTIONS,
|
||||
pad_value=0,
|
||||
return_tensors="np",
|
||||
location="cpu",
|
||||
)
|
||||
output_dict["boxes"] = output_dict.pop("normalized_boxes")
|
||||
if not TEST:
|
||||
output_dict["img_id"] = np.array(img_ids)
|
||||
batch = self.schema.encode_batch(output_dict)
|
||||
writer.write_batch(batch)
|
||||
if TEST:
|
||||
break
|
||||
# finalizer the writer
|
||||
if not TEST:
|
||||
num_examples, num_bytes = writer.finalize()
|
||||
print(f"Success! You wrote {num_examples} entry(s) and {num_bytes >> 20} mb")
|
||||
|
||||
|
||||
def tryload(stream):
|
||||
try:
|
||||
data = json.load(stream)
|
||||
try:
|
||||
data = list(data.keys())
|
||||
except Exception:
|
||||
data = [d["img_id"] for d in data]
|
||||
except Exception:
|
||||
try:
|
||||
data = eval(stream.read())
|
||||
except Exception:
|
||||
data = stream.read().split("\n")
|
||||
return data
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
extract = Extract(sys.argv[1:])
|
||||
extract()
|
||||
if not TEST:
|
||||
dataset = datasets.Dataset.from_file(extract.outputfile)
|
||||
# wala!
|
||||
# print(np.array(dataset[0:2]["roi_features"]).shape)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,147 +0,0 @@
|
||||
"""
|
||||
coding=utf-8
|
||||
Copyright 2018, Antonio Mendoza Hao Tan, Mohit Bansal
|
||||
Adapted From Facebook Inc, Detectron2
|
||||
|
||||
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 sys
|
||||
from typing import Tuple
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from PIL import Image
|
||||
|
||||
from utils import img_tensorize
|
||||
|
||||
|
||||
class ResizeShortestEdge:
|
||||
def __init__(self, short_edge_length, max_size=sys.maxsize):
|
||||
"""
|
||||
Args:
|
||||
short_edge_length (list[min, max])
|
||||
max_size (int): maximum allowed longest edge length.
|
||||
"""
|
||||
self.interp_method = "bilinear"
|
||||
self.max_size = max_size
|
||||
self.short_edge_length = short_edge_length
|
||||
|
||||
def __call__(self, imgs):
|
||||
img_augs = []
|
||||
for img in imgs:
|
||||
h, w = img.shape[:2]
|
||||
# later: provide list and randomly choose index for resize
|
||||
size = np.random.randint(self.short_edge_length[0], self.short_edge_length[1] + 1)
|
||||
if size == 0:
|
||||
return img
|
||||
scale = size * 1.0 / min(h, w)
|
||||
if h < w:
|
||||
newh, neww = size, scale * w
|
||||
else:
|
||||
newh, neww = scale * h, size
|
||||
if max(newh, neww) > self.max_size:
|
||||
scale = self.max_size * 1.0 / max(newh, neww)
|
||||
newh = newh * scale
|
||||
neww = neww * scale
|
||||
neww = int(neww + 0.5)
|
||||
newh = int(newh + 0.5)
|
||||
|
||||
if img.dtype == np.uint8:
|
||||
pil_image = Image.fromarray(img)
|
||||
pil_image = pil_image.resize((neww, newh), Image.BILINEAR)
|
||||
img = np.asarray(pil_image)
|
||||
else:
|
||||
img = img.permute(2, 0, 1).unsqueeze(0) # 3, 0, 1) # hw(c) -> nchw
|
||||
img = F.interpolate(img, (newh, neww), mode=self.interp_method, align_corners=False).squeeze(0)
|
||||
img_augs.append(img)
|
||||
|
||||
return img_augs
|
||||
|
||||
|
||||
class Preprocess:
|
||||
def __init__(self, cfg):
|
||||
self.aug = ResizeShortestEdge([cfg.INPUT.MIN_SIZE_TEST, cfg.INPUT.MIN_SIZE_TEST], cfg.INPUT.MAX_SIZE_TEST)
|
||||
self.input_format = cfg.INPUT.FORMAT
|
||||
self.size_divisibility = cfg.SIZE_DIVISIBILITY
|
||||
self.pad_value = cfg.PAD_VALUE
|
||||
self.max_image_size = cfg.INPUT.MAX_SIZE_TEST
|
||||
self.device = cfg.MODEL.DEVICE
|
||||
self.pixel_std = torch.tensor(cfg.MODEL.PIXEL_STD).to(self.device).view(len(cfg.MODEL.PIXEL_STD), 1, 1)
|
||||
self.pixel_mean = torch.tensor(cfg.MODEL.PIXEL_MEAN).to(self.device).view(len(cfg.MODEL.PIXEL_STD), 1, 1)
|
||||
self.normalizer = lambda x: (x - self.pixel_mean) / self.pixel_std
|
||||
|
||||
def pad(self, images):
|
||||
max_size = tuple(max(s) for s in zip(*[img.shape for img in images]))
|
||||
image_sizes = [im.shape[-2:] for im in images]
|
||||
images = [
|
||||
F.pad(
|
||||
im,
|
||||
[0, max_size[-1] - size[1], 0, max_size[-2] - size[0]],
|
||||
value=self.pad_value,
|
||||
)
|
||||
for size, im in zip(image_sizes, images)
|
||||
]
|
||||
|
||||
return torch.stack(images), torch.tensor(image_sizes)
|
||||
|
||||
def __call__(self, images, single_image=False):
|
||||
with torch.no_grad():
|
||||
if not isinstance(images, list):
|
||||
images = [images]
|
||||
if single_image:
|
||||
assert len(images) == 1
|
||||
for i in range(len(images)):
|
||||
if isinstance(images[i], torch.Tensor):
|
||||
images.insert(i, images.pop(i).to(self.device).float())
|
||||
elif not isinstance(images[i], torch.Tensor):
|
||||
images.insert(
|
||||
i,
|
||||
torch.as_tensor(img_tensorize(images.pop(i), input_format=self.input_format))
|
||||
.to(self.device)
|
||||
.float(),
|
||||
)
|
||||
# resize smallest edge
|
||||
raw_sizes = torch.tensor([im.shape[:2] for im in images])
|
||||
images = self.aug(images)
|
||||
# transpose images and convert to torch tensors
|
||||
# images = [torch.as_tensor(i.astype("float32")).permute(2, 0, 1).to(self.device) for i in images]
|
||||
# now normalize before pad to aoid useless arithmatic
|
||||
images = [self.normalizer(x) for x in images]
|
||||
# now pad them to do the following operations
|
||||
images, sizes = self.pad(images)
|
||||
# Normalize
|
||||
|
||||
if self.size_divisibility > 0:
|
||||
raise NotImplementedError()
|
||||
# pad
|
||||
scales_yx = torch.true_divide(raw_sizes, sizes)
|
||||
if single_image:
|
||||
return images[0], sizes[0], scales_yx[0]
|
||||
else:
|
||||
return images, sizes, scales_yx
|
||||
|
||||
|
||||
def _scale_box(boxes, scale_yx):
|
||||
boxes[:, 0::2] *= scale_yx[:, 1]
|
||||
boxes[:, 1::2] *= scale_yx[:, 0]
|
||||
return boxes
|
||||
|
||||
|
||||
def _clip_box(tensor, box_size: Tuple[int, int]):
|
||||
assert torch.isfinite(tensor).all(), "Box tensor contains infinite or NaN!"
|
||||
h, w = box_size
|
||||
tensor[:, 0].clamp_(min=0, max=w)
|
||||
tensor[:, 1].clamp_(min=0, max=h)
|
||||
tensor[:, 2].clamp_(min=0, max=w)
|
||||
tensor[:, 3].clamp_(min=0, max=h)
|
||||
@@ -1,99 +0,0 @@
|
||||
appdirs==1.4.3
|
||||
argon2-cffi==20.1.0
|
||||
async-generator==1.10
|
||||
attrs==20.2.0
|
||||
backcall==0.2.0
|
||||
bleach==3.1.5
|
||||
CacheControl==0.12.6
|
||||
certifi==2020.6.20
|
||||
cffi==1.14.2
|
||||
chardet==3.0.4
|
||||
click==7.1.2
|
||||
colorama==0.4.3
|
||||
contextlib2==0.6.0
|
||||
cycler==0.10.0
|
||||
datasets==1.0.0
|
||||
decorator==4.4.2
|
||||
defusedxml==0.6.0
|
||||
dill==0.3.2
|
||||
distlib==0.3.0
|
||||
distro==1.4.0
|
||||
entrypoints==0.3
|
||||
filelock==3.0.12
|
||||
future==0.18.2
|
||||
html5lib==1.0.1
|
||||
idna==2.8
|
||||
ipaddr==2.2.0
|
||||
ipykernel==5.3.4
|
||||
ipython
|
||||
ipython-genutils==0.2.0
|
||||
ipywidgets==7.5.1
|
||||
jedi==0.17.2
|
||||
Jinja2==2.11.2
|
||||
joblib==0.16.0
|
||||
jsonschema==3.2.0
|
||||
jupyter==1.0.0
|
||||
jupyter-client==6.1.7
|
||||
jupyter-console==6.2.0
|
||||
jupyter-core==4.6.3
|
||||
jupyterlab-pygments==0.1.1
|
||||
kiwisolver==1.2.0
|
||||
lockfile==0.12.2
|
||||
MarkupSafe==1.1.1
|
||||
matplotlib==3.3.1
|
||||
mistune==0.8.4
|
||||
msgpack==0.6.2
|
||||
nbclient==0.5.0
|
||||
nbconvert==6.0.1
|
||||
nbformat==5.0.7
|
||||
nest-asyncio==1.4.0
|
||||
notebook==6.1.4
|
||||
numpy==1.19.2
|
||||
opencv-python==4.4.0.42
|
||||
packaging==20.3
|
||||
pandas==1.1.2
|
||||
pandocfilters==1.4.2
|
||||
parso==0.7.1
|
||||
pep517==0.8.2
|
||||
pexpect==4.8.0
|
||||
pickleshare==0.7.5
|
||||
Pillow==7.2.0
|
||||
progress==1.5
|
||||
prometheus-client==0.8.0
|
||||
prompt-toolkit==3.0.7
|
||||
ptyprocess==0.6.0
|
||||
pyaml==20.4.0
|
||||
pyarrow==1.0.1
|
||||
pycparser==2.20
|
||||
Pygments==2.6.1
|
||||
pyparsing==2.4.6
|
||||
pyrsistent==0.16.0
|
||||
python-dateutil==2.8.1
|
||||
pytoml==0.1.21
|
||||
pytz==2020.1
|
||||
PyYAML==5.3.1
|
||||
pyzmq==19.0.2
|
||||
qtconsole==4.7.7
|
||||
QtPy==1.9.0
|
||||
regex==2020.7.14
|
||||
requests==2.22.0
|
||||
retrying==1.3.3
|
||||
sacremoses==0.0.43
|
||||
Send2Trash==1.5.0
|
||||
sentencepiece==0.1.91
|
||||
six==1.14.0
|
||||
terminado==0.8.3
|
||||
testpath==0.4.4
|
||||
tokenizers==0.8.1rc2
|
||||
torch==1.6.0
|
||||
torchvision==0.7.0
|
||||
tornado==6.0.4
|
||||
tqdm==4.48.2
|
||||
traitlets
|
||||
git+https://github.com/huggingface/transformers.git
|
||||
urllib3==1.25.8
|
||||
wcwidth==0.2.5
|
||||
webencodings==0.5.1
|
||||
wget==3.2
|
||||
widgetsnbextension==3.5.1
|
||||
xxhash==2.0.0
|
||||
@@ -1,559 +0,0 @@
|
||||
"""
|
||||
coding=utf-8
|
||||
Copyright 2018, Antonio Mendoza Hao Tan, Mohit Bansal, Huggingface team :)
|
||||
Adapted From Facebook Inc, Detectron2
|
||||
|
||||
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 copy
|
||||
import fnmatch
|
||||
import json
|
||||
import os
|
||||
import pickle as pkl
|
||||
import shutil
|
||||
import sys
|
||||
import tarfile
|
||||
import tempfile
|
||||
from collections import OrderedDict
|
||||
from contextlib import contextmanager
|
||||
from functools import partial
|
||||
from hashlib import sha256
|
||||
from io import BytesIO
|
||||
from pathlib import Path
|
||||
from urllib.parse import urlparse
|
||||
from zipfile import ZipFile, is_zipfile
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
import cv2
|
||||
import requests
|
||||
import wget
|
||||
from filelock import FileLock
|
||||
from yaml import Loader, dump, load
|
||||
|
||||
|
||||
try:
|
||||
import torch
|
||||
|
||||
_torch_available = True
|
||||
except ImportError:
|
||||
_torch_available = False
|
||||
|
||||
|
||||
try:
|
||||
from torch.hub import _get_torch_home
|
||||
|
||||
torch_cache_home = _get_torch_home()
|
||||
except ImportError:
|
||||
torch_cache_home = os.path.expanduser(
|
||||
os.getenv("TORCH_HOME", os.path.join(os.getenv("XDG_CACHE_HOME", "~/.cache"), "torch"))
|
||||
)
|
||||
|
||||
default_cache_path = os.path.join(torch_cache_home, "transformers")
|
||||
|
||||
CLOUDFRONT_DISTRIB_PREFIX = "https://cdn.huggingface.co"
|
||||
S3_BUCKET_PREFIX = "https://s3.amazonaws.com/models.huggingface.co/bert"
|
||||
PATH = "/".join(str(Path(__file__).resolve()).split("/")[:-1])
|
||||
CONFIG = os.path.join(PATH, "config.yaml")
|
||||
ATTRIBUTES = os.path.join(PATH, "attributes.txt")
|
||||
OBJECTS = os.path.join(PATH, "objects.txt")
|
||||
PYTORCH_PRETRAINED_BERT_CACHE = os.getenv("PYTORCH_PRETRAINED_BERT_CACHE", default_cache_path)
|
||||
PYTORCH_TRANSFORMERS_CACHE = os.getenv("PYTORCH_TRANSFORMERS_CACHE", PYTORCH_PRETRAINED_BERT_CACHE)
|
||||
TRANSFORMERS_CACHE = os.getenv("TRANSFORMERS_CACHE", PYTORCH_TRANSFORMERS_CACHE)
|
||||
WEIGHTS_NAME = "pytorch_model.bin"
|
||||
CONFIG_NAME = "config.yaml"
|
||||
|
||||
|
||||
def load_labels(objs=OBJECTS, attrs=ATTRIBUTES):
|
||||
vg_classes = []
|
||||
with open(objs) as f:
|
||||
for object in f.readlines():
|
||||
vg_classes.append(object.split(",")[0].lower().strip())
|
||||
|
||||
vg_attrs = []
|
||||
with open(attrs) as f:
|
||||
for object in f.readlines():
|
||||
vg_attrs.append(object.split(",")[0].lower().strip())
|
||||
return vg_classes, vg_attrs
|
||||
|
||||
|
||||
def load_checkpoint(ckp):
|
||||
r = OrderedDict()
|
||||
with open(ckp, "rb") as f:
|
||||
ckp = pkl.load(f)["model"]
|
||||
for k in copy.deepcopy(list(ckp.keys())):
|
||||
v = ckp.pop(k)
|
||||
if isinstance(v, np.ndarray):
|
||||
v = torch.tensor(v)
|
||||
else:
|
||||
assert isinstance(v, torch.tensor), type(v)
|
||||
r[k] = v
|
||||
return r
|
||||
|
||||
|
||||
class Config:
|
||||
_pointer = {}
|
||||
|
||||
def __init__(self, dictionary: dict, name: str = "root", level=0):
|
||||
self._name = name
|
||||
self._level = level
|
||||
d = {}
|
||||
for k, v in dictionary.items():
|
||||
if v is None:
|
||||
raise ValueError()
|
||||
k = copy.deepcopy(k)
|
||||
v = copy.deepcopy(v)
|
||||
if isinstance(v, dict):
|
||||
v = Config(v, name=k, level=level + 1)
|
||||
d[k] = v
|
||||
setattr(self, k, v)
|
||||
|
||||
self._pointer = d
|
||||
|
||||
def __repr__(self):
|
||||
return str(list((self._pointer.keys())))
|
||||
|
||||
def __setattr__(self, key, val):
|
||||
self.__dict__[key] = val
|
||||
self.__dict__[key.upper()] = val
|
||||
levels = key.split(".")
|
||||
last_level = len(levels) - 1
|
||||
pointer = self._pointer
|
||||
if len(levels) > 1:
|
||||
for i, l in enumerate(levels):
|
||||
if hasattr(self, l) and isinstance(getattr(self, l), Config):
|
||||
setattr(getattr(self, l), ".".join(levels[i:]), val)
|
||||
if l == last_level:
|
||||
pointer[l] = val
|
||||
else:
|
||||
pointer = pointer[l]
|
||||
|
||||
def to_dict(self):
|
||||
return self._pointer
|
||||
|
||||
def dump_yaml(self, data, file_name):
|
||||
with open(f"{file_name}", "w") as stream:
|
||||
dump(data, stream)
|
||||
|
||||
def dump_json(self, data, file_name):
|
||||
with open(f"{file_name}", "w") as stream:
|
||||
json.dump(data, stream)
|
||||
|
||||
@staticmethod
|
||||
def load_yaml(config):
|
||||
with open(config) as stream:
|
||||
data = load(stream, Loader=Loader)
|
||||
return data
|
||||
|
||||
def __str__(self):
|
||||
t = " "
|
||||
if self._name != "root":
|
||||
r = f"{t * (self._level-1)}{self._name}:\n"
|
||||
else:
|
||||
r = ""
|
||||
level = self._level
|
||||
for i, (k, v) in enumerate(self._pointer.items()):
|
||||
if isinstance(v, Config):
|
||||
r += f"{t * (self._level)}{v}\n"
|
||||
self._level += 1
|
||||
else:
|
||||
r += f"{t * (self._level)}{k}: {v} ({type(v).__name__})\n"
|
||||
self._level = level
|
||||
return r[:-1]
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, pretrained_model_name_or_path: str, **kwargs):
|
||||
config_dict, kwargs = cls.get_config_dict(pretrained_model_name_or_path, **kwargs)
|
||||
return cls(config_dict)
|
||||
|
||||
@classmethod
|
||||
def get_config_dict(cls, pretrained_model_name_or_path: str, **kwargs):
|
||||
|
||||
cache_dir = kwargs.pop("cache_dir", None)
|
||||
force_download = kwargs.pop("force_download", False)
|
||||
resume_download = kwargs.pop("resume_download", False)
|
||||
proxies = kwargs.pop("proxies", None)
|
||||
local_files_only = kwargs.pop("local_files_only", False)
|
||||
|
||||
if os.path.isdir(pretrained_model_name_or_path):
|
||||
config_file = os.path.join(pretrained_model_name_or_path, CONFIG_NAME)
|
||||
elif os.path.isfile(pretrained_model_name_or_path) or is_remote_url(pretrained_model_name_or_path):
|
||||
config_file = pretrained_model_name_or_path
|
||||
else:
|
||||
config_file = hf_bucket_url(pretrained_model_name_or_path, filename=CONFIG_NAME, use_cdn=False)
|
||||
|
||||
try:
|
||||
# Load from URL or cache if already cached
|
||||
resolved_config_file = cached_path(
|
||||
config_file,
|
||||
cache_dir=cache_dir,
|
||||
force_download=force_download,
|
||||
proxies=proxies,
|
||||
resume_download=resume_download,
|
||||
local_files_only=local_files_only,
|
||||
)
|
||||
# Load config dict
|
||||
if resolved_config_file is None:
|
||||
raise EnvironmentError
|
||||
|
||||
config_file = Config.load_yaml(resolved_config_file)
|
||||
|
||||
except EnvironmentError:
|
||||
msg = "Can't load config for"
|
||||
raise EnvironmentError(msg)
|
||||
|
||||
if resolved_config_file == config_file:
|
||||
print("loading configuration file from path")
|
||||
else:
|
||||
print("loading configuration file cache")
|
||||
|
||||
return Config.load_yaml(resolved_config_file), kwargs
|
||||
|
||||
|
||||
# quick compare tensors
|
||||
def compare(in_tensor):
|
||||
|
||||
out_tensor = torch.load("dump.pt", map_location=in_tensor.device)
|
||||
n1 = in_tensor.numpy()
|
||||
n2 = out_tensor.numpy()[0]
|
||||
print(n1.shape, n1[0, 0, :5])
|
||||
print(n2.shape, n2[0, 0, :5])
|
||||
assert np.allclose(
|
||||
n1, n2, rtol=0.01, atol=0.1
|
||||
), f"{sum([1 for x in np.isclose(n1, n2, rtol=0.01, atol=0.1).flatten() if x == False])/len(n1.flatten())*100:.4f} % element-wise mismatch"
|
||||
raise Exception("tensors are all good")
|
||||
|
||||
# Hugging face functiions below
|
||||
|
||||
|
||||
def is_remote_url(url_or_filename):
|
||||
parsed = urlparse(url_or_filename)
|
||||
return parsed.scheme in ("http", "https")
|
||||
|
||||
|
||||
def hf_bucket_url(model_id: str, filename: str, use_cdn=True) -> str:
|
||||
endpoint = CLOUDFRONT_DISTRIB_PREFIX if use_cdn else S3_BUCKET_PREFIX
|
||||
legacy_format = "/" not in model_id
|
||||
if legacy_format:
|
||||
return f"{endpoint}/{model_id}-{filename}"
|
||||
else:
|
||||
return f"{endpoint}/{model_id}/{filename}"
|
||||
|
||||
|
||||
def http_get(
|
||||
url,
|
||||
temp_file,
|
||||
proxies=None,
|
||||
resume_size=0,
|
||||
user_agent=None,
|
||||
):
|
||||
ua = "python/{}".format(sys.version.split()[0])
|
||||
if _torch_available:
|
||||
ua += "; torch/{}".format(torch.__version__)
|
||||
if isinstance(user_agent, dict):
|
||||
ua += "; " + "; ".join("{}/{}".format(k, v) for k, v in user_agent.items())
|
||||
elif isinstance(user_agent, str):
|
||||
ua += "; " + user_agent
|
||||
headers = {"user-agent": ua}
|
||||
if resume_size > 0:
|
||||
headers["Range"] = "bytes=%d-" % (resume_size,)
|
||||
response = requests.get(url, stream=True, proxies=proxies, headers=headers)
|
||||
if response.status_code == 416: # Range not satisfiable
|
||||
return
|
||||
content_length = response.headers.get("Content-Length")
|
||||
total = resume_size + int(content_length) if content_length is not None else None
|
||||
progress = tqdm(
|
||||
unit="B",
|
||||
unit_scale=True,
|
||||
total=total,
|
||||
initial=resume_size,
|
||||
desc="Downloading",
|
||||
)
|
||||
for chunk in response.iter_content(chunk_size=1024):
|
||||
if chunk: # filter out keep-alive new chunks
|
||||
progress.update(len(chunk))
|
||||
temp_file.write(chunk)
|
||||
progress.close()
|
||||
|
||||
|
||||
def get_from_cache(
|
||||
url,
|
||||
cache_dir=None,
|
||||
force_download=False,
|
||||
proxies=None,
|
||||
etag_timeout=10,
|
||||
resume_download=False,
|
||||
user_agent=None,
|
||||
local_files_only=False,
|
||||
):
|
||||
|
||||
if cache_dir is None:
|
||||
cache_dir = TRANSFORMERS_CACHE
|
||||
if isinstance(cache_dir, Path):
|
||||
cache_dir = str(cache_dir)
|
||||
|
||||
os.makedirs(cache_dir, exist_ok=True)
|
||||
|
||||
etag = None
|
||||
if not local_files_only:
|
||||
try:
|
||||
response = requests.head(url, allow_redirects=True, proxies=proxies, timeout=etag_timeout)
|
||||
if response.status_code == 200:
|
||||
etag = response.headers.get("ETag")
|
||||
except (EnvironmentError, requests.exceptions.Timeout):
|
||||
# etag is already None
|
||||
pass
|
||||
|
||||
filename = url_to_filename(url, etag)
|
||||
|
||||
# get cache path to put the file
|
||||
cache_path = os.path.join(cache_dir, filename)
|
||||
|
||||
# etag is None = we don't have a connection, or url doesn't exist, or is otherwise inaccessible.
|
||||
# try to get the last downloaded one
|
||||
if etag is None:
|
||||
if os.path.exists(cache_path):
|
||||
return cache_path
|
||||
else:
|
||||
matching_files = [
|
||||
file
|
||||
for file in fnmatch.filter(os.listdir(cache_dir), filename + ".*")
|
||||
if not file.endswith(".json") and not file.endswith(".lock")
|
||||
]
|
||||
if len(matching_files) > 0:
|
||||
return os.path.join(cache_dir, matching_files[-1])
|
||||
else:
|
||||
# If files cannot be found and local_files_only=True,
|
||||
# the models might've been found if local_files_only=False
|
||||
# Notify the user about that
|
||||
if local_files_only:
|
||||
raise ValueError(
|
||||
"Cannot find the requested files in the cached path and outgoing traffic has been"
|
||||
" disabled. To enable model look-ups and downloads online, set 'local_files_only'"
|
||||
" to False."
|
||||
)
|
||||
return None
|
||||
|
||||
# From now on, etag is not None.
|
||||
if os.path.exists(cache_path) and not force_download:
|
||||
return cache_path
|
||||
|
||||
# Prevent parallel downloads of the same file with a lock.
|
||||
lock_path = cache_path + ".lock"
|
||||
with FileLock(lock_path):
|
||||
|
||||
# If the download just completed while the lock was activated.
|
||||
if os.path.exists(cache_path) and not force_download:
|
||||
# Even if returning early like here, the lock will be released.
|
||||
return cache_path
|
||||
|
||||
if resume_download:
|
||||
incomplete_path = cache_path + ".incomplete"
|
||||
|
||||
@contextmanager
|
||||
def _resumable_file_manager():
|
||||
with open(incomplete_path, "a+b") as f:
|
||||
yield f
|
||||
|
||||
temp_file_manager = _resumable_file_manager
|
||||
if os.path.exists(incomplete_path):
|
||||
resume_size = os.stat(incomplete_path).st_size
|
||||
else:
|
||||
resume_size = 0
|
||||
else:
|
||||
temp_file_manager = partial(tempfile.NamedTemporaryFile, dir=cache_dir, delete=False)
|
||||
resume_size = 0
|
||||
|
||||
# Download to temporary file, then copy to cache dir once finished.
|
||||
# Otherwise you get corrupt cache entries if the download gets interrupted.
|
||||
with temp_file_manager() as temp_file:
|
||||
print(
|
||||
"%s not found in cache or force_download set to True, downloading to %s",
|
||||
url,
|
||||
temp_file.name,
|
||||
)
|
||||
|
||||
http_get(
|
||||
url,
|
||||
temp_file,
|
||||
proxies=proxies,
|
||||
resume_size=resume_size,
|
||||
user_agent=user_agent,
|
||||
)
|
||||
|
||||
os.replace(temp_file.name, cache_path)
|
||||
|
||||
meta = {"url": url, "etag": etag}
|
||||
meta_path = cache_path + ".json"
|
||||
with open(meta_path, "w") as meta_file:
|
||||
json.dump(meta, meta_file)
|
||||
|
||||
return cache_path
|
||||
|
||||
|
||||
def url_to_filename(url, etag=None):
|
||||
|
||||
url_bytes = url.encode("utf-8")
|
||||
url_hash = sha256(url_bytes)
|
||||
filename = url_hash.hexdigest()
|
||||
|
||||
if etag:
|
||||
etag_bytes = etag.encode("utf-8")
|
||||
etag_hash = sha256(etag_bytes)
|
||||
filename += "." + etag_hash.hexdigest()
|
||||
|
||||
if url.endswith(".h5"):
|
||||
filename += ".h5"
|
||||
|
||||
return filename
|
||||
|
||||
|
||||
def cached_path(
|
||||
url_or_filename,
|
||||
cache_dir=None,
|
||||
force_download=False,
|
||||
proxies=None,
|
||||
resume_download=False,
|
||||
user_agent=None,
|
||||
extract_compressed_file=False,
|
||||
force_extract=False,
|
||||
local_files_only=False,
|
||||
):
|
||||
if cache_dir is None:
|
||||
cache_dir = TRANSFORMERS_CACHE
|
||||
if isinstance(url_or_filename, Path):
|
||||
url_or_filename = str(url_or_filename)
|
||||
if isinstance(cache_dir, Path):
|
||||
cache_dir = str(cache_dir)
|
||||
|
||||
if is_remote_url(url_or_filename):
|
||||
# URL, so get it from the cache (downloading if necessary)
|
||||
output_path = get_from_cache(
|
||||
url_or_filename,
|
||||
cache_dir=cache_dir,
|
||||
force_download=force_download,
|
||||
proxies=proxies,
|
||||
resume_download=resume_download,
|
||||
user_agent=user_agent,
|
||||
local_files_only=local_files_only,
|
||||
)
|
||||
elif os.path.exists(url_or_filename):
|
||||
# File, and it exists.
|
||||
output_path = url_or_filename
|
||||
elif urlparse(url_or_filename).scheme == "":
|
||||
# File, but it doesn't exist.
|
||||
raise EnvironmentError("file {} not found".format(url_or_filename))
|
||||
else:
|
||||
# Something unknown
|
||||
raise ValueError("unable to parse {} as a URL or as a local path".format(url_or_filename))
|
||||
|
||||
if extract_compressed_file:
|
||||
if not is_zipfile(output_path) and not tarfile.is_tarfile(output_path):
|
||||
return output_path
|
||||
|
||||
# Path where we extract compressed archives
|
||||
# We avoid '.' in dir name and add "-extracted" at the end: "./model.zip" => "./model-zip-extracted/"
|
||||
output_dir, output_file = os.path.split(output_path)
|
||||
output_extract_dir_name = output_file.replace(".", "-") + "-extracted"
|
||||
output_path_extracted = os.path.join(output_dir, output_extract_dir_name)
|
||||
|
||||
if os.path.isdir(output_path_extracted) and os.listdir(output_path_extracted) and not force_extract:
|
||||
return output_path_extracted
|
||||
|
||||
# Prevent parallel extractions
|
||||
lock_path = output_path + ".lock"
|
||||
with FileLock(lock_path):
|
||||
shutil.rmtree(output_path_extracted, ignore_errors=True)
|
||||
os.makedirs(output_path_extracted)
|
||||
if is_zipfile(output_path):
|
||||
with ZipFile(output_path, "r") as zip_file:
|
||||
zip_file.extractall(output_path_extracted)
|
||||
zip_file.close()
|
||||
elif tarfile.is_tarfile(output_path):
|
||||
tar_file = tarfile.open(output_path)
|
||||
tar_file.extractall(output_path_extracted)
|
||||
tar_file.close()
|
||||
else:
|
||||
raise EnvironmentError("Archive format of {} could not be identified".format(output_path))
|
||||
|
||||
return output_path_extracted
|
||||
|
||||
return output_path
|
||||
|
||||
|
||||
def get_data(query, delim=","):
|
||||
assert isinstance(query, str)
|
||||
if os.path.isfile(query):
|
||||
with open(query) as f:
|
||||
data = eval(f.read())
|
||||
else:
|
||||
req = requests.get(query)
|
||||
try:
|
||||
data = requests.json()
|
||||
except Exception:
|
||||
data = req.content.decode()
|
||||
assert data is not None, "could not connect"
|
||||
try:
|
||||
data = eval(data)
|
||||
except Exception:
|
||||
data = data.split("\n")
|
||||
req.close()
|
||||
return data
|
||||
|
||||
|
||||
def get_image_from_url(url):
|
||||
response = requests.get(url)
|
||||
img = np.array(Image.open(BytesIO(response.content)))
|
||||
return img
|
||||
|
||||
|
||||
# to load legace frcnn checkpoint from detectron
|
||||
def load_frcnn_pkl_from_url(url):
|
||||
fn = url.split("/")[-1]
|
||||
if fn not in os.listdir(os.getcwd()):
|
||||
wget.download(url)
|
||||
with open(fn, "rb") as stream:
|
||||
weights = pkl.load(stream)
|
||||
model = weights.pop("model")
|
||||
new = {}
|
||||
for k, v in model.items():
|
||||
new[k] = torch.from_numpy(v)
|
||||
if "running_var" in k:
|
||||
zero = torch.Tensor([0])
|
||||
k2 = k.replace("running_var", "num_batches_tracked")
|
||||
new[k2] = zero
|
||||
return new
|
||||
|
||||
|
||||
def get_demo_path():
|
||||
print(f"{os.path.abspath(os.path.join(PATH, os.pardir))}/demo.ipynb")
|
||||
|
||||
|
||||
def img_tensorize(im, input_format="RGB"):
|
||||
assert isinstance(im, str)
|
||||
if os.path.isfile(im):
|
||||
img = cv2.imread(im)
|
||||
else:
|
||||
img = get_image_from_url(im)
|
||||
assert img is not None, f"could not connect to: {im}"
|
||||
img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
|
||||
if input_format == "RGB":
|
||||
img = img[:, :, ::-1]
|
||||
return img
|
||||
|
||||
|
||||
def chunk(images, batch=1):
|
||||
return (images[i : i + batch] for i in range(0, len(images), batch))
|
||||
@@ -1,499 +0,0 @@
|
||||
"""
|
||||
coding=utf-8
|
||||
Copyright 2018, Antonio Mendoza Hao Tan, Mohit Bansal
|
||||
Adapted From Facebook Inc, Detectron2
|
||||
|
||||
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 colorsys
|
||||
import io
|
||||
|
||||
import matplotlib as mpl
|
||||
import matplotlib.colors as mplc
|
||||
import matplotlib.figure as mplfigure
|
||||
import numpy as np
|
||||
import torch
|
||||
from matplotlib.backends.backend_agg import FigureCanvasAgg
|
||||
|
||||
import cv2
|
||||
from utils import img_tensorize
|
||||
|
||||
|
||||
_SMALL_OBJ = 1000
|
||||
|
||||
|
||||
class SingleImageViz:
|
||||
def __init__(
|
||||
self,
|
||||
img,
|
||||
scale=1.2,
|
||||
edgecolor="g",
|
||||
alpha=0.5,
|
||||
linestyle="-",
|
||||
saveas="test_out.jpg",
|
||||
rgb=True,
|
||||
pynb=False,
|
||||
id2obj=None,
|
||||
id2attr=None,
|
||||
pad=0.7,
|
||||
):
|
||||
"""
|
||||
img: an RGB image of shape (H, W, 3).
|
||||
"""
|
||||
if isinstance(img, torch.Tensor):
|
||||
img = img.numpy().astype("np.uint8")
|
||||
if isinstance(img, str):
|
||||
img = img_tensorize(img)
|
||||
assert isinstance(img, np.ndarray)
|
||||
|
||||
width, height = img.shape[1], img.shape[0]
|
||||
fig = mplfigure.Figure(frameon=False)
|
||||
dpi = fig.get_dpi()
|
||||
width_in = (width * scale + 1e-2) / dpi
|
||||
height_in = (height * scale + 1e-2) / dpi
|
||||
fig.set_size_inches(width_in, height_in)
|
||||
ax = fig.add_axes([0.0, 0.0, 1.0, 1.0])
|
||||
ax.axis("off")
|
||||
ax.set_xlim(0.0, width)
|
||||
ax.set_ylim(height)
|
||||
|
||||
self.saveas = saveas
|
||||
self.rgb = rgb
|
||||
self.pynb = pynb
|
||||
self.img = img
|
||||
self.edgecolor = edgecolor
|
||||
self.alpha = 0.5
|
||||
self.linestyle = linestyle
|
||||
self.font_size = int(np.sqrt(min(height, width)) * scale // 3)
|
||||
self.width = width
|
||||
self.height = height
|
||||
self.scale = scale
|
||||
self.fig = fig
|
||||
self.ax = ax
|
||||
self.pad = pad
|
||||
self.id2obj = id2obj
|
||||
self.id2attr = id2attr
|
||||
self.canvas = FigureCanvasAgg(fig)
|
||||
|
||||
def add_box(self, box, color=None):
|
||||
if color is None:
|
||||
color = self.edgecolor
|
||||
(x0, y0, x1, y1) = box
|
||||
width = x1 - x0
|
||||
height = y1 - y0
|
||||
self.ax.add_patch(
|
||||
mpl.patches.Rectangle(
|
||||
(x0, y0),
|
||||
width,
|
||||
height,
|
||||
fill=False,
|
||||
edgecolor=color,
|
||||
linewidth=self.font_size // 3,
|
||||
alpha=self.alpha,
|
||||
linestyle=self.linestyle,
|
||||
)
|
||||
)
|
||||
|
||||
def draw_boxes(self, boxes, obj_ids=None, obj_scores=None, attr_ids=None, attr_scores=None):
|
||||
if len(boxes.shape) > 2:
|
||||
boxes = boxes[0]
|
||||
if len(obj_ids.shape) > 1:
|
||||
obj_ids = obj_ids[0]
|
||||
if len(obj_scores.shape) > 1:
|
||||
obj_scores = obj_scores[0]
|
||||
if len(attr_ids.shape) > 1:
|
||||
attr_ids = attr_ids[0]
|
||||
if len(attr_scores.shape) > 1:
|
||||
attr_scores = attr_scores[0]
|
||||
if isinstance(boxes, torch.Tensor):
|
||||
boxes = boxes.numpy()
|
||||
if isinstance(boxes, list):
|
||||
boxes = np.array(boxes)
|
||||
assert isinstance(boxes, np.ndarray)
|
||||
areas = np.prod(boxes[:, 2:] - boxes[:, :2], axis=1)
|
||||
sorted_idxs = np.argsort(-areas).tolist()
|
||||
boxes = boxes[sorted_idxs] if boxes is not None else None
|
||||
obj_ids = obj_ids[sorted_idxs] if obj_ids is not None else None
|
||||
obj_scores = obj_scores[sorted_idxs] if obj_scores is not None else None
|
||||
attr_ids = attr_ids[sorted_idxs] if attr_ids is not None else None
|
||||
attr_scores = attr_scores[sorted_idxs] if attr_scores is not None else None
|
||||
|
||||
assigned_colors = [self._random_color(maximum=1) for _ in range(len(boxes))]
|
||||
assigned_colors = [assigned_colors[idx] for idx in sorted_idxs]
|
||||
if obj_ids is not None:
|
||||
labels = self._create_text_labels_attr(obj_ids, obj_scores, attr_ids, attr_scores)
|
||||
for i in range(len(boxes)):
|
||||
color = assigned_colors[i]
|
||||
self.add_box(boxes[i], color)
|
||||
self.draw_labels(labels[i], boxes[i], color)
|
||||
|
||||
def draw_labels(self, label, box, color):
|
||||
x0, y0, x1, y1 = box
|
||||
text_pos = (x0, y0)
|
||||
instance_area = (y1 - y0) * (x1 - x0)
|
||||
small = _SMALL_OBJ * self.scale
|
||||
if instance_area < small or y1 - y0 < 40 * self.scale:
|
||||
if y1 >= self.height - 5:
|
||||
text_pos = (x1, y0)
|
||||
else:
|
||||
text_pos = (x0, y1)
|
||||
|
||||
height_ratio = (y1 - y0) / np.sqrt(self.height * self.width)
|
||||
lighter_color = self._change_color_brightness(color, brightness_factor=0.7)
|
||||
font_size = np.clip((height_ratio - 0.02) / 0.08 + 1, 1.2, 2)
|
||||
font_size *= 0.75 * self.font_size
|
||||
|
||||
self.draw_text(
|
||||
text=label,
|
||||
position=text_pos,
|
||||
color=lighter_color,
|
||||
)
|
||||
|
||||
def draw_text(
|
||||
self,
|
||||
text,
|
||||
position,
|
||||
color="g",
|
||||
ha="left",
|
||||
):
|
||||
rotation = 0
|
||||
font_size = self.font_size
|
||||
color = np.maximum(list(mplc.to_rgb(color)), 0.2)
|
||||
color[np.argmax(color)] = max(0.8, np.max(color))
|
||||
bbox = {
|
||||
"facecolor": "black",
|
||||
"alpha": self.alpha,
|
||||
"pad": self.pad,
|
||||
"edgecolor": "none",
|
||||
}
|
||||
x, y = position
|
||||
self.ax.text(
|
||||
x,
|
||||
y,
|
||||
text,
|
||||
size=font_size * self.scale,
|
||||
family="sans-serif",
|
||||
bbox=bbox,
|
||||
verticalalignment="top",
|
||||
horizontalalignment=ha,
|
||||
color=color,
|
||||
zorder=10,
|
||||
rotation=rotation,
|
||||
)
|
||||
|
||||
def save(self, saveas=None):
|
||||
if saveas is None:
|
||||
saveas = self.saveas
|
||||
if saveas.lower().endswith(".jpg") or saveas.lower().endswith(".png"):
|
||||
cv2.imwrite(
|
||||
saveas,
|
||||
self._get_buffer()[:, :, ::-1],
|
||||
)
|
||||
else:
|
||||
self.fig.savefig(saveas)
|
||||
|
||||
def _create_text_labels_attr(self, classes, scores, attr_classes, attr_scores):
|
||||
labels = [self.id2obj[i] for i in classes]
|
||||
attr_labels = [self.id2attr[i] for i in attr_classes]
|
||||
labels = [
|
||||
f"{label} {score:.2f} {attr} {attr_score:.2f}"
|
||||
for label, score, attr, attr_score in zip(labels, scores, attr_labels, attr_scores)
|
||||
]
|
||||
return labels
|
||||
|
||||
def _create_text_labels(self, classes, scores):
|
||||
labels = [self.id2obj[i] for i in classes]
|
||||
if scores is not None:
|
||||
if labels is None:
|
||||
labels = ["{:.0f}%".format(s * 100) for s in scores]
|
||||
else:
|
||||
labels = ["{} {:.0f}%".format(li, s * 100) for li, s in zip(labels, scores)]
|
||||
return labels
|
||||
|
||||
def _random_color(self, maximum=255):
|
||||
idx = np.random.randint(0, len(_COLORS))
|
||||
ret = _COLORS[idx] * maximum
|
||||
if not self.rgb:
|
||||
ret = ret[::-1]
|
||||
return ret
|
||||
|
||||
def _get_buffer(self):
|
||||
if not self.pynb:
|
||||
s, (width, height) = self.canvas.print_to_buffer()
|
||||
if (width, height) != (self.width, self.height):
|
||||
img = cv2.resize(self.img, (width, height))
|
||||
else:
|
||||
img = self.img
|
||||
else:
|
||||
buf = io.BytesIO() # works for cairo backend
|
||||
self.canvas.print_rgba(buf)
|
||||
width, height = self.width, self.height
|
||||
s = buf.getvalue()
|
||||
img = self.img
|
||||
|
||||
buffer = np.frombuffer(s, dtype="uint8")
|
||||
img_rgba = buffer.reshape(height, width, 4)
|
||||
rgb, alpha = np.split(img_rgba, [3], axis=2)
|
||||
|
||||
try:
|
||||
import numexpr as ne # fuse them with numexpr
|
||||
|
||||
visualized_image = ne.evaluate("img * (1 - alpha / 255.0) + rgb * (alpha / 255.0)")
|
||||
except ImportError:
|
||||
alpha = alpha.astype("float32") / 255.0
|
||||
visualized_image = img * (1 - alpha) + rgb * alpha
|
||||
|
||||
return visualized_image.astype("uint8")
|
||||
|
||||
def _change_color_brightness(self, color, brightness_factor):
|
||||
assert brightness_factor >= -1.0 and brightness_factor <= 1.0
|
||||
color = mplc.to_rgb(color)
|
||||
polygon_color = colorsys.rgb_to_hls(*mplc.to_rgb(color))
|
||||
modified_lightness = polygon_color[1] + (brightness_factor * polygon_color[1])
|
||||
modified_lightness = 0.0 if modified_lightness < 0.0 else modified_lightness
|
||||
modified_lightness = 1.0 if modified_lightness > 1.0 else modified_lightness
|
||||
modified_color = colorsys.hls_to_rgb(polygon_color[0], modified_lightness, polygon_color[2])
|
||||
return modified_color
|
||||
|
||||
|
||||
# Color map
|
||||
_COLORS = (
|
||||
np.array(
|
||||
[
|
||||
0.000,
|
||||
0.447,
|
||||
0.741,
|
||||
0.850,
|
||||
0.325,
|
||||
0.098,
|
||||
0.929,
|
||||
0.694,
|
||||
0.125,
|
||||
0.494,
|
||||
0.184,
|
||||
0.556,
|
||||
0.466,
|
||||
0.674,
|
||||
0.188,
|
||||
0.301,
|
||||
0.745,
|
||||
0.933,
|
||||
0.635,
|
||||
0.078,
|
||||
0.184,
|
||||
0.300,
|
||||
0.300,
|
||||
0.300,
|
||||
0.600,
|
||||
0.600,
|
||||
0.600,
|
||||
1.000,
|
||||
0.000,
|
||||
0.000,
|
||||
1.000,
|
||||
0.500,
|
||||
0.000,
|
||||
0.749,
|
||||
0.749,
|
||||
0.000,
|
||||
0.000,
|
||||
1.000,
|
||||
0.000,
|
||||
0.000,
|
||||
0.000,
|
||||
1.000,
|
||||
0.667,
|
||||
0.000,
|
||||
1.000,
|
||||
0.333,
|
||||
0.333,
|
||||
0.000,
|
||||
0.333,
|
||||
0.667,
|
||||
0.000,
|
||||
0.333,
|
||||
1.000,
|
||||
0.000,
|
||||
0.667,
|
||||
0.333,
|
||||
0.000,
|
||||
0.667,
|
||||
0.667,
|
||||
0.000,
|
||||
0.667,
|
||||
1.000,
|
||||
0.000,
|
||||
1.000,
|
||||
0.333,
|
||||
0.000,
|
||||
1.000,
|
||||
0.667,
|
||||
0.000,
|
||||
1.000,
|
||||
1.000,
|
||||
0.000,
|
||||
0.000,
|
||||
0.333,
|
||||
0.500,
|
||||
0.000,
|
||||
0.667,
|
||||
0.500,
|
||||
0.000,
|
||||
1.000,
|
||||
0.500,
|
||||
0.333,
|
||||
0.000,
|
||||
0.500,
|
||||
0.333,
|
||||
0.333,
|
||||
0.500,
|
||||
0.333,
|
||||
0.667,
|
||||
0.500,
|
||||
0.333,
|
||||
1.000,
|
||||
0.500,
|
||||
0.667,
|
||||
0.000,
|
||||
0.500,
|
||||
0.667,
|
||||
0.333,
|
||||
0.500,
|
||||
0.667,
|
||||
0.667,
|
||||
0.500,
|
||||
0.667,
|
||||
1.000,
|
||||
0.500,
|
||||
1.000,
|
||||
0.000,
|
||||
0.500,
|
||||
1.000,
|
||||
0.333,
|
||||
0.500,
|
||||
1.000,
|
||||
0.667,
|
||||
0.500,
|
||||
1.000,
|
||||
1.000,
|
||||
0.500,
|
||||
0.000,
|
||||
0.333,
|
||||
1.000,
|
||||
0.000,
|
||||
0.667,
|
||||
1.000,
|
||||
0.000,
|
||||
1.000,
|
||||
1.000,
|
||||
0.333,
|
||||
0.000,
|
||||
1.000,
|
||||
0.333,
|
||||
0.333,
|
||||
1.000,
|
||||
0.333,
|
||||
0.667,
|
||||
1.000,
|
||||
0.333,
|
||||
1.000,
|
||||
1.000,
|
||||
0.667,
|
||||
0.000,
|
||||
1.000,
|
||||
0.667,
|
||||
0.333,
|
||||
1.000,
|
||||
0.667,
|
||||
0.667,
|
||||
1.000,
|
||||
0.667,
|
||||
1.000,
|
||||
1.000,
|
||||
1.000,
|
||||
0.000,
|
||||
1.000,
|
||||
1.000,
|
||||
0.333,
|
||||
1.000,
|
||||
1.000,
|
||||
0.667,
|
||||
1.000,
|
||||
0.333,
|
||||
0.000,
|
||||
0.000,
|
||||
0.500,
|
||||
0.000,
|
||||
0.000,
|
||||
0.667,
|
||||
0.000,
|
||||
0.000,
|
||||
0.833,
|
||||
0.000,
|
||||
0.000,
|
||||
1.000,
|
||||
0.000,
|
||||
0.000,
|
||||
0.000,
|
||||
0.167,
|
||||
0.000,
|
||||
0.000,
|
||||
0.333,
|
||||
0.000,
|
||||
0.000,
|
||||
0.500,
|
||||
0.000,
|
||||
0.000,
|
||||
0.667,
|
||||
0.000,
|
||||
0.000,
|
||||
0.833,
|
||||
0.000,
|
||||
0.000,
|
||||
1.000,
|
||||
0.000,
|
||||
0.000,
|
||||
0.000,
|
||||
0.167,
|
||||
0.000,
|
||||
0.000,
|
||||
0.333,
|
||||
0.000,
|
||||
0.000,
|
||||
0.500,
|
||||
0.000,
|
||||
0.000,
|
||||
0.667,
|
||||
0.000,
|
||||
0.000,
|
||||
0.833,
|
||||
0.000,
|
||||
0.000,
|
||||
1.000,
|
||||
0.000,
|
||||
0.000,
|
||||
0.000,
|
||||
0.143,
|
||||
0.143,
|
||||
0.143,
|
||||
0.857,
|
||||
0.857,
|
||||
0.857,
|
||||
1.000,
|
||||
1.000,
|
||||
1.000,
|
||||
]
|
||||
)
|
||||
.astype(np.float32)
|
||||
.reshape(-1, 3)
|
||||
)
|
||||
@@ -934,7 +934,7 @@ def main():
|
||||
checkpoints = list(
|
||||
os.path.dirname(c) for c in sorted(glob.glob(args.output_dir + "/**/" + WEIGHTS_NAME, recursive=True))
|
||||
)
|
||||
|
||||
logging.getLogger("transformers.modeling_utils").setLevel(logging.WARN) # Reduce logging
|
||||
logger.info("Evaluate the following checkpoints: %s", checkpoints)
|
||||
for checkpoint in checkpoints:
|
||||
global_step = checkpoint.split("-")[-1] if len(checkpoints) > 1 else ""
|
||||
|
||||
@@ -1098,7 +1098,7 @@ def main():
|
||||
os.path.dirname(c)
|
||||
for c in sorted(glob.glob(args.output_dir + "/**/" + WEIGHTS_NAME, recursive=True))
|
||||
)
|
||||
|
||||
logging.getLogger("transformers.modeling_utils").setLevel(logging.WARN) # Reduce model loading logs
|
||||
else:
|
||||
logger.info("Loading checkpoint %s for evaluation", args.model_name_or_path)
|
||||
checkpoints = [args.model_name_or_path]
|
||||
|
||||
@@ -792,7 +792,7 @@ def main():
|
||||
os.path.dirname(c)
|
||||
for c in sorted(glob.glob(args.output_dir + "/**/" + WEIGHTS_NAME, recursive=True))
|
||||
)
|
||||
|
||||
logging.getLogger("transformers.modeling_utils").setLevel(logging.WARN) # Reduce model loading logs
|
||||
else:
|
||||
logger.info("Loading checkpoint %s for evaluation", args.model_name_or_path)
|
||||
checkpoints = [args.model_name_or_path]
|
||||
|
||||
@@ -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 - `RagSequenceForGeneration. and `RagTokenForGeneration`. 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_filename retrieval_preds.tsv \ # name of file in which predictions will be stored
|
||||
--eval_mode retrieval \ # indicates whether we're performing retrieval evaluation or e2e evaluation
|
||||
--recalculate # if predictions_filename already exists, and this option is set - we regenerate the answers, otherwise we reuse the predicsion file to calculate metrics.
|
||||
```
|
||||
|
||||
|
||||
## 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_filename e2e_preds.txt \
|
||||
--eval_mode e2e \ # indicates whether we're performing retrieval evaluation or e2e evaluation (default)
|
||||
--n_docs 5 \ # You can experiment with retrieving different number of documents at evaluation time
|
||||
--print_predictions
|
||||
```
|
||||
@@ -0,0 +1,30 @@
|
||||
import logging
|
||||
import os
|
||||
|
||||
from pytorch_lightning.callbacks import ModelCheckpoint
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def get_checkpoint_callback(output_dir, metric):
|
||||
"""Saves the best model by validation ROUGE2 score."""
|
||||
if metric == "rouge2":
|
||||
exp = "{val_avg_rouge2:.4f}-{step_count}"
|
||||
elif metric == "bleu":
|
||||
exp = "{val_avg_bleu:.4f}-{step_count}"
|
||||
elif metric == "em":
|
||||
exp = "{val_avg_em:.4f}-{step_count}"
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f"seq2seq callbacks only support rouge2 and bleu, got {metric}, You can make your own by adding to this function."
|
||||
)
|
||||
|
||||
checkpoint_callback = ModelCheckpoint(
|
||||
filepath=os.path.join(output_dir, exp),
|
||||
monitor=f"val_{metric}",
|
||||
mode="max",
|
||||
save_top_k=10,
|
||||
period=0, # maybe save a checkpoint every time val is run, not just end of epoch.
|
||||
)
|
||||
return checkpoint_callback
|
||||
@@ -0,0 +1,317 @@
|
||||
""" 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,
|
||||
RagSequenceForGeneration,
|
||||
RagTokenForGeneration,
|
||||
)
|
||||
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.numpy(), 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_filename",
|
||||
type=str,
|
||||
default="predictions.txt",
|
||||
help="Name of the predictions file, to be stored in the checkpoints directry",
|
||||
)
|
||||
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 = RagTokenForGeneration if args.model_type == "rag_token" else RagSequenceForGeneration
|
||||
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:
|
||||
predictions_path = os.path.join(checkpoint, args.predictions_filename)
|
||||
if os.path.exists(predictions_path) and (not args.recalculate):
|
||||
logger.info("Calculating metrics based on an existing predictions file: {}".format(predictions_path))
|
||||
score_fn(args, predictions_path, args.gold_data_path)
|
||||
continue
|
||||
|
||||
logger.info("***** Running evaluation for {} *****".format(checkpoint))
|
||||
logger.info(" Batch size = %d", args.eval_batch_size)
|
||||
logger.info(" Predictions will be stored under {}".format(predictions_path))
|
||||
|
||||
model = model_class.from_pretrained(checkpoint, **model_kwargs)
|
||||
model.to(args.device)
|
||||
retriever = RagRetriever(model.config) # TODO: add tokenizers
|
||||
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")
|
||||
)
|
||||
|
||||
with open(args.evaluation_set, "r") as eval_file, open(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, predictions_path, args.gold_data_path)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
args = get_args()
|
||||
main(args)
|
||||
@@ -0,0 +1,476 @@
|
||||
"""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,
|
||||
RagPyTorchDistributedRetriever,
|
||||
RagSequenceForGeneration,
|
||||
RagTokenForGeneration,
|
||||
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 = RagSequenceForGeneration
|
||||
elif hparams.model_type == "rag_token":
|
||||
self.model_class = RagTokenForGeneration
|
||||
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:
|
||||
generator_config = AutoConfig.from_pretrained(
|
||||
config.pretrained_generator_name_or_path, prefix=config.prefix
|
||||
)
|
||||
hparams, generator_config = set_extra_model_params(extra_model_params, hparams, generator_config)
|
||||
model = self.model_class.from_pretrained(
|
||||
hparams.model_name_or_path, config=config, generator_config=generator_config
|
||||
)
|
||||
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 = (
|
||||
RagPyTorchDistributedRetriever(self.model.config) if self.is_rag_model else None
|
||||
) # TODO add tokenizers
|
||||
|
||||
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
|
||||
@@ -15,4 +15,4 @@ pandas
|
||||
datasets
|
||||
fire
|
||||
pytest
|
||||
conllu
|
||||
conllu
|
||||
@@ -46,7 +46,7 @@ export DATA_DIR=${PWD}/wmt_en_de
|
||||
|
||||
#### Private Data
|
||||
|
||||
If you are using your own data, it must be formatted as one directory with 6 files:
|
||||
If you are using your own data, it must be formatted as one directory with 6 files:
|
||||
```
|
||||
train.source
|
||||
train.target
|
||||
@@ -227,81 +227,6 @@ python run_eval.py sshleifer/distilbart-cnn-12-6 $DATA_DIR/val.source dbart_val_
|
||||
--fp16 \
|
||||
--bs 32
|
||||
```
|
||||
### Multi-GPU Evalulation
|
||||
here is a command to run xsum evaluation on 8 GPUS. It is more than linearly faster than run_eval.py in some cases
|
||||
because it uses SortishSampler to minimize padding. You can also use it on 1 GPU. `data_dir` must have
|
||||
`{type_path}.source` and `{type_path}.target`. Run `python run_distributed_eval.py --help` for all clargs.
|
||||
|
||||
```bash
|
||||
python -m torch.distributed.launch --nproc_per_node=8 run_distributed_eval.py \
|
||||
--model_name sshleifer/distilbart-large-xsum-12-3 \
|
||||
--save_dir xsum_generations \
|
||||
--data_dir xsum \
|
||||
--fp16 # you can pass generate kwargs like num_beams here, just like run_eval.py
|
||||
```
|
||||
|
||||
Contributions that implement this command for other distributed hardware setups are welcome!
|
||||
|
||||
#### run_eval tips and tricks
|
||||
|
||||
When using `run_eval.py`, the following features can be useful:
|
||||
|
||||
* if you running the script multiple times and want to make it easier to track what arguments produced that output, use `--dump-args`. Along with the results it will also dump any custom params that were passed to the script. For example if you used: `--num_beams 8 --early_stopping true`, the output will be:
|
||||
```
|
||||
{'bleu': 26.887, 'n_obs': 10, 'runtime': 1, 'seconds_per_sample': 0.1, 'num_beams': 8, 'early_stopping': True}
|
||||
```
|
||||
|
||||
`--info` is an additional argument available for the same purpose of tracking the conditions of the experiment. It's useful to pass things that weren't in the argument list, e.g. a language pair `--info "lang:en-ru"`. But also if you pass `--info` without a value it will fallback to the current date/time string, e.g. `2020-09-13 18:44:43`.
|
||||
|
||||
If using `--dump-args --info`, the output will be:
|
||||
|
||||
```
|
||||
{'bleu': 26.887, 'n_obs': 10, 'runtime': 1, 'seconds_per_sample': 0.1, 'num_beams': 8, 'early_stopping': True, 'info': '2020-09-13 18:44:43'}
|
||||
```
|
||||
|
||||
If using `--dump-args --info "pair:en-ru chkpt=best`, the output will be:
|
||||
|
||||
```
|
||||
{'bleu': 26.887, 'n_obs': 10, 'runtime': 1, 'seconds_per_sample': 0.1, 'num_beams': 8, 'early_stopping': True, 'info': 'pair=en-ru chkpt=best'}
|
||||
```
|
||||
|
||||
|
||||
* if you need to perform a parametric search in order to find the best ones that lead to the highest BLEU score, let `run_eval_search.py` to do the searching for you.
|
||||
|
||||
The script accepts the exact same arguments as `run_eval.py`, plus an additional argument `--search`. The value of `--search` is parsed, reformatted and fed to ``run_eval.py`` as additional args.
|
||||
|
||||
The format for the `--search` value is a simple string with hparams and colon separated values to try, e.g.:
|
||||
```
|
||||
--search "num_beams=5:10 length_penalty=0.8:1.0:1.2 early_stopping=true:false"
|
||||
```
|
||||
which will generate `12` `(2*3*2)` searches for a product of each hparam. For example the example that was just used will invoke `run_eval.py` repeatedly with:
|
||||
|
||||
```
|
||||
--num_beams 5 --length_penalty 0.8 --early_stopping true
|
||||
--num_beams 5 --length_penalty 0.8 --early_stopping false
|
||||
[...]
|
||||
--num_beams 10 --length_penalty 1.2 --early_stopping false
|
||||
```
|
||||
|
||||
On completion, this function prints a markdown table of the results sorted by the best BLEU score and the winning arguments.
|
||||
|
||||
```
|
||||
bleu | num_beams | length_penalty | early_stopping
|
||||
----- | --------- | -------------- | --------------
|
||||
26.71 | 5 | 1.1 | 1
|
||||
26.66 | 5 | 0.9 | 1
|
||||
26.66 | 5 | 0.9 | 0
|
||||
26.41 | 5 | 1.1 | 0
|
||||
21.94 | 1 | 0.9 | 1
|
||||
21.94 | 1 | 0.9 | 0
|
||||
21.94 | 1 | 1.1 | 1
|
||||
21.94 | 1 | 1.1 | 0
|
||||
|
||||
Best score args:
|
||||
stas/wmt19-en-ru data/en-ru/val.source data/en-ru/test_translations.txt --reference_path data/en-ru/val.target --score_path data/en-ru/test_bleu.json --bs 8 --task translation --num_beams 5 --length_penalty 1.1 --early_stopping True
|
||||
```
|
||||
|
||||
If you pass `--info "some experiment-specific info"` it will get printed before the results table - this is useful for scripting and multiple runs, so one can tell the different sets of results from each other.
|
||||
|
||||
|
||||
### DistilBART
|
||||
|
||||
@@ -472,7 +472,6 @@ LAYERS_TO_COPY = {
|
||||
6: [0, 3, 6, 9, 12, 15],
|
||||
8: [0, 2, 4, 6, 8, 10, 12, 15],
|
||||
9: [0, 1, 3, 5, 7, 9, 11, 13, 15],
|
||||
12: [0, 1, 2, 3, 4, 5, 6, 7, 9, 11, 13, 15],
|
||||
16: list(range(16)),
|
||||
},
|
||||
6: {1: [0], 2: [0, 5], 3: [0, 2, 5], 4: [0, 1, 3, 5], 6: list(range(6))},
|
||||
|
||||
@@ -12,7 +12,7 @@ Note: You need to have your test_generations.txt before you start this process.
|
||||
cd $HOME
|
||||
git clone git@github.com:moses-smt/mosesdecoder.git
|
||||
cd mosesdecoder
|
||||
git clone git@github.com:rsennrich/wmt16-scripts.git
|
||||
git@github.com:rsennrich/wmt16-scripts.git
|
||||
```
|
||||
|
||||
(2) define a function for post processing.
|
||||
|
||||
@@ -1,223 +0,0 @@
|
||||
import argparse
|
||||
import shutil
|
||||
import time
|
||||
from json import JSONDecodeError
|
||||
from logging import getLogger
|
||||
from pathlib import Path
|
||||
from typing import Dict, List
|
||||
|
||||
import torch
|
||||
from torch.utils.data import DataLoader
|
||||
from tqdm import tqdm
|
||||
|
||||
from transformers import AutoModelForSeq2SeqLM, AutoTokenizer
|
||||
|
||||
|
||||
logger = getLogger(__name__)
|
||||
|
||||
try:
|
||||
from .utils import (
|
||||
Seq2SeqDataset,
|
||||
calculate_bleu,
|
||||
calculate_rouge,
|
||||
lmap,
|
||||
load_json,
|
||||
parse_numeric_n_bool_cl_kwargs,
|
||||
save_json,
|
||||
use_task_specific_params,
|
||||
write_txt_file,
|
||||
)
|
||||
except ImportError:
|
||||
from utils import (
|
||||
Seq2SeqDataset,
|
||||
calculate_bleu,
|
||||
calculate_rouge,
|
||||
lmap,
|
||||
load_json,
|
||||
parse_numeric_n_bool_cl_kwargs,
|
||||
save_json,
|
||||
use_task_specific_params,
|
||||
write_txt_file,
|
||||
)
|
||||
|
||||
|
||||
def eval_data_dir(
|
||||
data_dir,
|
||||
save_dir: str,
|
||||
model_name: str,
|
||||
bs: int = 8,
|
||||
max_source_length: int = 1024,
|
||||
type_path="val",
|
||||
n_obs=None,
|
||||
fp16=False,
|
||||
task="summarization",
|
||||
local_rank=None,
|
||||
**generate_kwargs,
|
||||
) -> Dict:
|
||||
"""Run evaluation on part of the data for one gpu and save to {save_dir}/rank_{rank}_output.json"""
|
||||
model_name = str(model_name)
|
||||
assert local_rank is not None
|
||||
torch.distributed.init_process_group(backend="nccl", rank=local_rank)
|
||||
|
||||
save_dir = Path(save_dir)
|
||||
save_path = save_dir.joinpath(f"rank_{local_rank}_output.json")
|
||||
torch.cuda.set_device(local_rank)
|
||||
model = AutoModelForSeq2SeqLM.from_pretrained(model_name).cuda()
|
||||
if fp16:
|
||||
model = model.half()
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(model_name)
|
||||
logger.info(f"Inferred tokenizer type: {tokenizer.__class__}") # if this is wrong, check config.model_type.
|
||||
use_task_specific_params(model, task) # update config with task specific params
|
||||
if max_source_length is None:
|
||||
max_source_length = tokenizer.model_max_length
|
||||
ds = Seq2SeqDataset(
|
||||
tokenizer,
|
||||
data_dir,
|
||||
max_source_length,
|
||||
max_target_length=1024,
|
||||
type_path=type_path,
|
||||
n_obs=n_obs,
|
||||
prefix=model.config.prefix,
|
||||
)
|
||||
# I set shuffle=True for a more accurate progress bar.
|
||||
# If all the longest samples are first, the prog bar estimate is too high at the beginning.
|
||||
sampler = ds.make_sortish_sampler(bs, distributed=True, add_extra_examples=False, shuffle=True)
|
||||
data_loader = DataLoader(ds, sampler=sampler, batch_size=bs, collate_fn=ds.collate_fn)
|
||||
results = []
|
||||
for batch in tqdm(data_loader):
|
||||
summaries = model.generate(
|
||||
input_ids=batch["input_ids"].to(model.device),
|
||||
attention_mask=batch["attention_mask"].to(model.device),
|
||||
**generate_kwargs,
|
||||
)
|
||||
preds = tokenizer.batch_decode(summaries, skip_special_tokens=True, clean_up_tokenization_spaces=False)
|
||||
ids = batch["ids"]
|
||||
for i, pred in enumerate(preds):
|
||||
results.append(dict(pred=pred, id=ids[i].item()))
|
||||
save_json(results, save_path)
|
||||
return results, sampler.num_replicas
|
||||
|
||||
|
||||
def run_generate():
|
||||
parser = argparse.ArgumentParser(
|
||||
epilog="Unspecified args like --num_beams=2 --decoder_start_token_id=4 are passed to model.generate"
|
||||
)
|
||||
parser.add_argument("--data_dir", type=str, help="like cnn_dm/test.source")
|
||||
parser.add_argument(
|
||||
"--model_name",
|
||||
type=str,
|
||||
help="like facebook/bart-large-cnn,t5-base, etc.",
|
||||
default="sshleifer/distilbart-xsum-12-3",
|
||||
)
|
||||
parser.add_argument("--save_dir", type=str, help="where to save", default="tmp_gen")
|
||||
parser.add_argument("--max_source_length", type=int, default=None)
|
||||
parser.add_argument(
|
||||
"--type_path", type=str, default="test", help="which subset to evaluate typically train/val/test"
|
||||
)
|
||||
parser.add_argument("--reference_path", type=str, required=False, help="like cnn_dm/test.target")
|
||||
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(
|
||||
"--local_rank", type=int, default=-1, required=False, help="should be passed by distributed.launch"
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--n_obs", type=int, default=None, required=False, help="How many observations. Defaults to all."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--sync_timeout",
|
||||
type=int,
|
||||
default=600,
|
||||
required=False,
|
||||
help="How long should master process wait for other processes to finish.",
|
||||
)
|
||||
parser.add_argument("--fp16", action="store_true")
|
||||
parser.add_argument("--debug", action="store_true")
|
||||
start_time = time.time()
|
||||
args, rest = parser.parse_known_args()
|
||||
generate_kwargs = parse_numeric_n_bool_cl_kwargs(rest)
|
||||
if generate_kwargs and args.local_rank <= 0:
|
||||
print(f"parsed the following generate kwargs: {generate_kwargs}")
|
||||
json_save_dir = Path(args.save_dir + "_tmp")
|
||||
Path(json_save_dir).mkdir(exist_ok=True) # this handles locking.
|
||||
intermediate_files = list(json_save_dir.glob("rank_*.json"))
|
||||
if intermediate_files:
|
||||
raise ValueError(f"Found files at {json_save_dir} please move or remove them.")
|
||||
# In theory, a node could finish and save before another node hits this. If this happens, we can address later.
|
||||
|
||||
Path(args.save_dir).mkdir(exist_ok=True)
|
||||
results, num_replicas = eval_data_dir(
|
||||
args.data_dir,
|
||||
json_save_dir,
|
||||
args.model_name,
|
||||
type_path=args.type_path,
|
||||
bs=args.bs,
|
||||
fp16=args.fp16,
|
||||
task=args.task,
|
||||
local_rank=args.local_rank,
|
||||
n_obs=args.n_obs,
|
||||
max_source_length=args.max_source_length,
|
||||
**generate_kwargs,
|
||||
)
|
||||
|
||||
if args.local_rank <= 0:
|
||||
save_dir = Path(args.save_dir)
|
||||
save_dir.mkdir(exist_ok=True)
|
||||
partial_results = gather_results_from_each_node(num_replicas, json_save_dir, args.sync_timeout)
|
||||
preds = combine_partial_results(partial_results)
|
||||
tgt_file = Path(args.data_dir).joinpath(args.type_path + ".target")
|
||||
labels = [x.rstrip() for x in open(tgt_file).readlines()][: len(preds)]
|
||||
|
||||
# Calculate metrics, save metrics, and save _generations.txt
|
||||
calc_bleu = "translation" in args.task
|
||||
score_fn = calculate_bleu if calc_bleu else calculate_rouge
|
||||
metric_name = "bleu" if calc_bleu else "rouge"
|
||||
metrics: Dict = score_fn(preds, labels)
|
||||
metrics["n_obs"] = len(preds)
|
||||
runtime = time.time() - start_time
|
||||
metrics["seconds_per_sample"] = round(runtime / metrics["n_obs"], 2)
|
||||
# TODO(@stas00): add whatever metadata to metrics
|
||||
metrics_save_path = save_dir.joinpath(f"{args.type_path}_{metric_name}.json")
|
||||
save_json(metrics, metrics_save_path, indent=None)
|
||||
print(metrics)
|
||||
write_txt_file(preds, save_dir.joinpath(f"{args.type_path}_generations.txt"))
|
||||
if args.debug:
|
||||
write_txt_file(labels, save_dir.joinpath(f"{args.type_path}.target"))
|
||||
else:
|
||||
shutil.rmtree(json_save_dir)
|
||||
|
||||
|
||||
def combine_partial_results(partial_results) -> List:
|
||||
"""Concatenate partial results into one file, then sort it by id."""
|
||||
records = []
|
||||
for partial_result in partial_results:
|
||||
records.extend(partial_result)
|
||||
records = list(sorted(records, key=lambda x: x["id"]))
|
||||
preds = [x["pred"] for x in records]
|
||||
return preds
|
||||
|
||||
|
||||
def gather_results_from_each_node(num_replicas, save_dir, timeout) -> List[Dict[str, List]]:
|
||||
# WAIT FOR lots of .json files
|
||||
start_wait = time.time()
|
||||
logger.info("waiting for all nodes to finish")
|
||||
json_data = None
|
||||
while (time.time() - start_wait) < timeout:
|
||||
json_files = list(save_dir.glob("rank_*.json"))
|
||||
if len(json_files) < num_replicas:
|
||||
continue
|
||||
try:
|
||||
# make sure all json files are fully saved
|
||||
json_data = lmap(load_json, json_files)
|
||||
return json_data
|
||||
except JSONDecodeError:
|
||||
continue
|
||||
else:
|
||||
raise TimeoutError("Rank 0 gave up on waiting for other processes")
|
||||
# Unreachable
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Usage for MT:
|
||||
run_generate()
|
||||
@@ -1,5 +1,4 @@
|
||||
import argparse
|
||||
import datetime
|
||||
import json
|
||||
import time
|
||||
import warnings
|
||||
@@ -16,9 +15,9 @@ from transformers import AutoModelForSeq2SeqLM, AutoTokenizer
|
||||
logger = getLogger(__name__)
|
||||
|
||||
try:
|
||||
from .utils import calculate_bleu, calculate_rouge, parse_numeric_n_bool_cl_kwargs, use_task_specific_params
|
||||
from .utils import calculate_bleu, calculate_rouge, parse_numeric_cl_kwargs, use_task_specific_params
|
||||
except ImportError:
|
||||
from utils import calculate_bleu, calculate_rouge, parse_numeric_n_bool_cl_kwargs, use_task_specific_params
|
||||
from utils import calculate_bleu, calculate_rouge, parse_numeric_cl_kwargs, use_task_specific_params
|
||||
|
||||
DEFAULT_DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
@@ -73,26 +72,7 @@ def generate_summaries_or_translations(
|
||||
return dict(n_obs=n_obs, runtime=runtime, seconds_per_sample=round(runtime / n_obs, 4))
|
||||
|
||||
|
||||
def datetime_now():
|
||||
return datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
||||
|
||||
|
||||
def run_generate(verbose=True):
|
||||
"""
|
||||
|
||||
Takes input text, generates output, and then using reference calculates the BLEU scores.
|
||||
|
||||
The results are saved to a file and returned to the caller, and printed out unless ``verbose=False`` is passed.
|
||||
|
||||
Args:
|
||||
verbose (:obj:`bool`, `optional`, defaults to :obj:`True`): print results to stdout
|
||||
|
||||
Returns:
|
||||
a tuple: ``(scores, params}``
|
||||
- ``scores``: a dict of scores data ``{'bleu': 39.6501, 'n_obs': 2000, 'runtime': 186, 'seconds_per_sample': 0.093}``
|
||||
- ``params``: a dict of custom params, e.g. ``{'num_beams': 5, 'length_penalty': 0.8}``
|
||||
"""
|
||||
|
||||
def run_generate():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("model_name", type=str, help="like facebook/bart-large-cnn,t5-base, etc.")
|
||||
parser.add_argument("input_path", type=str, help="like cnn_dm/test.source")
|
||||
@@ -109,19 +89,11 @@ def run_generate(verbose=True):
|
||||
"--n_obs", type=int, default=-1, required=False, help="How many observations. Defaults to all."
|
||||
)
|
||||
parser.add_argument("--fp16", action="store_true")
|
||||
parser.add_argument("--dump-args", action="store_true", help="print the custom hparams with the results")
|
||||
parser.add_argument(
|
||||
"--info",
|
||||
nargs="?",
|
||||
type=str,
|
||||
const=datetime_now(),
|
||||
help="use in conjunction w/ --dump-args to print with the results whatever other info you'd like, e.g. lang=en-ru. If no value is passed, the current datetime string will be used.",
|
||||
)
|
||||
# Unspecified args like --num_beams=2 --decoder_start_token_id=4 are passed to model.generate
|
||||
args, rest = parser.parse_known_args()
|
||||
parsed_args = parse_numeric_n_bool_cl_kwargs(rest)
|
||||
if parsed_args and verbose:
|
||||
print(f"parsed the following generate kwargs: {parsed_args}")
|
||||
parsed = parse_numeric_cl_kwargs(rest)
|
||||
if parsed:
|
||||
print(f"parsed the following generate kwargs: {parsed}")
|
||||
examples = [" " + x.rstrip() if "t5" in args.model_name else x.rstrip() for x in open(args.input_path).readlines()]
|
||||
if args.n_obs > 0:
|
||||
examples = examples[: args.n_obs]
|
||||
@@ -137,35 +109,23 @@ def run_generate(verbose=True):
|
||||
fp16=args.fp16,
|
||||
task=args.task,
|
||||
prefix=args.prefix,
|
||||
**parsed_args,
|
||||
**parsed,
|
||||
)
|
||||
|
||||
if args.reference_path is None:
|
||||
return {}
|
||||
|
||||
return
|
||||
# Compute scores
|
||||
score_fn = calculate_bleu if "translation" in args.task else calculate_rouge
|
||||
output_lns = [x.rstrip() for x in open(args.save_path).readlines()]
|
||||
reference_lns = [x.rstrip() for x in open(args.reference_path).readlines()][: len(output_lns)]
|
||||
scores: dict = score_fn(output_lns, reference_lns)
|
||||
scores.update(runtime_metrics)
|
||||
|
||||
if args.dump_args:
|
||||
scores.update(parsed_args)
|
||||
if args.info:
|
||||
scores["info"] = args.info
|
||||
|
||||
if verbose:
|
||||
print(scores)
|
||||
|
||||
print(scores)
|
||||
if args.score_path is not None:
|
||||
path = args.score_path
|
||||
json.dump(scores, open(path, "w"))
|
||||
|
||||
json.dump(scores, open(args.score_path, "w"))
|
||||
return scores
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Usage for MT:
|
||||
# python run_eval.py MODEL_NAME $DATA_DIR/test.source $save_dir/test_translations.txt --reference_path $DATA_DIR/test.target --score_path $save_dir/test_bleu.json --task translation $@
|
||||
run_generate(verbose=True)
|
||||
run_generate()
|
||||
|
||||
@@ -1,139 +0,0 @@
|
||||
import argparse
|
||||
import itertools
|
||||
import operator
|
||||
import sys
|
||||
from collections import OrderedDict
|
||||
|
||||
|
||||
try:
|
||||
from .run_eval import datetime_now, run_generate
|
||||
except ImportError:
|
||||
from run_eval import datetime_now, run_generate
|
||||
|
||||
|
||||
# A table of supported tasks and the list of scores in the order of importance to be sorted by.
|
||||
# To add a new task, simply list the score names that `run_eval.run_generate()` returns
|
||||
task_score_names = {
|
||||
"translation": ["bleu"],
|
||||
"translation_en_to_de": ["bleu"],
|
||||
"summarization": ["rouge1", "rouge2", "rougeL"],
|
||||
}
|
||||
|
||||
|
||||
def parse_search_arg(search):
|
||||
groups = search.split()
|
||||
entries = {k: vs for k, vs in (g.split("=") for g in groups)}
|
||||
entry_names = list(entries.keys())
|
||||
sets = [list((f"--{k} {v}") for v in vs.split(":")) for k, vs in entries.items()]
|
||||
matrix = [list(x) for x in itertools.product(*sets)]
|
||||
return matrix, entry_names
|
||||
|
||||
|
||||
def run_search():
|
||||
"""
|
||||
Run parametric search over the desired hparam space with help of ``run_eval.py``.
|
||||
|
||||
All the arguments except ``--search`` are passed to ``run_eval.py`` as is. The values inside of "--search" are parsed, reformatted and fed to ``run_eval.py`` as additional args.
|
||||
|
||||
The format for the ``--search`` value is a simple string with hparams and colon separated values to try, e.g.:
|
||||
```
|
||||
--search "num_beams=5:10 length_penalty=0.8:1.0:1.2 early_stopping=true:false"
|
||||
```
|
||||
which will generate ``12`` ``(2*3*2)`` searches for a product of each hparam. For example the example that was just used will invoke ``run_eval.py`` repeatedly with:
|
||||
|
||||
```
|
||||
--num_beams 5 --length_penalty 0.8 --early_stopping true
|
||||
--num_beams 5 --length_penalty 0.8 --early_stopping false
|
||||
[...]
|
||||
--num_beams 10 --length_penalty 1.2 --early_stopping false
|
||||
```
|
||||
|
||||
On completion, this function prints a markdown table of the results sorted by the best BLEU score and the winning arguments.
|
||||
|
||||
|
||||
"""
|
||||
prog = sys.argv[0]
|
||||
|
||||
parser = argparse.ArgumentParser(
|
||||
usage="\n\nImportant: this script accepts all arguments `run_eval.py` accepts and then a few extra, therefore refer to `run_eval.py -h` for the complete list."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--search",
|
||||
type=str,
|
||||
required=False,
|
||||
help='param space to search, e.g. "num_beams=5:10 length_penalty=0.8:1.0:1.2"',
|
||||
)
|
||||
parser.add_argument(
|
||||
"--bs", type=int, default=8, required=False, help="initial batch size (may get reduced if it's too big)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--task", type=str, help="used for task_specific_params + metrics", choices=task_score_names.keys()
|
||||
)
|
||||
parser.add_argument(
|
||||
"--info",
|
||||
nargs="?",
|
||||
type=str,
|
||||
const=datetime_now(),
|
||||
help="add custom notes to be printed before the results table. If no value is passed, the current datetime string will be used.",
|
||||
)
|
||||
args, args_main = parser.parse_known_args()
|
||||
# we share some of the args
|
||||
args_main.extend(["--task", args.task])
|
||||
args_normal = [prog] + args_main
|
||||
|
||||
matrix, col_names = parse_search_arg(args.search)
|
||||
col_names[0:0] = task_score_names[args.task] # score cols first
|
||||
col_widths = {col: len(str(col)) for col in col_names}
|
||||
results = []
|
||||
for r in matrix:
|
||||
hparams = {k: v for k, v in (x.replace("--", "").split() for x in r)}
|
||||
args_exp = " ".join(r).split()
|
||||
args_exp.extend(["--bs", str(args.bs)]) # in case we need to reduce its size due to CUDA OOM
|
||||
sys.argv = args_normal + args_exp
|
||||
|
||||
# XXX: need to trap CUDA OOM and lower args.bs if that happens and retry
|
||||
|
||||
scores = run_generate(verbose=False)
|
||||
# make sure scores are first in the table
|
||||
result = OrderedDict()
|
||||
for score in task_score_names[args.task]:
|
||||
result[score] = scores[score]
|
||||
result.update(hparams)
|
||||
results.append(result)
|
||||
|
||||
# find widest entries
|
||||
for k, v in result.items():
|
||||
l = len(str(v))
|
||||
if l > col_widths[k]:
|
||||
col_widths[k] = l
|
||||
|
||||
results_sorted = sorted(results, key=operator.itemgetter(*task_score_names[args.task]), reverse=True)
|
||||
print(" | ".join([f"{col:{col_widths[col]}}" for col in col_names]))
|
||||
print(" | ".join([f"{'-'*col_widths[col]}" for col in col_names]))
|
||||
for row in results_sorted:
|
||||
print(" | ".join([f"{row[col]:{col_widths[col]}}" for col in col_names]))
|
||||
|
||||
best = results_sorted[0]
|
||||
for score in task_score_names[args.task]:
|
||||
del best[score]
|
||||
best_args = [f"--{k} {v}" for k, v in best.items()]
|
||||
dyn_args = ["--bs", str(args.bs)]
|
||||
if args.info:
|
||||
print(f"\nInfo: {args.info}")
|
||||
print("\nBest score args:")
|
||||
print(" ".join(args_main + best_args + dyn_args))
|
||||
|
||||
return results_sorted
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Usage:
|
||||
# [normal-run_eval_search.py cmd plus] \
|
||||
# --search="num_beams=1:5:10 length_penalty=0.8:1:1.2 early_stopping=true:false"
|
||||
#
|
||||
# Example:
|
||||
# PYTHONPATH="src:examples/seq2seq" python examples/seq2seq/run_eval_search.py $MODEL_NAME \
|
||||
# $DATA_DIR/val.source $SAVE_DIR/test_translations.txt --reference_path $DATA_DIR/val.target \
|
||||
# --score_path $SAVE_DIR/test_bleu.json --bs $BS --task translation \
|
||||
# --search="num_beams=1:5:10 length_penalty=0.8:1:1.2 early_stopping=true:false"
|
||||
run_search()
|
||||
@@ -23,7 +23,6 @@ from .distillation import distill_main, evaluate_checkpoint
|
||||
from .finetune import SummarizationModule, main
|
||||
from .pack_dataset import pack_data_dir
|
||||
from .run_eval import generate_summaries_or_translations, run_generate
|
||||
from .run_eval_search import run_search
|
||||
from .utils import LegacySeq2SeqDataset, Seq2SeqDataset, label_smoothed_nll_loss, lmap, load_json
|
||||
|
||||
|
||||
@@ -284,7 +283,7 @@ class TestSummarizationDistiller(unittest.TestCase):
|
||||
return model
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", [pytest.param(T5_TINY), pytest.param(BART_TINY), pytest.param(MBART_TINY)])
|
||||
@pytest.mark.parametrize(["model"], [pytest.param(T5_TINY), pytest.param(BART_TINY), pytest.param(MBART_TINY)])
|
||||
def test_run_eval(model):
|
||||
input_file_name = Path(tempfile.mkdtemp()) / "utest_input.source"
|
||||
output_file_name = input_file_name.parent / "utest_output.txt"
|
||||
@@ -313,59 +312,6 @@ def test_run_eval(model):
|
||||
os.remove(Path(output_file_name))
|
||||
|
||||
|
||||
@slow
|
||||
@pytest.mark.parametrize("model", [pytest.param(T5_TINY)])
|
||||
def test_run_eval_search(model):
|
||||
input_file_name = Path(tempfile.mkdtemp()) / "utest_input.source"
|
||||
output_file_name = input_file_name.parent / "utest_output.txt"
|
||||
assert not output_file_name.exists()
|
||||
|
||||
text = {
|
||||
"en": ["Machine learning is great, isn't it?", "I like to eat bananas", "Tomorrow is another great day!"],
|
||||
"de": [
|
||||
"Maschinelles Lernen ist großartig, oder?",
|
||||
"Ich esse gerne Bananen",
|
||||
"Morgen ist wieder ein toller Tag!",
|
||||
],
|
||||
}
|
||||
|
||||
tmp_dir = Path(tempfile.mkdtemp())
|
||||
score_path = str(tmp_dir / "scores.json")
|
||||
reference_path = str(tmp_dir / "val.target")
|
||||
_dump_articles(input_file_name, text["en"])
|
||||
_dump_articles(reference_path, text["de"])
|
||||
task = "translation_en_to_de" if model == T5_TINY else "summarization"
|
||||
testargs = [
|
||||
"run_eval_search.py",
|
||||
model,
|
||||
str(input_file_name),
|
||||
str(output_file_name),
|
||||
"--score_path",
|
||||
score_path,
|
||||
"--reference_path",
|
||||
reference_path,
|
||||
"--task",
|
||||
task,
|
||||
"--search",
|
||||
"num_beams=1:2 length_penalty=0.9:1.0",
|
||||
]
|
||||
with patch.object(sys, "argv", testargs):
|
||||
with CaptureStdout() as cs:
|
||||
run_search()
|
||||
expected_strings = [" num_beams | length_penalty", model, "Best score args"]
|
||||
un_expected_strings = ["Info"]
|
||||
if "translation" in task:
|
||||
expected_strings.append("bleu")
|
||||
else:
|
||||
expected_strings.extend(["rouge1", "rouge2", "rougeL"])
|
||||
for w in expected_strings:
|
||||
assert w in cs.out
|
||||
for w in un_expected_strings:
|
||||
assert w not in cs.out
|
||||
assert Path(output_file_name).exists()
|
||||
os.remove(Path(output_file_name))
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
["model"],
|
||||
[pytest.param(T5_TINY), pytest.param(BART_TINY), pytest.param(MBART_TINY), pytest.param(MARIAN_TINY)],
|
||||
|
||||
+31
-54
@@ -18,7 +18,6 @@ from torch import nn
|
||||
from torch.utils.data import Dataset, Sampler
|
||||
|
||||
from transformers import BartTokenizer
|
||||
from transformers.file_utils import cached_property
|
||||
|
||||
|
||||
def label_smoothed_nll_loss(lprobs, target, epsilon, ignore_index=-100):
|
||||
@@ -99,8 +98,7 @@ class AbstractSeq2SeqDataset(Dataset):
|
||||
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 prefix is not None else ""
|
||||
|
||||
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
|
||||
@@ -115,11 +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, distributed=False, shuffle=True, **kwargs):
|
||||
def make_sortish_sampler(self, batch_size, distributed=False):
|
||||
if distributed:
|
||||
return DistributedSortishSampler(self, batch_size, shuffle=shuffle, **kwargs)
|
||||
return DistributedSortishSampler(self, batch_size)
|
||||
else:
|
||||
return SortishSampler(self.src_lens, batch_size, shuffle=shuffle)
|
||||
return SortishSampler(self.src_lens, batch_size)
|
||||
|
||||
def __getitem__(self, item):
|
||||
raise NotImplementedError("You must implement this")
|
||||
@@ -172,11 +170,14 @@ class Seq2SeqDataset(AbstractSeq2SeqDataset):
|
||||
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}"
|
||||
return {"tgt_texts": tgt_line, "src_texts": source_line, "id": index - 1}
|
||||
return {
|
||||
"tgt_texts": tgt_line,
|
||||
"src_texts": source_line,
|
||||
}
|
||||
|
||||
def collate_fn(self, batch) -> Dict[str, torch.Tensor]:
|
||||
"""Call prepare_seq2seq_batch."""
|
||||
batch_encoding: Dict[str, torch.Tensor] = self.tokenizer.prepare_seq2seq_batch(
|
||||
batch_encoding = self.tokenizer.prepare_seq2seq_batch(
|
||||
[x["src_texts"] for x in batch],
|
||||
src_lang=self.src_lang,
|
||||
tgt_texts=[x["tgt_texts"] for x in batch],
|
||||
@@ -185,28 +186,25 @@ class Seq2SeqDataset(AbstractSeq2SeqDataset):
|
||||
max_target_length=self.max_target_length,
|
||||
return_tensors="pt",
|
||||
add_prefix_space=self.add_prefix_space,
|
||||
).data
|
||||
batch_encoding["ids"] = torch.tensor([x["id"] for x in batch])
|
||||
return batch_encoding
|
||||
)
|
||||
return batch_encoding.data
|
||||
|
||||
|
||||
class SortishSampler(Sampler):
|
||||
"Go through the text data by order of src length with a bit of randomness. From fastai repo."
|
||||
|
||||
def __init__(self, data, batch_size, shuffle=True):
|
||||
self.data, self.bs, self.shuffle = data, batch_size, shuffle
|
||||
def __init__(self, data, batch_size):
|
||||
self.data, self.bs = data, batch_size
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self.data)
|
||||
|
||||
def __iter__(self):
|
||||
return iter(sortish_sampler_indices(self.data, self.bs, shuffle=self.shuffle))
|
||||
return iter(sortish_sampler_indices(self.data, self.bs))
|
||||
|
||||
|
||||
def sortish_sampler_indices(data: List, bs: int, shuffle=True) -> np.array:
|
||||
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."
|
||||
if not shuffle:
|
||||
return np.argsort(np.array(data) * -1)
|
||||
|
||||
def key_fn(i):
|
||||
return data[i]
|
||||
@@ -227,7 +225,7 @@ def sortish_sampler_indices(data: List, bs: int, shuffle=True) -> np.array:
|
||||
class DistributedSortishSampler(Sampler):
|
||||
"""Copied from torch DistributedSampler"""
|
||||
|
||||
def __init__(self, dataset, batch_size, num_replicas=None, rank=None, add_extra_examples=True, shuffle=True):
|
||||
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")
|
||||
@@ -240,28 +238,22 @@ class DistributedSortishSampler(Sampler):
|
||||
self.num_replicas = num_replicas
|
||||
self.rank = rank
|
||||
self.epoch = 0
|
||||
if add_extra_examples:
|
||||
self.num_samples = int(math.ceil(len(self.dataset) * 1.0 / self.num_replicas))
|
||||
self.total_size = self.num_samples * self.num_replicas
|
||||
else:
|
||||
self.total_size = len(dataset)
|
||||
self.num_samples = len(self.available_indices)
|
||||
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
|
||||
self.add_extra_examples = add_extra_examples
|
||||
self.shuffle = shuffle
|
||||
|
||||
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 self.available_indices]
|
||||
sortish_indices = sortish_sampler_indices(sortish_data, self.batch_size, shuffle=self.shuffle)
|
||||
indices = [self.available_indices[i] for i in sortish_indices]
|
||||
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)
|
||||
|
||||
@cached_property
|
||||
def available_indices(self) -> np.array:
|
||||
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))]
|
||||
@@ -312,9 +304,9 @@ def save_git_info(folder_path: str) -> None:
|
||||
save_json(repo_infos, os.path.join(folder_path, "git_log.json"))
|
||||
|
||||
|
||||
def save_json(content, path, indent=4, **json_dump_kwargs):
|
||||
def save_json(content, path):
|
||||
with open(path, "w") as f:
|
||||
json.dump(content, f, indent=indent, **json_dump_kwargs)
|
||||
json.dump(content, f, indent=4)
|
||||
|
||||
|
||||
def load_json(path):
|
||||
@@ -380,33 +372,18 @@ def assert_not_all_frozen(model):
|
||||
# CLI Parsing utils
|
||||
|
||||
|
||||
def parse_numeric_n_bool_cl_kwargs(unparsed_args: List[str]) -> Dict[str, Union[int, float, bool]]:
|
||||
"""
|
||||
Parse an argv list of unspecified command line args to a dict.
|
||||
Assumes all values are either numeric or boolean in the form of true/false.
|
||||
"""
|
||||
def parse_numeric_cl_kwargs(unparsed_args: List[str]) -> Dict[str, Union[int, float]]:
|
||||
"""Parse an argv list of unspecified command line args to a dict. Assumes all values are numeric."""
|
||||
result = {}
|
||||
assert len(unparsed_args) % 2 == 0, f"got odd number of unparsed args: {unparsed_args}"
|
||||
num_pairs = len(unparsed_args) // 2
|
||||
for pair_num in range(num_pairs):
|
||||
i = 2 * pair_num
|
||||
assert unparsed_args[i].startswith("--")
|
||||
if unparsed_args[i + 1].lower() == "true":
|
||||
value = True
|
||||
elif unparsed_args[i + 1].lower() == "false":
|
||||
value = False
|
||||
else:
|
||||
try:
|
||||
value = int(unparsed_args[i + 1])
|
||||
except ValueError:
|
||||
value = float(unparsed_args[i + 1]) # this can raise another informative ValueError
|
||||
try:
|
||||
value = int(unparsed_args[i + 1])
|
||||
except ValueError:
|
||||
value = float(unparsed_args[i + 1]) # this can raise another informative ValueError
|
||||
|
||||
result[unparsed_args[i][2:]] = value
|
||||
return result
|
||||
|
||||
|
||||
def write_txt_file(ordered_tgt, path):
|
||||
f = Path(path).open("w")
|
||||
for ln in ordered_tgt:
|
||||
f.write(ln + "\n")
|
||||
f.flush()
|
||||
|
||||
@@ -80,7 +80,10 @@ class ExamplesTests(TestCasePlus):
|
||||
--warmup_steps=2
|
||||
--seed=42
|
||||
--max_seq_length=128
|
||||
""".split()
|
||||
"""
|
||||
output_dir = "./tests/fixtures/tests_samples/temp_dir_{}".format(hash(testargs))
|
||||
testargs += "--output_dir " + output_dir
|
||||
testargs = testargs.split()
|
||||
|
||||
if is_cuda_and_apex_available():
|
||||
testargs.append("--fp16")
|
||||
@@ -146,14 +149,17 @@ class ExamplesTests(TestCasePlus):
|
||||
--do_train
|
||||
--do_eval
|
||||
--num_train_epochs=1
|
||||
""".split()
|
||||
"""
|
||||
output_dir = "./tests/fixtures/tests_samples/temp_dir_{}".format(hash(testargs))
|
||||
testargs += "--output_dir " + output_dir
|
||||
testargs = testargs.split()
|
||||
|
||||
if torch_device != "cuda":
|
||||
testargs.append("--no_cuda")
|
||||
|
||||
with patch.object(sys, "argv", testargs):
|
||||
result = run_language_modeling.main()
|
||||
self.assertLess(result["perplexity"], 42)
|
||||
self.assertLess(result["perplexity"], 35)
|
||||
|
||||
def test_run_squad(self):
|
||||
stream_handler = logging.StreamHandler(sys.stdout)
|
||||
|
||||
@@ -602,7 +602,7 @@ def main():
|
||||
checkpoints = list(
|
||||
os.path.dirname(c) for c in sorted(glob.glob(args.output_dir + "/**/" + WEIGHTS_NAME, recursive=True))
|
||||
)
|
||||
|
||||
logging.getLogger("transformers.modeling_utils").setLevel(logging.WARN) # Reduce logging
|
||||
logger.info("Evaluate the following checkpoints: %s", checkpoints)
|
||||
for checkpoint in checkpoints:
|
||||
global_step = checkpoint.split("-")[-1] if len(checkpoints) > 1 else ""
|
||||
|
||||
@@ -1,53 +0,0 @@
|
||||
---
|
||||
language: "fr"
|
||||
---
|
||||
|
||||
# BelGPT-2
|
||||
|
||||
**BelGPT-2** (*Belgian GPT-2* 🇧🇪) is a "small" GPT-2 model pre-trained on a very large and heterogeneous French corpus (around 60Gb). Please check [antoiloui/gpt2-french](https://github.com/antoiloui/gpt2-french) for more information about the pre-trained model, the data, the code to use the model and the code to pre-train it.
|
||||
|
||||
|
||||
## Using BelGPT-2 for Text Generation in French
|
||||
|
||||
You can use BelGPT-2 with [🤗 transformers](https://github.com/huggingface/transformers) library as follows:
|
||||
|
||||
```python
|
||||
import torch
|
||||
from transformers import GPT2Tokenizer, GPT2LMHeadModel
|
||||
|
||||
# Load pretrained model and tokenizer
|
||||
model = GPT2LMHeadModel.from_pretrained("antoiloui/belgpt2")
|
||||
tokenizer = GPT2Tokenizer.from_pretrained("antoiloui/belgpt2")
|
||||
|
||||
# Generate a sample of text
|
||||
model.eval()
|
||||
output = model.generate(
|
||||
bos_token_id=random.randint(1,50000),
|
||||
do_sample=True,
|
||||
top_k=50,
|
||||
max_length=100,
|
||||
top_p=0.95,
|
||||
num_return_sequences=1
|
||||
)
|
||||
|
||||
# Decode it
|
||||
decoded_output = []
|
||||
for sample in output:
|
||||
decoded_output.append(tokenizer.decode(sample, skip_special_tokens=True))
|
||||
print(decoded_output)
|
||||
```
|
||||
|
||||
## Data
|
||||
|
||||
Below is the list of all French copora used to pre-trained the model:
|
||||
|
||||
| Dataset | `$corpus_name` | Raw size | Cleaned size |
|
||||
| :------| :--- | :---: | :---: |
|
||||
| CommonCrawl | `common_crawl` | 200.2 GB | 40.4 GB |
|
||||
| NewsCrawl | `news_crawl` | 10.4 GB | 9.8 GB |
|
||||
| Wikipedia | `wiki` | 19.4 GB | 4.1 GB |
|
||||
| Wikisource | `wikisource` | 4.6 GB | 2.3 GB |
|
||||
| Project Gutenberg | `gutenberg` | 1.3 GB | 1.1 GB |
|
||||
| EuroParl | `europarl` | 289.9 MB | 278.7 MB |
|
||||
| NewsCommentary | `news_commentary` | 61.4 MB | 58.1 MB |
|
||||
| **Total** | | **236.3 GB** | **57.9 GB** |
|
||||
@@ -1,24 +0,0 @@
|
||||
The model can be loaded and used as follows on [this branch](https://github.com/huggingface/transformers/tree/finalize_rag) as follows.
|
||||
|
||||
|
||||
# Load model
|
||||
|
||||
```python
|
||||
from transformers import RagTokenizer, RagTokenForGeneration, RagRetriever
|
||||
|
||||
# create Retriever augmented model
|
||||
retriever = RagRetriever.from_pretrained("facebook/rag-token-nq_new", use_dummy_dataset=True)
|
||||
model = RagTokenForGeneration.from_pretrained("facebook/rag-token-nq_new", retriever=retriever)
|
||||
|
||||
tokenizer = RagTokenizer.from_pretrained("facebook/rag-token-nq_new")
|
||||
|
||||
# create input ids and labels
|
||||
input_ids = tokenizer("who sings does he love me with reba", return_tensors="pt").input_ids
|
||||
|
||||
# use labels
|
||||
labels = tokenizer.generator("Linda Davis", return_tensors="pt").input_ids
|
||||
|
||||
|
||||
# compute loss
|
||||
outputs = model(input_ids, labels=labels)
|
||||
```
|
||||
@@ -1,94 +0,0 @@
|
||||
---
|
||||
language: en
|
||||
license: apache-2.0
|
||||
datasets:
|
||||
- bookcorpus
|
||||
- wikipedia
|
||||
- gigaword
|
||||
---
|
||||
|
||||
# Funnel Transformer intermediate model (B6-6-6 without decoder)
|
||||
|
||||
Pretrained model on English language using a similar objective objective as [ELECTRA](https://huggingface.co/transformers/model_doc/electra.html). It was introduced in
|
||||
[this paper](https://arxiv.org/pdf/2006.03236.pdf) and first released in
|
||||
[this repository](https://github.com/laiguokun/Funnel-Transformer). This model is uncased: it does not make a difference
|
||||
between english and English.
|
||||
|
||||
Disclaimer: The team releasing Funnel Transformer did not write a model card for this model so this model card has been
|
||||
written by the Hugging Face team.
|
||||
|
||||
## Model description
|
||||
|
||||
Funnel Transformer is a transformers model pretrained on a large corpus of English data in a self-supervised fashion. This means it
|
||||
was pretrained on the raw texts only, with no humans labelling them in any way (which is why it can use lots of
|
||||
publicly available data) with an automatic process to generate inputs and labels from those texts.
|
||||
|
||||
More precisely, a small language model corrupts the input texts and serves as a generator of inputs for this model, and
|
||||
the pretraining objective is to predict which token is an original and which one has been replaced, a bit like a GAN training.
|
||||
|
||||
This way, the model learns an inner representation of the English language that can then be used to extract features
|
||||
useful for downstream tasks: if you have a dataset of labeled sentences for instance, you can train a standard
|
||||
classifier using the features produced by the BERT model as inputs.
|
||||
|
||||
**Note:** This model does not contain the decoder, so it ouputs hidden states that have a sequence length of one fourth
|
||||
of the inputs. It's good to use for tasks requiring a summary of the sentence (like sentence classification) but not if
|
||||
you need one input per initial token. You should use the `intermediate` model in that case.
|
||||
|
||||
## Intended uses & limitations
|
||||
|
||||
You can use the raw model to extract a vector representation of a given text, but it's mostly intended to
|
||||
be fine-tuned on a downstream task. See the [model hub](https://huggingface.co/models?filter=funnel-transformer) to look for
|
||||
fine-tuned versions on a task that interests you.
|
||||
|
||||
Note that this model is primarily aimed at being fine-tuned on tasks that use the whole sentence (potentially masked)
|
||||
to make decisions, such as sequence classification, token classification or question answering. For tasks such as text
|
||||
generation you should look at model like GPT2.
|
||||
|
||||
### How to use
|
||||
|
||||
|
||||
Here is how to use this model to get the features of a given text in PyTorch:
|
||||
|
||||
```python
|
||||
from transformers import FunnelTokenizer, FunnelBaseModel
|
||||
tokenizer = FunnelTokenizer.from_pretrained("funnel-transformer/intermediate-base")
|
||||
model = FunnelBaseModel.from_pretrained("funnel-transformer/intermediate-base")
|
||||
text = "Replace me by any text you'd like."
|
||||
encoded_input = tokenizer(text, return_tensors='pt')
|
||||
output = model(**encoded_input)
|
||||
```
|
||||
|
||||
and in TensorFlow:
|
||||
|
||||
```python
|
||||
from transformers import FunnelTokenizer, TFFunnelBaseModel
|
||||
tokenizer = FunnelTokenizer.from_pretrained("funnel-transformer/intermediate-base")
|
||||
model = TFFunnelBaseModel.from_pretrained("funnel-transformer/intermediate-base")
|
||||
text = "Replace me by any text you'd like."
|
||||
encoded_input = tokenizer(text, return_tensors='tf')
|
||||
output = model(encoded_input)
|
||||
```
|
||||
|
||||
## Training data
|
||||
|
||||
The BERT model was pretrained on:
|
||||
- [BookCorpus](https://yknzhu.wixsite.com/mbweb), a dataset consisting of 11,038 unpublished books,
|
||||
- [English Wikipedia](https://en.wikipedia.org/wiki/English_Wikipedia) (excluding lists, tables and headers),
|
||||
- [Clue Web](https://lemurproject.org/clueweb12/), a dataset of 733,019,372 English web pages,
|
||||
- [GigaWord](https://catalog.ldc.upenn.edu/LDC2011T07), an archive of newswire text data,
|
||||
- [Common Crawl](https://commoncrawl.org/), a dataset of raw web pages.
|
||||
|
||||
|
||||
### BibTeX entry and citation info
|
||||
|
||||
```bibtex
|
||||
@misc{dai2020funneltransformer,
|
||||
title={Funnel-Transformer: Filtering out Sequential Redundancy for Efficient Language Processing},
|
||||
author={Zihang Dai and Guokun Lai and Yiming Yang and Quoc V. Le},
|
||||
year={2020},
|
||||
eprint={2006.03236},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.LG}
|
||||
}
|
||||
```
|
||||
|
||||
@@ -1,90 +0,0 @@
|
||||
---
|
||||
language: en
|
||||
license: apache-2.0
|
||||
datasets:
|
||||
- bookcorpus
|
||||
- wikipedia
|
||||
- gigaword
|
||||
---
|
||||
|
||||
# Funnel Transformer intermediate model (B6-6-6 with decoder)
|
||||
|
||||
Pretrained model on English language using a similar objective objective as [ELECTRA](https://huggingface.co/transformers/model_doc/electra.html). It was introduced in
|
||||
[this paper](https://arxiv.org/pdf/2006.03236.pdf) and first released in
|
||||
[this repository](https://github.com/laiguokun/Funnel-Transformer). This model is uncased: it does not make a difference
|
||||
between english and English.
|
||||
|
||||
Disclaimer: The team releasing Funnel Transformer did not write a model card for this model so this model card has been
|
||||
written by the Hugging Face team.
|
||||
|
||||
## Model description
|
||||
|
||||
Funnel Transformer is a transformers model pretrained on a large corpus of English data in a self-supervised fashion. This means it
|
||||
was pretrained on the raw texts only, with no humans labelling them in any way (which is why it can use lots of
|
||||
publicly available data) with an automatic process to generate inputs and labels from those texts.
|
||||
|
||||
More precisely, a small language model corrupts the input texts and serves as a generator of inputs for this model, and
|
||||
the pretraining objective is to predict which token is an original and which one has been replaced, a bit like a GAN training.
|
||||
|
||||
This way, the model learns an inner representation of the English language that can then be used to extract features
|
||||
useful for downstream tasks: if you have a dataset of labeled sentences for instance, you can train a standard
|
||||
classifier using the features produced by the BERT model as inputs.
|
||||
|
||||
## Intended uses & limitations
|
||||
|
||||
You can use the raw model to extract a vector representation of a given text, but it's mostly intended to
|
||||
be fine-tuned on a downstream task. See the [model hub](https://huggingface.co/models?filter=funnel-transformer) to look for
|
||||
fine-tuned versions on a task that interests you.
|
||||
|
||||
Note that this model is primarily aimed at being fine-tuned on tasks that use the whole sentence (potentially masked)
|
||||
to make decisions, such as sequence classification, token classification or question answering. For tasks such as text
|
||||
generation you should look at model like GPT2.
|
||||
|
||||
### How to use
|
||||
|
||||
|
||||
Here is how to use this model to get the features of a given text in PyTorch:
|
||||
|
||||
```python
|
||||
from transformers import FunnelTokenizer, FunnelModel
|
||||
tokenizer = FunnelTokenizer.from_pretrained("funnel-transformer/intermediate")
|
||||
model = FunneModel.from_pretrained("funnel-transformer/intermediate")
|
||||
text = "Replace me by any text you'd like."
|
||||
encoded_input = tokenizer(text, return_tensors='pt')
|
||||
output = model(**encoded_input)
|
||||
```
|
||||
|
||||
and in TensorFlow:
|
||||
|
||||
```python
|
||||
from transformers import FunnelTokenizer, TFFunnelModel
|
||||
tokenizer = FunnelTokenizer.from_pretrained("funnel-transformer/intermediate")
|
||||
model = TFFunnelModel.from_pretrained("funnel-transformer/intermediatesmall")
|
||||
text = "Replace me by any text you'd like."
|
||||
encoded_input = tokenizer(text, return_tensors='tf')
|
||||
output = model(encoded_input)
|
||||
```
|
||||
|
||||
## Training data
|
||||
|
||||
The BERT model was pretrained on:
|
||||
- [BookCorpus](https://yknzhu.wixsite.com/mbweb), a dataset consisting of 11,038 unpublished books,
|
||||
- [English Wikipedia](https://en.wikipedia.org/wiki/English_Wikipedia) (excluding lists, tables and headers),
|
||||
- [Clue Web](https://lemurproject.org/clueweb12/), a dataset of 733,019,372 English web pages,
|
||||
- [GigaWord](https://catalog.ldc.upenn.edu/LDC2011T07), an archive of newswire text data,
|
||||
- [Common Crawl](https://commoncrawl.org/), a dataset of raw web pages.
|
||||
|
||||
|
||||
### BibTeX entry and citation info
|
||||
|
||||
```bibtex
|
||||
@misc{dai2020funneltransformer,
|
||||
title={Funnel-Transformer: Filtering out Sequential Redundancy for Efficient Language Processing},
|
||||
author={Zihang Dai and Guokun Lai and Yiming Yang and Quoc V. Le},
|
||||
year={2020},
|
||||
eprint={2006.03236},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.LG}
|
||||
}
|
||||
```
|
||||
|
||||
@@ -1,94 +0,0 @@
|
||||
---
|
||||
language: en
|
||||
license: apache-2.0
|
||||
datasets:
|
||||
- bookcorpus
|
||||
- wikipedia
|
||||
- gigaword
|
||||
---
|
||||
|
||||
# Funnel Transformer large model (B8-8-8 without decoder)
|
||||
|
||||
Pretrained model on English language using a similar objective objective as [ELECTRA](https://huggingface.co/transformers/model_doc/electra.html). It was introduced in
|
||||
[this paper](https://arxiv.org/pdf/2006.03236.pdf) and first released in
|
||||
[this repository](https://github.com/laiguokun/Funnel-Transformer). This model is uncased: it does not make a difference
|
||||
between english and English.
|
||||
|
||||
Disclaimer: The team releasing Funnel Transformer did not write a model card for this model so this model card has been
|
||||
written by the Hugging Face team.
|
||||
|
||||
## Model description
|
||||
|
||||
Funnel Transformer is a transformers model pretrained on a large corpus of English data in a self-supervised fashion. This means it
|
||||
was pretrained on the raw texts only, with no humans labelling them in any way (which is why it can use lots of
|
||||
publicly available data) with an automatic process to generate inputs and labels from those texts.
|
||||
|
||||
More precisely, a small language model corrupts the input texts and serves as a generator of inputs for this model, and
|
||||
the pretraining objective is to predict which token is an original and which one has been replaced, a bit like a GAN training.
|
||||
|
||||
This way, the model learns an inner representation of the English language that can then be used to extract features
|
||||
useful for downstream tasks: if you have a dataset of labeled sentences for instance, you can train a standard
|
||||
classifier using the features produced by the BERT model as inputs.
|
||||
|
||||
**Note:** This model does not contain the decoder, so it ouputs hidden states that have a sequence length of one fourth
|
||||
of the inputs. It's good to use for tasks requiring a summary of the sentence (like sentence classification) but not if
|
||||
you need one input per initial token. You should use the `large` model in that case.
|
||||
|
||||
## Intended uses & limitations
|
||||
|
||||
You can use the raw model to extract a vector representation of a given text, but it's mostly intended to
|
||||
be fine-tuned on a downstream task. See the [model hub](https://huggingface.co/models?filter=funnel-transformer) to look for
|
||||
fine-tuned versions on a task that interests you.
|
||||
|
||||
Note that this model is primarily aimed at being fine-tuned on tasks that use the whole sentence (potentially masked)
|
||||
to make decisions, such as sequence classification, token classification or question answering. For tasks such as text
|
||||
generation you should look at model like GPT2.
|
||||
|
||||
### How to use
|
||||
|
||||
|
||||
Here is how to use this model to get the features of a given text in PyTorch:
|
||||
|
||||
```python
|
||||
from transformers import FunnelTokenizer, FunnelBaseModel
|
||||
tokenizer = FunnelTokenizer.from_pretrained("funnel-transformer/large-base")
|
||||
model = FunnelBaseModel.from_pretrained("funnel-transformer/large-base")
|
||||
text = "Replace me by any text you'd like."
|
||||
encoded_input = tokenizer(text, return_tensors='pt')
|
||||
output = model(**encoded_input)
|
||||
```
|
||||
|
||||
and in TensorFlow:
|
||||
|
||||
```python
|
||||
from transformers import FunnelTokenizer, TFFunnelBaseModel
|
||||
tokenizer = FunnelTokenizer.from_pretrained("funnel-transformer/large-base")
|
||||
model = TFFunnelBaseModel.from_pretrained("funnel-transformer/large-base")
|
||||
text = "Replace me by any text you'd like."
|
||||
encoded_input = tokenizer(text, return_tensors='tf')
|
||||
output = model(encoded_input)
|
||||
```
|
||||
|
||||
## Training data
|
||||
|
||||
The BERT model was pretrained on:
|
||||
- [BookCorpus](https://yknzhu.wixsite.com/mbweb), a dataset consisting of 11,038 unpublished books,
|
||||
- [English Wikipedia](https://en.wikipedia.org/wiki/English_Wikipedia) (excluding lists, tables and headers),
|
||||
- [Clue Web](https://lemurproject.org/clueweb12/), a dataset of 733,019,372 English web pages,
|
||||
- [GigaWord](https://catalog.ldc.upenn.edu/LDC2011T07), an archive of newswire text data,
|
||||
- [Common Crawl](https://commoncrawl.org/), a dataset of raw web pages.
|
||||
|
||||
|
||||
### BibTeX entry and citation info
|
||||
|
||||
```bibtex
|
||||
@misc{dai2020funneltransformer,
|
||||
title={Funnel-Transformer: Filtering out Sequential Redundancy for Efficient Language Processing},
|
||||
author={Zihang Dai and Guokun Lai and Yiming Yang and Quoc V. Le},
|
||||
year={2020},
|
||||
eprint={2006.03236},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.LG}
|
||||
}
|
||||
```
|
||||
|
||||
@@ -1,90 +0,0 @@
|
||||
---
|
||||
language: en
|
||||
license: apache-2.0
|
||||
datasets:
|
||||
- bookcorpus
|
||||
- wikipedia
|
||||
- gigaword
|
||||
---
|
||||
|
||||
# Funnel Transformer large model (B8-8-8 with decoder)
|
||||
|
||||
Pretrained model on English language using a similar objective objective as [ELECTRA](https://huggingface.co/transformers/model_doc/electra.html). It was introduced in
|
||||
[this paper](https://arxiv.org/pdf/2006.03236.pdf) and first released in
|
||||
[this repository](https://github.com/laiguokun/Funnel-Transformer). This model is uncased: it does not make a difference
|
||||
between english and English.
|
||||
|
||||
Disclaimer: The team releasing Funnel Transformer did not write a model card for this model so this model card has been
|
||||
written by the Hugging Face team.
|
||||
|
||||
## Model description
|
||||
|
||||
Funnel Transformer is a transformers model pretrained on a large corpus of English data in a self-supervised fashion. This means it
|
||||
was pretrained on the raw texts only, with no humans labelling them in any way (which is why it can use lots of
|
||||
publicly available data) with an automatic process to generate inputs and labels from those texts.
|
||||
|
||||
More precisely, a small language model corrupts the input texts and serves as a generator of inputs for this model, and
|
||||
the pretraining objective is to predict which token is an original and which one has been replaced, a bit like a GAN training.
|
||||
|
||||
This way, the model learns an inner representation of the English language that can then be used to extract features
|
||||
useful for downstream tasks: if you have a dataset of labeled sentences for instance, you can train a standard
|
||||
classifier using the features produced by the BERT model as inputs.
|
||||
|
||||
## Intended uses & limitations
|
||||
|
||||
You can use the raw model to extract a vector representation of a given text, but it's mostly intended to
|
||||
be fine-tuned on a downstream task. See the [model hub](https://huggingface.co/models?filter=funnel-transformer) to look for
|
||||
fine-tuned versions on a task that interests you.
|
||||
|
||||
Note that this model is primarily aimed at being fine-tuned on tasks that use the whole sentence (potentially masked)
|
||||
to make decisions, such as sequence classification, token classification or question answering. For tasks such as text
|
||||
generation you should look at model like GPT2.
|
||||
|
||||
### How to use
|
||||
|
||||
|
||||
Here is how to use this model to get the features of a given text in PyTorch:
|
||||
|
||||
```python
|
||||
from transformers import FunnelTokenizer, FunnelModel
|
||||
tokenizer = FunnelTokenizer.from_pretrained("funnel-transformer/large")
|
||||
model = FunneModel.from_pretrained("funnel-transformer/large")
|
||||
text = "Replace me by any text you'd like."
|
||||
encoded_input = tokenizer(text, return_tensors='pt')
|
||||
output = model(**encoded_input)
|
||||
```
|
||||
|
||||
and in TensorFlow:
|
||||
|
||||
```python
|
||||
from transformers import FunnelTokenizer, TFFunnelModel
|
||||
tokenizer = FunnelTokenizer.from_pretrained("funnel-transformer/large")
|
||||
model = TFFunnelModel.from_pretrained("funnel-transformer/large")
|
||||
text = "Replace me by any text you'd like."
|
||||
encoded_input = tokenizer(text, return_tensors='tf')
|
||||
output = model(encoded_input)
|
||||
```
|
||||
|
||||
## Training data
|
||||
|
||||
The BERT model was pretrained on:
|
||||
- [BookCorpus](https://yknzhu.wixsite.com/mbweb), a dataset consisting of 11,038 unpublished books,
|
||||
- [English Wikipedia](https://en.wikipedia.org/wiki/English_Wikipedia) (excluding lists, tables and headers),
|
||||
- [Clue Web](https://lemurproject.org/clueweb12/), a dataset of 733,019,372 English web pages,
|
||||
- [GigaWord](https://catalog.ldc.upenn.edu/LDC2011T07), an archive of newswire text data,
|
||||
- [Common Crawl](https://commoncrawl.org/), a dataset of raw web pages.
|
||||
|
||||
|
||||
### BibTeX entry and citation info
|
||||
|
||||
```bibtex
|
||||
@misc{dai2020funneltransformer,
|
||||
title={Funnel-Transformer: Filtering out Sequential Redundancy for Efficient Language Processing},
|
||||
author={Zihang Dai and Guokun Lai and Yiming Yang and Quoc V. Le},
|
||||
year={2020},
|
||||
eprint={2006.03236},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.LG}
|
||||
}
|
||||
```
|
||||
|
||||
@@ -1,94 +0,0 @@
|
||||
---
|
||||
language: en
|
||||
license: apache-2.0
|
||||
datasets:
|
||||
- bookcorpus
|
||||
- wikipedia
|
||||
- gigaword
|
||||
---
|
||||
|
||||
# Funnel Transformer medium model (B6-3x2-3x2 without decoder)
|
||||
|
||||
Pretrained model on English language using a similar objective objective as [ELECTRA](https://huggingface.co/transformers/model_doc/electra.html). It was introduced in
|
||||
[this paper](https://arxiv.org/pdf/2006.03236.pdf) and first released in
|
||||
[this repository](https://github.com/laiguokun/Funnel-Transformer). This model is uncased: it does not make a difference
|
||||
between english and English.
|
||||
|
||||
Disclaimer: The team releasing Funnel Transformer did not write a model card for this model so this model card has been
|
||||
written by the Hugging Face team.
|
||||
|
||||
## Model description
|
||||
|
||||
Funnel Transformer is a transformers model pretrained on a large corpus of English data in a self-supervised fashion. This means it
|
||||
was pretrained on the raw texts only, with no humans labelling them in any way (which is why it can use lots of
|
||||
publicly available data) with an automatic process to generate inputs and labels from those texts.
|
||||
|
||||
More precisely, a small language model corrupts the input texts and serves as a generator of inputs for this model, and
|
||||
the pretraining objective is to predict which token is an original and which one has been replaced, a bit like a GAN training.
|
||||
|
||||
This way, the model learns an inner representation of the English language that can then be used to extract features
|
||||
useful for downstream tasks: if you have a dataset of labeled sentences for instance, you can train a standard
|
||||
classifier using the features produced by the BERT model as inputs.
|
||||
|
||||
**Note:** This model does not contain the decoder, so it ouputs hidden states that have a sequence length of one fourth
|
||||
of the inputs. It's good to use for tasks requiring a summary of the sentence (like sentence classification) but not if
|
||||
you need one input per initial token. You should use the `medium` model in that case.
|
||||
|
||||
## Intended uses & limitations
|
||||
|
||||
You can use the raw model to extract a vector representation of a given text, but it's mostly intended to
|
||||
be fine-tuned on a downstream task. See the [model hub](https://huggingface.co/models?filter=funnel-transformer) to look for
|
||||
fine-tuned versions on a task that interests you.
|
||||
|
||||
Note that this model is primarily aimed at being fine-tuned on tasks that use the whole sentence (potentially masked)
|
||||
to make decisions, such as sequence classification, token classification or question answering. For tasks such as text
|
||||
generation you should look at model like GPT2.
|
||||
|
||||
### How to use
|
||||
|
||||
|
||||
Here is how to use this model to get the features of a given text in PyTorch:
|
||||
|
||||
```python
|
||||
from transformers import FunnelTokenizer, FunnelBaseModel
|
||||
tokenizer = FunnelTokenizer.from_pretrained("funnel-transformer/medium-base")
|
||||
model = FunnelBaseModel.from_pretrained("funnel-transformer/medium-base")
|
||||
text = "Replace me by any text you'd like."
|
||||
encoded_input = tokenizer(text, return_tensors='pt')
|
||||
output = model(**encoded_input)
|
||||
```
|
||||
|
||||
and in TensorFlow:
|
||||
|
||||
```python
|
||||
from transformers import FunnelTokenizer, TFFunnelBaseModel
|
||||
tokenizer = FunnelTokenizer.from_pretrained("funnel-transformer/medium-base")
|
||||
model = TFFunnelBaseModel.from_pretrained("funnel-transformer/medium-base")
|
||||
text = "Replace me by any text you'd like."
|
||||
encoded_input = tokenizer(text, return_tensors='tf')
|
||||
output = model(encoded_input)
|
||||
```
|
||||
|
||||
## Training data
|
||||
|
||||
The BERT model was pretrained on:
|
||||
- [BookCorpus](https://yknzhu.wixsite.com/mbweb), a dataset consisting of 11,038 unpublished books,
|
||||
- [English Wikipedia](https://en.wikipedia.org/wiki/English_Wikipedia) (excluding lists, tables and headers),
|
||||
- [Clue Web](https://lemurproject.org/clueweb12/), a dataset of 733,019,372 English web pages,
|
||||
- [GigaWord](https://catalog.ldc.upenn.edu/LDC2011T07), an archive of newswire text data,
|
||||
- [Common Crawl](https://commoncrawl.org/), a dataset of raw web pages.
|
||||
|
||||
|
||||
### BibTeX entry and citation info
|
||||
|
||||
```bibtex
|
||||
@misc{dai2020funneltransformer,
|
||||
title={Funnel-Transformer: Filtering out Sequential Redundancy for Efficient Language Processing},
|
||||
author={Zihang Dai and Guokun Lai and Yiming Yang and Quoc V. Le},
|
||||
year={2020},
|
||||
eprint={2006.03236},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.LG}
|
||||
}
|
||||
```
|
||||
|
||||
@@ -1,90 +0,0 @@
|
||||
---
|
||||
language: en
|
||||
license: apache-2.0
|
||||
datasets:
|
||||
- bookcorpus
|
||||
- wikipedia
|
||||
- gigaword
|
||||
---
|
||||
|
||||
# Funnel Transformer medium model (B6-3x2-3x2 with decoder)
|
||||
|
||||
Pretrained model on English language using a similar objective objective as [ELECTRA](https://huggingface.co/transformers/model_doc/electra.html). It was introduced in
|
||||
[this paper](https://arxiv.org/pdf/2006.03236.pdf) and first released in
|
||||
[this repository](https://github.com/laiguokun/Funnel-Transformer). This model is uncased: it does not make a difference
|
||||
between english and English.
|
||||
|
||||
Disclaimer: The team releasing Funnel Transformer did not write a model card for this model so this model card has been
|
||||
written by the Hugging Face team.
|
||||
|
||||
## Model description
|
||||
|
||||
Funnel Transformer is a transformers model pretrained on a large corpus of English data in a self-supervised fashion. This means it
|
||||
was pretrained on the raw texts only, with no humans labelling them in any way (which is why it can use lots of
|
||||
publicly available data) with an automatic process to generate inputs and labels from those texts.
|
||||
|
||||
More precisely, a small language model corrupts the input texts and serves as a generator of inputs for this model, and
|
||||
the pretraining objective is to predict which token is an original and which one has been replaced, a bit like a GAN training.
|
||||
|
||||
This way, the model learns an inner representation of the English language that can then be used to extract features
|
||||
useful for downstream tasks: if you have a dataset of labeled sentences for instance, you can train a standard
|
||||
classifier using the features produced by the BERT model as inputs.
|
||||
|
||||
## Intended uses & limitations
|
||||
|
||||
You can use the raw model to extract a vector representation of a given text, but it's mostly intended to
|
||||
be fine-tuned on a downstream task. See the [model hub](https://huggingface.co/models?filter=funnel-transformer) to look for
|
||||
fine-tuned versions on a task that interests you.
|
||||
|
||||
Note that this model is primarily aimed at being fine-tuned on tasks that use the whole sentence (potentially masked)
|
||||
to make decisions, such as sequence classification, token classification or question answering. For tasks such as text
|
||||
generation you should look at model like GPT2.
|
||||
|
||||
### How to use
|
||||
|
||||
|
||||
Here is how to use this model to get the features of a given text in PyTorch:
|
||||
|
||||
```python
|
||||
from transformers import FunnelTokenizer, FunnelModel
|
||||
tokenizer = FunnelTokenizer.from_pretrained("funnel-transformer/medium")
|
||||
model = FunneModel.from_pretrained("funnel-transformer/medium")
|
||||
text = "Replace me by any text you'd like."
|
||||
encoded_input = tokenizer(text, return_tensors='pt')
|
||||
output = model(**encoded_input)
|
||||
```
|
||||
|
||||
and in TensorFlow:
|
||||
|
||||
```python
|
||||
from transformers import FunnelTokenizer, TFFunnelModel
|
||||
tokenizer = FunnelTokenizer.from_pretrained("funnel-transformer/medium")
|
||||
model = TFFunnelModel.from_pretrained("funnel-transformer/medium")
|
||||
text = "Replace me by any text you'd like."
|
||||
encoded_input = tokenizer(text, return_tensors='tf')
|
||||
output = model(encoded_input)
|
||||
```
|
||||
|
||||
## Training data
|
||||
|
||||
The BERT model was pretrained on:
|
||||
- [BookCorpus](https://yknzhu.wixsite.com/mbweb), a dataset consisting of 11,038 unpublished books,
|
||||
- [English Wikipedia](https://en.wikipedia.org/wiki/English_Wikipedia) (excluding lists, tables and headers),
|
||||
- [Clue Web](https://lemurproject.org/clueweb12/), a dataset of 733,019,372 English web pages,
|
||||
- [GigaWord](https://catalog.ldc.upenn.edu/LDC2011T07), an archive of newswire text data,
|
||||
- [Common Crawl](https://commoncrawl.org/), a dataset of raw web pages.
|
||||
|
||||
|
||||
### BibTeX entry and citation info
|
||||
|
||||
```bibtex
|
||||
@misc{dai2020funneltransformer,
|
||||
title={Funnel-Transformer: Filtering out Sequential Redundancy for Efficient Language Processing},
|
||||
author={Zihang Dai and Guokun Lai and Yiming Yang and Quoc V. Le},
|
||||
year={2020},
|
||||
eprint={2006.03236},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.LG}
|
||||
}
|
||||
```
|
||||
|
||||
@@ -1,94 +0,0 @@
|
||||
---
|
||||
language: en
|
||||
license: apache-2.0
|
||||
datasets:
|
||||
- bookcorpus
|
||||
- wikipedia
|
||||
- gigaword
|
||||
---
|
||||
|
||||
# Funnel Transformer small model (B4-4-4 without decoder)
|
||||
|
||||
Pretrained model on English language using a similar objective objective as [ELECTRA](https://huggingface.co/transformers/model_doc/electra.html). It was introduced in
|
||||
[this paper](https://arxiv.org/pdf/2006.03236.pdf) and first released in
|
||||
[this repository](https://github.com/laiguokun/Funnel-Transformer). This model is uncased: it does not make a difference
|
||||
between english and English.
|
||||
|
||||
Disclaimer: The team releasing Funnel Transformer did not write a model card for this model so this model card has been
|
||||
written by the Hugging Face team.
|
||||
|
||||
## Model description
|
||||
|
||||
Funnel Transformer is a transformers model pretrained on a large corpus of English data in a self-supervised fashion. This means it
|
||||
was pretrained on the raw texts only, with no humans labelling them in any way (which is why it can use lots of
|
||||
publicly available data) with an automatic process to generate inputs and labels from those texts.
|
||||
|
||||
More precisely, a small language model corrupts the input texts and serves as a generator of inputs for this model, and
|
||||
the pretraining objective is to predict which token is an original and which one has been replaced, a bit like a GAN training.
|
||||
|
||||
This way, the model learns an inner representation of the English language that can then be used to extract features
|
||||
useful for downstream tasks: if you have a dataset of labeled sentences for instance, you can train a standard
|
||||
classifier using the features produced by the BERT model as inputs.
|
||||
|
||||
**Note:** This model does not contain the decoder, so it ouputs hidden states that have a sequence length of one fourth
|
||||
of the inputs. It's good to use for tasks requiring a summary of the sentence (like sentence classification) but not if
|
||||
you need one input per initial token. You should use the `small` model in that case.
|
||||
|
||||
## Intended uses & limitations
|
||||
|
||||
You can use the raw model to extract a vector representation of a given text, but it's mostly intended to
|
||||
be fine-tuned on a downstream task. See the [model hub](https://huggingface.co/models?filter=funnel-transformer) to look for
|
||||
fine-tuned versions on a task that interests you.
|
||||
|
||||
Note that this model is primarily aimed at being fine-tuned on tasks that use the whole sentence (potentially masked)
|
||||
to make decisions, such as sequence classification, token classification or question answering. For tasks such as text
|
||||
generation you should look at model like GPT2.
|
||||
|
||||
### How to use
|
||||
|
||||
|
||||
Here is how to use this model to get the features of a given text in PyTorch:
|
||||
|
||||
```python
|
||||
from transformers import FunnelTokenizer, FunnelBaseModel
|
||||
tokenizer = FunnelTokenizer.from_pretrained("funnel-transformer/small-base")
|
||||
model = FunnelBaseModel.from_pretrained("funnel-transformer/small-base")
|
||||
text = "Replace me by any text you'd like."
|
||||
encoded_input = tokenizer(text, return_tensors='pt')
|
||||
output = model(**encoded_input)
|
||||
```
|
||||
|
||||
and in TensorFlow:
|
||||
|
||||
```python
|
||||
from transformers import FunnelTokenizer, TFFunnelBaseModel
|
||||
tokenizer = FunnelTokenizer.from_pretrained("funnel-transformer/small-base")
|
||||
model = TFFunnelBaseModel.from_pretrained("funnel-transformer/small-base")
|
||||
text = "Replace me by any text you'd like."
|
||||
encoded_input = tokenizer(text, return_tensors='tf')
|
||||
output = model(encoded_input)
|
||||
```
|
||||
|
||||
## Training data
|
||||
|
||||
The BERT model was pretrained on:
|
||||
- [BookCorpus](https://yknzhu.wixsite.com/mbweb), a dataset consisting of 11,038 unpublished books,
|
||||
- [English Wikipedia](https://en.wikipedia.org/wiki/English_Wikipedia) (excluding lists, tables and headers),
|
||||
- [Clue Web](https://lemurproject.org/clueweb12/), a dataset of 733,019,372 English web pages,
|
||||
- [GigaWord](https://catalog.ldc.upenn.edu/LDC2011T07), an archive of newswire text data,
|
||||
- [Common Crawl](https://commoncrawl.org/), a dataset of raw web pages.
|
||||
|
||||
|
||||
### BibTeX entry and citation info
|
||||
|
||||
```bibtex
|
||||
@misc{dai2020funneltransformer,
|
||||
title={Funnel-Transformer: Filtering out Sequential Redundancy for Efficient Language Processing},
|
||||
author={Zihang Dai and Guokun Lai and Yiming Yang and Quoc V. Le},
|
||||
year={2020},
|
||||
eprint={2006.03236},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.LG}
|
||||
}
|
||||
```
|
||||
|
||||
@@ -1,90 +0,0 @@
|
||||
---
|
||||
language: en
|
||||
license: apache-2.0
|
||||
datasets:
|
||||
- bookcorpus
|
||||
- wikipedia
|
||||
- gigaword
|
||||
---
|
||||
|
||||
# Funnel Transformer small model (B4-4-4 with decoder)
|
||||
|
||||
Pretrained model on English language using a similar objective objective as [ELECTRA](https://huggingface.co/transformers/model_doc/electra.html). It was introduced in
|
||||
[this paper](https://arxiv.org/pdf/2006.03236.pdf) and first released in
|
||||
[this repository](https://github.com/laiguokun/Funnel-Transformer). This model is uncased: it does not make a difference
|
||||
between english and English.
|
||||
|
||||
Disclaimer: The team releasing Funnel Transformer did not write a model card for this model so this model card has been
|
||||
written by the Hugging Face team.
|
||||
|
||||
## Model description
|
||||
|
||||
Funnel Transformer is a transformers model pretrained on a large corpus of English data in a self-supervised fashion. This means it
|
||||
was pretrained on the raw texts only, with no humans labelling them in any way (which is why it can use lots of
|
||||
publicly available data) with an automatic process to generate inputs and labels from those texts.
|
||||
|
||||
More precisely, a small language model corrupts the input texts and serves as a generator of inputs for this model, and
|
||||
the pretraining objective is to predict which token is an original and which one has been replaced, a bit like a GAN training.
|
||||
|
||||
This way, the model learns an inner representation of the English language that can then be used to extract features
|
||||
useful for downstream tasks: if you have a dataset of labeled sentences for instance, you can train a standard
|
||||
classifier using the features produced by the BERT model as inputs.
|
||||
|
||||
## Intended uses & limitations
|
||||
|
||||
You can use the raw model to extract a vector representation of a given text, but it's mostly intended to
|
||||
be fine-tuned on a downstream task. See the [model hub](https://huggingface.co/models?filter=funnel-transformer) to look for
|
||||
fine-tuned versions on a task that interests you.
|
||||
|
||||
Note that this model is primarily aimed at being fine-tuned on tasks that use the whole sentence (potentially masked)
|
||||
to make decisions, such as sequence classification, token classification or question answering. For tasks such as text
|
||||
generation you should look at model like GPT2.
|
||||
|
||||
### How to use
|
||||
|
||||
|
||||
Here is how to use this model to get the features of a given text in PyTorch:
|
||||
|
||||
```python
|
||||
from transformers import FunnelTokenizer, FunnelModel
|
||||
tokenizer = FunnelTokenizer.from_pretrained("funnel-transformer/small")
|
||||
model = FunneModel.from_pretrained("funnel-transformer/small")
|
||||
text = "Replace me by any text you'd like."
|
||||
encoded_input = tokenizer(text, return_tensors='pt')
|
||||
output = model(**encoded_input)
|
||||
```
|
||||
|
||||
and in TensorFlow:
|
||||
|
||||
```python
|
||||
from transformers import FunnelTokenizer, TFFunnelModel
|
||||
tokenizer = FunnelTokenizer.from_pretrained("funnel-transformer/small")
|
||||
model = TFFunnelModel.from_pretrained("funnel-transformer/small")
|
||||
text = "Replace me by any text you'd like."
|
||||
encoded_input = tokenizer(text, return_tensors='tf')
|
||||
output = model(encoded_input)
|
||||
```
|
||||
|
||||
## Training data
|
||||
|
||||
The BERT model was pretrained on:
|
||||
- [BookCorpus](https://yknzhu.wixsite.com/mbweb), a dataset consisting of 11,038 unpublished books,
|
||||
- [English Wikipedia](https://en.wikipedia.org/wiki/English_Wikipedia) (excluding lists, tables and headers),
|
||||
- [Clue Web](https://lemurproject.org/clueweb12/), a dataset of 733,019,372 English web pages,
|
||||
- [GigaWord](https://catalog.ldc.upenn.edu/LDC2011T07), an archive of newswire text data,
|
||||
- [Common Crawl](https://commoncrawl.org/), a dataset of raw web pages.
|
||||
|
||||
|
||||
### BibTeX entry and citation info
|
||||
|
||||
```bibtex
|
||||
@misc{dai2020funneltransformer,
|
||||
title={Funnel-Transformer: Filtering out Sequential Redundancy for Efficient Language Processing},
|
||||
author={Zihang Dai and Guokun Lai and Yiming Yang and Quoc V. Le},
|
||||
year={2020},
|
||||
eprint={2006.03236},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.LG}
|
||||
}
|
||||
```
|
||||
|
||||
@@ -1,94 +0,0 @@
|
||||
---
|
||||
language: en
|
||||
license: apache-2.0
|
||||
datasets:
|
||||
- bookcorpus
|
||||
- wikipedia
|
||||
- gigaword
|
||||
---
|
||||
|
||||
# Funnel Transformer xlarge model (B10-10-10 without decoder)
|
||||
|
||||
Pretrained model on English language using a similar objective objective as [ELECTRA](https://huggingface.co/transformers/model_doc/electra.html). It was introduced in
|
||||
[this paper](https://arxiv.org/pdf/2006.03236.pdf) and first released in
|
||||
[this repository](https://github.com/laiguokun/Funnel-Transformer). This model is uncased: it does not make a difference
|
||||
between english and English.
|
||||
|
||||
Disclaimer: The team releasing Funnel Transformer did not write a model card for this model so this model card has been
|
||||
written by the Hugging Face team.
|
||||
|
||||
## Model description
|
||||
|
||||
Funnel Transformer is a transformers model pretrained on a large corpus of English data in a self-supervised fashion. This means it
|
||||
was pretrained on the raw texts only, with no humans labelling them in any way (which is why it can use lots of
|
||||
publicly available data) with an automatic process to generate inputs and labels from those texts.
|
||||
|
||||
More precisely, a small language model corrupts the input texts and serves as a generator of inputs for this model, and
|
||||
the pretraining objective is to predict which token is an original and which one has been replaced, a bit like a GAN training.
|
||||
|
||||
This way, the model learns an inner representation of the English language that can then be used to extract features
|
||||
useful for downstream tasks: if you have a dataset of labeled sentences for instance, you can train a standard
|
||||
classifier using the features produced by the BERT model as inputs.
|
||||
|
||||
**Note:** This model does not contain the decoder, so it ouputs hidden states that have a sequence length of one fourth
|
||||
of the inputs. It's good to use for tasks requiring a summary of the sentence (like sentence classification) but not if
|
||||
you need one input per initial token. You should use the `xlarge` model in that case.
|
||||
|
||||
## Intended uses & limitations
|
||||
|
||||
You can use the raw model to extract a vector representation of a given text, but it's mostly intended to
|
||||
be fine-tuned on a downstream task. See the [model hub](https://huggingface.co/models?filter=funnel-transformer) to look for
|
||||
fine-tuned versions on a task that interests you.
|
||||
|
||||
Note that this model is primarily aimed at being fine-tuned on tasks that use the whole sentence (potentially masked)
|
||||
to make decisions, such as sequence classification, token classification or question answering. For tasks such as text
|
||||
generation you should look at model like GPT2.
|
||||
|
||||
### How to use
|
||||
|
||||
|
||||
Here is how to use this model to get the features of a given text in PyTorch:
|
||||
|
||||
```python
|
||||
from transformers import FunnelTokenizer, FunnelBaseModel
|
||||
tokenizer = FunnelTokenizer.from_pretrained("funnel-transformer/xlarge-base")
|
||||
model = FunnelBaseModel.from_pretrained("funnel-transformer/xlarge-base")
|
||||
text = "Replace me by any text you'd like."
|
||||
encoded_input = tokenizer(text, return_tensors='pt')
|
||||
output = model(**encoded_input)
|
||||
```
|
||||
|
||||
and in TensorFlow:
|
||||
|
||||
```python
|
||||
from transformers import FunnelTokenizer, TFFunnelBaseModel
|
||||
tokenizer = FunnelTokenizer.from_pretrained("funnel-transformer/xlarge-base")
|
||||
model = TFFunnelBaseModel.from_pretrained("funnel-transformer/xlarge-base")
|
||||
text = "Replace me by any text you'd like."
|
||||
encoded_input = tokenizer(text, return_tensors='tf')
|
||||
output = model(encoded_input)
|
||||
```
|
||||
|
||||
## Training data
|
||||
|
||||
The BERT model was pretrained on:
|
||||
- [BookCorpus](https://yknzhu.wixsite.com/mbweb), a dataset consisting of 11,038 unpublished books,
|
||||
- [English Wikipedia](https://en.wikipedia.org/wiki/English_Wikipedia) (excluding lists, tables and headers),
|
||||
- [Clue Web](https://lemurproject.org/clueweb12/), a dataset of 733,019,372 English web pages,
|
||||
- [GigaWord](https://catalog.ldc.upenn.edu/LDC2011T07), an archive of newswire text data,
|
||||
- [Common Crawl](https://commoncrawl.org/), a dataset of raw web pages.
|
||||
|
||||
|
||||
### BibTeX entry and citation info
|
||||
|
||||
```bibtex
|
||||
@misc{dai2020funneltransformer,
|
||||
title={Funnel-Transformer: Filtering out Sequential Redundancy for Efficient Language Processing},
|
||||
author={Zihang Dai and Guokun Lai and Yiming Yang and Quoc V. Le},
|
||||
year={2020},
|
||||
eprint={2006.03236},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.LG}
|
||||
}
|
||||
```
|
||||
|
||||
@@ -1,90 +0,0 @@
|
||||
---
|
||||
language: en
|
||||
license: apache-2.0
|
||||
datasets:
|
||||
- bookcorpus
|
||||
- wikipedia
|
||||
- gigaword
|
||||
---
|
||||
|
||||
# Funnel Transformer xlarge model (B10-10-10 with decoder)
|
||||
|
||||
Pretrained model on English language using a similar objective objective as [ELECTRA](https://huggingface.co/transformers/model_doc/electra.html). It was introduced in
|
||||
[this paper](https://arxiv.org/pdf/2006.03236.pdf) and first released in
|
||||
[this repository](https://github.com/laiguokun/Funnel-Transformer). This model is uncased: it does not make a difference
|
||||
between english and English.
|
||||
|
||||
Disclaimer: The team releasing Funnel Transformer did not write a model card for this model so this model card has been
|
||||
written by the Hugging Face team.
|
||||
|
||||
## Model description
|
||||
|
||||
Funnel Transformer is a transformers model pretrained on a large corpus of English data in a self-supervised fashion. This means it
|
||||
was pretrained on the raw texts only, with no humans labelling them in any way (which is why it can use lots of
|
||||
publicly available data) with an automatic process to generate inputs and labels from those texts.
|
||||
|
||||
More precisely, a small language model corrupts the input texts and serves as a generator of inputs for this model, and
|
||||
the pretraining objective is to predict which token is an original and which one has been replaced, a bit like a GAN training.
|
||||
|
||||
This way, the model learns an inner representation of the English language that can then be used to extract features
|
||||
useful for downstream tasks: if you have a dataset of labeled sentences for instance, you can train a standard
|
||||
classifier using the features produced by the BERT model as inputs.
|
||||
|
||||
## Intended uses & limitations
|
||||
|
||||
You can use the raw model to extract a vector representation of a given text, but it's mostly intended to
|
||||
be fine-tuned on a downstream task. See the [model hub](https://huggingface.co/models?filter=funnel-transformer) to look for
|
||||
fine-tuned versions on a task that interests you.
|
||||
|
||||
Note that this model is primarily aimed at being fine-tuned on tasks that use the whole sentence (potentially masked)
|
||||
to make decisions, such as sequence classification, token classification or question answering. For tasks such as text
|
||||
generation you should look at model like GPT2.
|
||||
|
||||
### How to use
|
||||
|
||||
|
||||
Here is how to use this model to get the features of a given text in PyTorch:
|
||||
|
||||
```python
|
||||
from transformers import FunnelTokenizer, FunnelModel
|
||||
tokenizer = FunnelTokenizer.from_pretrained("funnel-transformer/xlarge")
|
||||
model = FunneModel.from_pretrained("funnel-transformer/xlarge")
|
||||
text = "Replace me by any text you'd like."
|
||||
encoded_input = tokenizer(text, return_tensors='pt')
|
||||
output = model(**encoded_input)
|
||||
```
|
||||
|
||||
and in TensorFlow:
|
||||
|
||||
```python
|
||||
from transformers import FunnelTokenizer, TFFunnelModel
|
||||
tokenizer = FunnelTokenizer.from_pretrained("funnel-transformer/xlarge")
|
||||
model = TFFunnelModel.from_pretrained("funnel-transformer/xlarge")
|
||||
text = "Replace me by any text you'd like."
|
||||
encoded_input = tokenizer(text, return_tensors='tf')
|
||||
output = model(encoded_input)
|
||||
```
|
||||
|
||||
## Training data
|
||||
|
||||
The BERT model was pretrained on:
|
||||
- [BookCorpus](https://yknzhu.wixsite.com/mbweb), a dataset consisting of 11,038 unpublished books,
|
||||
- [English Wikipedia](https://en.wikipedia.org/wiki/English_Wikipedia) (excluding lists, tables and headers),
|
||||
- [Clue Web](https://lemurproject.org/clueweb12/), a dataset of 733,019,372 English web pages,
|
||||
- [GigaWord](https://catalog.ldc.upenn.edu/LDC2011T07), an archive of newswire text data,
|
||||
- [Common Crawl](https://commoncrawl.org/), a dataset of raw web pages.
|
||||
|
||||
|
||||
### BibTeX entry and citation info
|
||||
|
||||
```bibtex
|
||||
@misc{dai2020funneltransformer,
|
||||
title={Funnel-Transformer: Filtering out Sequential Redundancy for Efficient Language Processing},
|
||||
author={Zihang Dai and Guokun Lai and Yiming Yang and Quoc V. Le},
|
||||
year={2020},
|
||||
eprint={2006.03236},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.LG}
|
||||
}
|
||||
```
|
||||
|
||||
@@ -7,66 +7,73 @@ tags:
|
||||
- commoncrawl
|
||||
- uncased
|
||||
- umlaute
|
||||
- umlauts
|
||||
- german
|
||||
- deutsch
|
||||
---
|
||||
|
||||
# German Electra Uncased
|
||||
<img width="300px" src="https://raw.githubusercontent.com/German-NLP-Group/german-transformer-training/master/model_cards/german-electra-logo.png">
|
||||
[¹]
|
||||
|
||||
|
||||
# Model Info
|
||||
|
||||
This Model is suitable for Training on many downstream tasks in German (Q&A, Sentiment Analysis, etc.).
|
||||
|
||||
It can be used as a drop-in Replacement for **BERT** in most down-stream tasks (**ELECTRA** is even implemented as an extended **BERT** Class).
|
||||
|
||||
At the time of release (August 2020) this Model is the best performing publicly available German NLP Model on various German Evaluation Metrics (CONLL03-DE, GermEval18 Coarse, GermEval18 Fine). For GermEval18 Coarse results see below. More will be published soon.
|
||||
|
||||
# Installation
|
||||
This model has the special feature that it is **uncased** but does **not strip accents**.
|
||||
This possibility was added by us with [PR #6280](https://github.com/huggingface/transformers/pull/6280).
|
||||
To use it you have to use Transformers version 3.1.0 or newer.
|
||||
|
||||
```bash
|
||||
pip install transformers -U
|
||||
```
|
||||
## Installation
|
||||
|
||||
# Uncase and Umlauts ('Ö', 'Ä', 'Ü')
|
||||
---
|
||||
This model is **uncased** but does not use **strip accents**.
|
||||
The necessary parameter is `strip_accents=False`.
|
||||
|
||||
This needs to be set for the tokenizer otherwise the model will perform slightly worse.
|
||||
It was added to Transformers with [PR #6280](https://github.com/huggingface/transformers/pull/6280).
|
||||
|
||||
Since Transformers has not been released since the PR #6280 was merged, you have to install directly from source:
|
||||
|
||||
`pip install git+https://github.com/huggingface/transformers.git -U`
|
||||
|
||||
---
|
||||
|
||||
|
||||
## Uncase and Umlauts ('Ö', 'Ä', 'Ü')
|
||||
This model is uncased. This helps especially for domains where colloquial terms with uncorrect capitalization is often used.
|
||||
|
||||
The special characters 'ö', 'ü', 'ä' are included through the `strip_accent=False` option, as this leads to an improved precision.
|
||||
|
||||
# Creators
|
||||
## Creators
|
||||
This model was trained and open sourced in conjunction with the [**German NLP Group**](https://github.com/German-NLP-Group) in equal parts by:
|
||||
- [**Philip May**](https://eniak.de) - [T-Systems on site services GmbH](https://www.t-systems-onsite.de/)
|
||||
- [**Philipp Reißel**](https://www.reissel.eu) - [ambeRoad](https://amberoad.de/)
|
||||
|
||||
# Evaluation: GermEval18 Coarse
|
||||
## Evaluation: GermEval18 Coarse
|
||||
|
||||
| Model Name | F1 macro<br/>Mean | F1 macro<br/>Median | F1 macro<br/>Std |
|
||||
| Model Name |</br>F1 macro<br/> Mean | </br>F1 macro<br/>Median | </br>F1 macro<br/>Std |
|
||||
|---|---|---|---|
|
||||
| dbmdz-bert-base-german-europeana-cased | 0.727 | 0.729 | 0.00674 |
|
||||
| dbmdz-bert-base-german-europeana-uncased | 0.736 | 0.737 | 0.00476 |
|
||||
| dbmdz/electra-base-german-europeana-cased-discriminator | 0.745 | 0.745 | 0.00498 |
|
||||
| distilbert-base-german-cased | 0.752 | 0.752 | 0.00341 |
|
||||
| bert-base-german-cased | 0.762 | 0.761 | 0.00597 |
|
||||
| dbmdz/bert-base-german-cased | 0.765 | 0.765 | 0.00523 |
|
||||
| dbmdz/bert-base-german-uncased | 0.770 | 0.770 | 0.00572 |
|
||||
| **ELECTRA-base-german-uncased (this model)** | **0.778** | **0.778** | **0.00392** |
|
||||
| <span style="color:red">**ELECTRA-base-german-uncased** (this model) | <span style="color:red">**0.778** | **0.778** | **0.00392** |
|
||||
| dbmdz/bert-base-german-uncased | 0.770 | 0.770 | 0.00572 |
|
||||
| dbmdz/bert-base-german-cased | 0.765 | 0.765 | 0.00523 |
|
||||
| bert-base-german-cased | 0.762 | 0.761 | 0.00597 |
|
||||
| distilbert-base-german-cased | 0.752 | 0.752 | 0.00341 |
|
||||
| dbmdz/electra-base-german-europeana-cased-discriminator | 0.745 | 0.745 | 0.00498 |
|
||||
| dbmdz-bert-base-german-europeana-uncased | 0.736 | 0.737 | 0.00476 |
|
||||
| dbmdz-bert-base-german-europeana-cased | 0.727 | 0.729 | 0.00674 |
|
||||
|
||||
- (1): Hyperparameters taken from the [FARM project](https://farm.deepset.ai/) "[germEval18Coarse_config.json](https://github.com/deepset-ai/FARM/blob/master/experiments/german-bert2.0-eval/germEval18Coarse_config.json)"
|
||||
|
||||

|
||||
|
||||
# Checkpoint evaluation
|
||||
## Checkpoint evaluation
|
||||
Since it it not guaranteed that the last checkpoint is the best, we evaluated the checkpoints on GermEval18. We found that the last checkpoint is indeed the best. The training was stable and did not overfit the text corpus. Below is a boxplot chart showing the different checkpoints.
|
||||
|
||||

|
||||
|
||||
# Pre-training details
|
||||
## Pre-training details
|
||||
|
||||
## Data
|
||||
### Data
|
||||
- Cleaned Common Crawl Corpus 2019-09 German: [CC_net](https://github.com/facebookresearch/cc_net) (Only head coprus and filtered for language_score > 0.98) - 62 GB
|
||||
- German Wikipedia Article Pages Dump (20200701) - 5.5 GB
|
||||
- German Wikipedia Talk Pages Dump (20200620) - 1.1 GB
|
||||
@@ -77,7 +84,7 @@ The sentences were split with [SojaMo](https://github.com/tsproisl/SoMaJo). We t
|
||||
|
||||
More Details can be found here [Preperaing Datasets for German Electra Github](https://github.com/German-NLP-Group/german-transformer-training)
|
||||
|
||||
## Electra Branch no_strip_accents
|
||||
### Electra Branch no_strip_accents
|
||||
Because we do not want to stip accents in our training data we made a change to Electra and used this repo [Electra no_strip_accents](https://github.com/PhilipMay/electra/tree/no_strip_accents) (branch `no_strip_accents`). Then created the tf dataset with:
|
||||
|
||||
```bash
|
||||
@@ -85,11 +92,13 @@ python build_pretraining_dataset.py --corpus-dir <corpus_dir> --vocab-file <dir>
|
||||
```
|
||||
|
||||
## The training
|
||||
|
||||
The training itself can be performed with the Original Electra Repo (No special case for this needed).
|
||||
We run it with the following Config:
|
||||
|
||||
|
||||
<details>
|
||||
<summary>The exact Training Config</summary>
|
||||
<summary>The exact Training Config</summary>
|
||||
<br/>debug False
|
||||
<br/>disallow_correct False
|
||||
<br/>disc_weight 50.0
|
||||
@@ -134,6 +143,7 @@ We run it with the following Config:
|
||||
<br/>vocab_file gs://XXX
|
||||
<br/>vocab_size 32767
|
||||
<br/>weight_decay_rate 0.01
|
||||
|
||||
</details>
|
||||
|
||||

|
||||
@@ -145,7 +155,7 @@ Special thanks to [Stefan Schweter](https://github.com/stefan-it) for your feedb
|
||||
|
||||
[¹]: Source for the picture [Pinterest](https://www.pinterest.cl/pin/371828512984142193/)
|
||||
|
||||
# Negative Results
|
||||
## Negative Results
|
||||
We tried the following approaches which we found had no positive influence:
|
||||
|
||||
- **Increased Vocab Size**: Leads to more parameters and thus reduced examples/sec while no visible Performance gains were measured
|
||||
|
||||
@@ -1,47 +0,0 @@
|
||||
---
|
||||
language: en
|
||||
thumbnail:
|
||||
tags:
|
||||
- bert
|
||||
- embeddings
|
||||
license: Apache-2.0
|
||||
---
|
||||
|
||||
# LABSE BERT
|
||||
|
||||
## Model description
|
||||
|
||||
Model for "Language-agnostic BERT Sentence Embedding" paper from Fangxiaoyu Feng, Yinfei Yang, Daniel Cer, Naveen Arivazhagan, Wei Wang. Model available in [TensorFlow Hub](https://tfhub.dev/google/LaBSE/1).
|
||||
|
||||
## Intended uses & limitations
|
||||
|
||||
#### How to use
|
||||
|
||||
```python
|
||||
from transformers import AutoTokenizer, AutoModel
|
||||
import torch
|
||||
|
||||
# from sentence-transformers
|
||||
def mean_pooling(model_output, attention_mask):
|
||||
token_embeddings = model_output[0] #First element of model_output contains all token embeddings
|
||||
input_mask_expanded = attention_mask.unsqueeze(-1).expand(token_embeddings.size()).float()
|
||||
sum_embeddings = torch.sum(token_embeddings * input_mask_expanded, 1)
|
||||
sum_mask = torch.clamp(input_mask_expanded.sum(1), min=1e-9)
|
||||
return sum_embeddings / sum_mask
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained("pvl/labse_bert", do_lower_case=False)
|
||||
model = AutoModel.from_pretrained("pvl/labse_bert")
|
||||
|
||||
sentences = ['This framework generates embeddings for each input sentence',
|
||||
'Sentences are passed as a list of string.',
|
||||
'The quick brown fox jumps over the lazy dog.']
|
||||
|
||||
encoded_input = tokenizer(sentences, padding=True, truncation=True, max_length=128, return_tensors='pt')
|
||||
|
||||
with torch.no_grad():
|
||||
model_output = model(**encoded_input)
|
||||
|
||||
sentence_embeddings = mean_pooling(model_output, encoded_input['attention_mask'])
|
||||
|
||||
|
||||
```
|
||||
@@ -1,53 +0,0 @@
|
||||
# Pegasus for Paraphrasing
|
||||
Pegasus model fine-tuned for paraphrasing
|
||||
|
||||
## Model in Action 🚀
|
||||
```
|
||||
import torch
|
||||
from transformers import PegasusForConditionalGeneration, PegasusTokenizer
|
||||
model_name = 'tuner007/pegasus_paraphrase'
|
||||
torch_device = 'cuda' if torch.cuda.is_available() else 'cpu'
|
||||
tokenizer = PegasusTokenizer.from_pretrained(model_name)
|
||||
model = PegasusForConditionalGeneration.from_pretrained(model_name).to(torch_device)
|
||||
|
||||
def get_response(input_text,num_return_sequences):
|
||||
batch = tokenizer.prepare_seq2seq_batch([input_text],truncation=True,padding='longest',max_length=60).to(torch_device)
|
||||
translated = model.generate(**batch,max_length=60,num_beams=10, num_return_sequences=num_return_sequences, temperature=1.5)
|
||||
tgt_text = tokenizer.batch_decode(translated, skip_special_tokens=True)
|
||||
return tgt_text
|
||||
```
|
||||
#### Example 1:
|
||||
```
|
||||
context = "The ultimate test of your knowledge is your capacity to convey it to another."
|
||||
get_response(context,10)
|
||||
# output:
|
||||
['The test of your knowledge is your ability to convey it.',
|
||||
'The ability to convey your knowledge is the ultimate test of your knowledge.',
|
||||
'The ability to convey your knowledge is the most important test of your knowledge.',
|
||||
'Your capacity to convey your knowledge is the ultimate test of it.',
|
||||
'The test of your knowledge is your ability to communicate it.',
|
||||
'Your capacity to convey your knowledge is the ultimate test of your knowledge.',
|
||||
'Your capacity to convey your knowledge to another is the ultimate test of your knowledge.',
|
||||
'Your capacity to convey your knowledge is the most important test of your knowledge.',
|
||||
'The test of your knowledge is how well you can convey it.',
|
||||
'Your capacity to convey your knowledge is the ultimate test.']
|
||||
```
|
||||
#### Example 2: Question paraphrasing (was not trained on quora dataset)
|
||||
```
|
||||
context = "Which course should I take to get started in data science?"
|
||||
get_response(context,10)
|
||||
# output:
|
||||
['Which data science course should I take?',
|
||||
'Which data science course should I take first?',
|
||||
'Should I take a data science course?',
|
||||
'Which data science class should I take?',
|
||||
'Which data science course should I attend?',
|
||||
'I want to get started in data science.',
|
||||
'Which data science course should I enroll in?',
|
||||
'Which data science course is right for me?',
|
||||
'Which data science course is best for me?',
|
||||
'Which course should I take to get started?']
|
||||
```
|
||||
|
||||
> Created by Arpit Rajauria
|
||||
[](https://twitter.com/arpit_rajauria)
|
||||
@@ -44,4 +44,3 @@ Pull Request so it can be included under the Community notebooks.
|
||||
|[Expand and Fine Tune Sci-BERT](https://github.com/lordtt13/word-embeddings/blob/master/COVID-19%20Research%20Data/COVID-SciBERT.ipynb)| How to increase vocabulary of a pretrained SciBERT model from AllenAI on the CORD dataset and pipeline it. | [Tanmay Thakur](https://github.com/lordtt13) | [](https://colab.research.google.com/drive/1rqAR40goxbAfez1xvF3hBJphSCsvXmh8)|
|
||||
|[Fine-tune Electra and interpret with Integrated Gradients](https://github.com/elsanns/xai-nlp-notebooks/blob/master/electra_fine_tune_interpret_captum_ig.ipynb) | How to fine-tune Electra for sentiment analysis and interpret predictions with Captum Integrated Gradients | [Eliza Szczechla](https://elsanns.github.io) | [](https://colab.research.google.com/github/elsanns/xai-nlp-notebooks/blob/master/electra_fine_tune_interpret_captum_ig.ipynb)|
|
||||
|[fine-tune a non-English GPT-2 Model with Trainer class](https://github.com/philschmid/fine-tune-GPT-2/blob/master/Fine_tune_a_non_English_GPT_2_Model_with_Huggingface.ipynb) | How to fine-tune a non-English GPT-2 Model with Trainer class | [Philipp Schmid](https://www.philschmid.de) | [](https://colab.research.google.com/github/philschmid/fine-tune-GPT-2/blob/master/Fine_tune_a_non_English_GPT_2_Model_with_Huggingface.ipynb)|
|
||||
|[Fine-tune a DistilBERT Model for Multi Label Classification task](https://github.com/DhavalTaunk08/Transformers_scripts/blob/master/Transformers_multilabel_distilbert.ipynb) | How to fine-tune a DistilBERT Model for Multi Label Classification task | [Dhaval Taunk](https://github.com/DhavalTaunk08) | [](https://colab.research.google.com/github/DhavalTaunk08/Transformers_scripts/blob/master/Transformers_multilabel_distilbert.ipynb)|
|
||||
@@ -1,62 +0,0 @@
|
||||
#/usr/bin/env bash
|
||||
|
||||
# this script acquires data and converts it to fsmt model
|
||||
# it covers:
|
||||
# - allenai/wmt16-en-de-dist-12-1
|
||||
# - allenai/wmt16-en-de-dist-6-1
|
||||
# - allenai/wmt16-en-de-12-1
|
||||
|
||||
# this script needs to be run from the top level of the transformers repo
|
||||
if [ ! -d "src/transformers" ]; then
|
||||
echo "Error: This script needs to be run from the top of the transformers repo"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
mkdir data
|
||||
|
||||
# get data (run once)
|
||||
|
||||
cd data
|
||||
gdown 'https://drive.google.com/uc?id=1x_G2cjvM1nW5hjAB8-vWxRqtQTlmIaQU'
|
||||
gdown 'https://drive.google.com/uc?id=1oA2aqZlVNj5FarxBlNXEHpBS4lRetTzU'
|
||||
gdown 'https://drive.google.com/uc?id=1Wup2D318QYBFPW_NKI1mfP_hXOfmUI9r'
|
||||
tar -xvzf trans_ende_12-1_0.2.tar.gz
|
||||
tar -xvzf trans_ende-dist_12-1_0.2.tar.gz
|
||||
tar -xvzf trans_ende-dist_6-1_0.2.tar.gz
|
||||
gdown 'https://drive.google.com/uc?id=1mNufoynJ9-Zy1kJh2TA_lHm2squji0i9'
|
||||
gdown 'https://drive.google.com/uc?id=1iO7um-HWoNoRKDtw27YUSgyeubn9uXqj'
|
||||
tar -xvzf wmt16.en-de.deep-shallow.dist.tar.gz
|
||||
tar -xvzf wmt16.en-de.deep-shallow.tar.gz
|
||||
cp wmt16.en-de.deep-shallow/data-bin/dict.*.txt trans_ende_12-1_0.2
|
||||
cp wmt16.en-de.deep-shallow.dist/data-bin/dict.*.txt trans_ende-dist_12-1_0.2
|
||||
cp wmt16.en-de.deep-shallow.dist/data-bin/dict.*.txt trans_ende-dist_6-1_0.2
|
||||
cp wmt16.en-de.deep-shallow/bpecodes trans_ende_12-1_0.2
|
||||
cp wmt16.en-de.deep-shallow.dist/bpecodes trans_ende-dist_12-1_0.2
|
||||
cp wmt16.en-de.deep-shallow.dist/bpecodes trans_ende-dist_6-1_0.2
|
||||
cd -
|
||||
|
||||
# run conversions and uploads
|
||||
|
||||
PYTHONPATH="src" python src/transformers/convert_fsmt_original_pytorch_checkpoint_to_pytorch.py --fsmt_checkpoint_path data/trans_ende-dist_12-1_0.2/checkpoint_top5_average.pt --pytorch_dump_folder_path data/wmt16-en-de-dist-12-1
|
||||
|
||||
PYTHONPATH="src" python src/transformers/convert_fsmt_original_pytorch_checkpoint_to_pytorch.py --fsmt_checkpoint_path data/trans_ende-dist_6-1_0.2/checkpoint_top5_average.pt --pytorch_dump_folder_path data/wmt16-en-de-dist-6-1
|
||||
|
||||
PYTHONPATH="src" python src/transformers/convert_fsmt_original_pytorch_checkpoint_to_pytorch.py --fsmt_checkpoint_path data/trans_ende_12-1_0.2/checkpoint_top5_average.pt --pytorch_dump_folder_path data/wmt16-en-de-12-1
|
||||
|
||||
|
||||
# upload
|
||||
cd data
|
||||
transformers-cli upload -y wmt16-en-de-dist-12-1
|
||||
transformers-cli upload -y wmt16-en-de-dist-6-1
|
||||
transformers-cli upload -y wmt16-en-de-12-1
|
||||
cd -
|
||||
|
||||
|
||||
# if updating just small files and not the large models, here is a script to generate the right commands:
|
||||
perl -le 'for $f (@ARGV) { print qq[transformers-cli upload -y $_/$f --filename $_/$f] for ("wmt16-en-de-dist-12-1", "wmt16-en-de-dist-6-1", "wmt16-en-de-12-1")}' vocab-src.json vocab-tgt.json tokenizer_config.json config.json
|
||||
# add/remove files as needed
|
||||
|
||||
# Caching note: Unfortunately due to CDN caching the uploaded model may be unavailable for up to 24hs after upload
|
||||
# So the only way to start using the new model sooner is either:
|
||||
# 1. download it to a local path and use that path as model_name
|
||||
# 2. make sure you use: from_pretrained(..., use_cdn=False) everywhere
|
||||
@@ -1,50 +0,0 @@
|
||||
#/usr/bin/env bash
|
||||
|
||||
# this script acquires data and converts it to fsmt model
|
||||
# it covers:
|
||||
# - allenai/wmt19-de-en-6-6-base
|
||||
# - allenai/wmt19-de-en-6-6-big
|
||||
|
||||
# this script needs to be run from the top level of the transformers repo
|
||||
if [ ! -d "src/transformers" ]; then
|
||||
echo "Error: This script needs to be run from the top of the transformers repo"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
mkdir data
|
||||
|
||||
# get data (run once)
|
||||
|
||||
cd data
|
||||
gdown 'https://drive.google.com/uc?id=1j6z9fYdlUyOYsh7KJoumRlr1yHczxR5T'
|
||||
gdown 'https://drive.google.com/uc?id=1yT7ZjqfvUYOBXvMjeY8uGRHQFWoSo8Q5'
|
||||
gdown 'https://drive.google.com/uc?id=15gAzHeRUCs-QV8vHeTReMPEh1j8excNE'
|
||||
tar -xvzf wmt19.de-en.tar.gz
|
||||
tar -xvzf wmt19_deen_base_dr0.1_1.tar.gz
|
||||
tar -xvzf wmt19_deen_big_dr0.1_2.tar.gz
|
||||
cp wmt19.de-en/data-bin/dict.*.txt wmt19_deen_base_dr0.1_1
|
||||
cp wmt19.de-en/data-bin/dict.*.txt wmt19_deen_big_dr0.1_2
|
||||
cd -
|
||||
|
||||
# run conversions and uploads
|
||||
|
||||
PYTHONPATH="src" python src/transformers/convert_fsmt_original_pytorch_checkpoint_to_pytorch.py --fsmt_checkpoint_path data/wmt19_deen_base_dr0.1_1/checkpoint_last3_avg.pt --pytorch_dump_folder_path data/wmt19-de-en-6-6-base
|
||||
|
||||
PYTHONPATH="src" python src/transformers/convert_fsmt_original_pytorch_checkpoint_to_pytorch.py --fsmt_checkpoint_path data/wmt19_deen_big_dr0.1_2/checkpoint_last3_avg.pt --pytorch_dump_folder_path data/wmt19-de-en-6-6-big
|
||||
|
||||
|
||||
# upload
|
||||
cd data
|
||||
transformers-cli upload -y wmt19-de-en-6-6-base
|
||||
transformers-cli upload -y wmt19-de-en-6-6-big
|
||||
cd -
|
||||
|
||||
|
||||
# if updating just small files and not the large models, here is a script to generate the right commands:
|
||||
perl -le 'for $f (@ARGV) { print qq[transformers-cli upload -y $_/$f --filename $_/$f] for ("wmt19-de-en-6-6-base", "wmt19-de-en-6-6-big")}' vocab-src.json vocab-tgt.json tokenizer_config.json config.json
|
||||
# add/remove files as needed
|
||||
|
||||
# Caching note: Unfortunately due to CDN caching the uploaded model may be unavailable for up to 24hs after upload
|
||||
# So the only way to start using the new model sooner is either:
|
||||
# 1. download it to a local path and use that path as model_name
|
||||
# 2. make sure you use: from_pretrained(..., use_cdn=False) everywhere
|
||||
@@ -1,61 +0,0 @@
|
||||
#/usr/bin/env bash
|
||||
|
||||
# this script acquires data and converts it to fsmt model
|
||||
# it covers:
|
||||
# - facebook/wmt19-ru-en
|
||||
# - facebook/wmt19-en-ru
|
||||
# - facebook/wmt19-de-en
|
||||
# - facebook/wmt19-en-de
|
||||
|
||||
# this script needs to be run from the top level of the transformers repo
|
||||
if [ ! -d "src/transformers" ]; then
|
||||
echo "Error: This script needs to be run from the top of the transformers repo"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
mkdir data
|
||||
|
||||
# get data (run once)
|
||||
|
||||
cd data
|
||||
wget https://dl.fbaipublicfiles.com/fairseq/models/wmt19.en-de.joined-dict.ensemble.tar.gz
|
||||
wget https://dl.fbaipublicfiles.com/fairseq/models/wmt19.de-en.joined-dict.ensemble.tar.gz
|
||||
wget https://dl.fbaipublicfiles.com/fairseq/models/wmt19.en-ru.ensemble.tar.gz
|
||||
wget https://dl.fbaipublicfiles.com/fairseq/models/wmt19.ru-en.ensemble.tar.gz
|
||||
tar -xvzf wmt19.en-de.joined-dict.ensemble.tar.gz
|
||||
tar -xvzf wmt19.de-en.joined-dict.ensemble.tar.gz
|
||||
tar -xvzf wmt19.en-ru.ensemble.tar.gz
|
||||
tar -xvzf wmt19.ru-en.ensemble.tar.gz
|
||||
cd -
|
||||
|
||||
# run conversions and uploads
|
||||
|
||||
export PAIR=ru-en
|
||||
PYTHONPATH="src" python src/transformers/convert_fsmt_original_pytorch_checkpoint_to_pytorch.py --fsmt_checkpoint_path data/wmt19.$PAIR.ensemble/model4.pt --pytorch_dump_folder_path data/wmt19-$PAIR
|
||||
|
||||
export PAIR=en-ru
|
||||
PYTHONPATH="src" python src/transformers/convert_fsmt_original_pytorch_checkpoint_to_pytorch.py --fsmt_checkpoint_path data/wmt19.$PAIR.ensemble/model4.pt --pytorch_dump_folder_path data/wmt19-$PAIR
|
||||
|
||||
export PAIR=de-en
|
||||
PYTHONPATH="src" python src/transformers/convert_fsmt_original_pytorch_checkpoint_to_pytorch.py --fsmt_checkpoint_path data/wmt19.$PAIR.joined-dict.ensemble/model4.pt --pytorch_dump_folder_path data/wmt19-$PAIR
|
||||
|
||||
export PAIR=en-de
|
||||
PYTHONPATH="src" python src/transformers/convert_fsmt_original_pytorch_checkpoint_to_pytorch.py --fsmt_checkpoint_path data/wmt19.$PAIR.joined-dict.ensemble/model4.pt --pytorch_dump_folder_path data/wmt19-$PAIR
|
||||
|
||||
|
||||
# upload
|
||||
cd data
|
||||
transformers-cli upload -y wmt19-ru-en
|
||||
transformers-cli upload -y wmt19-en-ru
|
||||
transformers-cli upload -y wmt19-de-en
|
||||
transformers-cli upload -y wmt19-en-de
|
||||
cd -
|
||||
|
||||
# if updating just small files and not the large models, here is a script to generate the right commands:
|
||||
perl -le 'for $f (@ARGV) { print qq[transformers-cli upload -y $_/$f --filename $_/$f] for map { "wmt19-$_" } ("en-ru", "ru-en", "de-en", "en-de")}' vocab-src.json vocab-tgt.json tokenizer_config.json config.json
|
||||
# add/remove files as needed
|
||||
|
||||
# Caching note: Unfortunately due to CDN caching the uploaded model may be unavailable for up to 24hs after upload
|
||||
# So the only way to start using the new model sooner is either:
|
||||
# 1. download it to a local path and use that path as model_name
|
||||
# 2. make sure you use: from_pretrained(..., use_cdn=False) everywhere
|
||||
@@ -1,66 +0,0 @@
|
||||
#/usr/bin/env bash
|
||||
|
||||
# this script evals the following fsmt models
|
||||
# it covers:
|
||||
# - allenai/wmt16-en-de-dist-12-1
|
||||
# - allenai/wmt16-en-de-dist-6-1
|
||||
# - allenai/wmt16-en-de-12-1
|
||||
|
||||
# this script needs to be run from the top level of the transformers repo
|
||||
if [ ! -d "src/transformers" ]; then
|
||||
echo "Error: This script needs to be run from the top of the transformers repo"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# In these scripts you may have to lower BS if you get CUDA OOM (or increase it if you have a large GPU)
|
||||
|
||||
### Normal eval ###
|
||||
|
||||
export PAIR=en-de
|
||||
export DATA_DIR=data/$PAIR
|
||||
export SAVE_DIR=data/$PAIR
|
||||
export BS=64
|
||||
export NUM_BEAMS=5
|
||||
mkdir -p $DATA_DIR
|
||||
sacrebleu -t wmt19 -l $PAIR --echo src > $DATA_DIR/val.source
|
||||
sacrebleu -t wmt19 -l $PAIR --echo ref > $DATA_DIR/val.target
|
||||
|
||||
MODEL_PATH=allenai/wmt16-en-de-dist-12-1
|
||||
echo $PAIR $MODEL_PATH
|
||||
PYTHONPATH="src:examples/seq2seq" python examples/seq2seq/run_eval.py $MODEL_PATH $DATA_DIR/val.source $SAVE_DIR/test_translations.txt --reference_path $DATA_DIR/val.target --score_path $SAVE_DIR/test_bleu.json --bs $BS --task translation --num_beams $NUM_BEAMS
|
||||
|
||||
MODEL_PATH=allenai/wmt16-en-de-dist-6-1
|
||||
echo $PAIR $MODEL_PATH
|
||||
PYTHONPATH="src:examples/seq2seq" python examples/seq2seq/run_eval.py $MODEL_PATH $DATA_DIR/val.source $SAVE_DIR/test_translations.txt --reference_path $DATA_DIR/val.target --score_path $SAVE_DIR/test_bleu.json --bs $BS --task translation --num_beams $NUM_BEAMS
|
||||
|
||||
MODEL_PATH=allenai/wmt16-en-de-12-1
|
||||
echo $PAIR $MODEL_PATH
|
||||
PYTHONPATH="src:examples/seq2seq" python examples/seq2seq/run_eval.py $MODEL_PATH $DATA_DIR/val.source $SAVE_DIR/test_translations.txt --reference_path $DATA_DIR/val.target --score_path $SAVE_DIR/test_bleu.json --bs $BS --task translation --num_beams $NUM_BEAMS
|
||||
|
||||
|
||||
|
||||
### Searching hparams eval ###
|
||||
|
||||
|
||||
export PAIR=en-de
|
||||
export DATA_DIR=data/$PAIR
|
||||
export SAVE_DIR=data/$PAIR
|
||||
export BS=32
|
||||
export NUM_BEAMS=5
|
||||
mkdir -p $DATA_DIR
|
||||
sacrebleu -t wmt19 -l $PAIR --echo src > $DATA_DIR/val.source
|
||||
sacrebleu -t wmt19 -l $PAIR --echo ref > $DATA_DIR/val.target
|
||||
|
||||
MODEL_PATH=allenai/wmt16-en-de-dist-12-1
|
||||
echo $PAIR $MODEL_PATH
|
||||
PYTHONPATH="src:examples/seq2seq" python examples/seq2seq/run_eval_search.py $MODEL_PATH $DATA_DIR/val.source $SAVE_DIR/test_translations.txt --reference_path $DATA_DIR/val.target --score_path $SAVE_DIR/test_bleu.json --bs $BS --task translation --search="num_beams=5:10:15 length_penalty=0.6:0.7:0.8:0.9:1.0:1.1"
|
||||
|
||||
|
||||
MODEL_PATH=allenai/wmt16-en-de-dist-6-1
|
||||
echo $PAIR $MODEL_PATH
|
||||
PYTHONPATH="src:examples/seq2seq" python examples/seq2seq/run_eval_search.py $MODEL_PATH $DATA_DIR/val.source $SAVE_DIR/test_translations.txt --reference_path $DATA_DIR/val.target --score_path $SAVE_DIR/test_bleu.json --bs $BS --task translation --search="num_beams=5:10:15 length_penalty=0.6:0.7:0.8:0.9:1.0:1.1"
|
||||
|
||||
|
||||
MODEL_PATH=allenai/wmt16-en-de-12-1
|
||||
echo $PAIR $MODEL_PATH
|
||||
PYTHONPATH="src:examples/seq2seq" python examples/seq2seq/run_eval_search.py $MODEL_PATH $DATA_DIR/val.source $SAVE_DIR/test_translations.txt --reference_path $DATA_DIR/val.target --score_path $SAVE_DIR/test_bleu.json --bs $BS --task translation --search="num_beams=5:10:15 length_penalty=0.6:0.7:0.8:0.9:1.0:1.1"
|
||||
@@ -1,54 +0,0 @@
|
||||
#/usr/bin/env bash
|
||||
|
||||
# this script evals the following fsmt models
|
||||
# it covers:
|
||||
# - allenai/wmt19-de-en-6-6-base
|
||||
# - allenai/wmt19-de-en-6-6-big
|
||||
|
||||
# this script needs to be run from the top level of the transformers repo
|
||||
if [ ! -d "src/transformers" ]; then
|
||||
echo "Error: This script needs to be run from the top of the transformers repo"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# In these scripts you may have to lower BS if you get CUDA OOM (or increase it if you have a large GPU)
|
||||
|
||||
### Normal eval ###
|
||||
|
||||
export PAIR=de-en
|
||||
export DATA_DIR=data/$PAIR
|
||||
export SAVE_DIR=data/$PAIR
|
||||
export BS=64
|
||||
export NUM_BEAMS=5
|
||||
mkdir -p $DATA_DIR
|
||||
sacrebleu -t wmt19 -l $PAIR --echo src > $DATA_DIR/val.source
|
||||
sacrebleu -t wmt19 -l $PAIR --echo ref > $DATA_DIR/val.target
|
||||
|
||||
MODEL_PATH=allenai/wmt19-de-en-6-6-base
|
||||
echo $PAIR $MODEL_PATH
|
||||
PYTHONPATH="src:examples/seq2seq" python examples/seq2seq/run_eval.py $MODEL_PATH $DATA_DIR/val.source $SAVE_DIR/test_translations.txt --reference_path $DATA_DIR/val.target --score_path $SAVE_DIR/test_bleu.json --bs $BS --task translation --num_beams $NUM_BEAMS
|
||||
|
||||
MODEL_PATH=allenai/wmt19-de-en-6-6-big
|
||||
echo $PAIR $MODEL_PATH
|
||||
PYTHONPATH="src:examples/seq2seq" python examples/seq2seq/run_eval.py $MODEL_PATH $DATA_DIR/val.source $SAVE_DIR/test_translations.txt --reference_path $DATA_DIR/val.target --score_path $SAVE_DIR/test_bleu.json --bs $BS --task translation --num_beams $NUM_BEAMS
|
||||
|
||||
|
||||
|
||||
### Searching hparams eval ###
|
||||
|
||||
export PAIR=de-en
|
||||
export DATA_DIR=data/$PAIR
|
||||
export SAVE_DIR=data/$PAIR
|
||||
export BS=16
|
||||
export NUM_BEAMS=5
|
||||
mkdir -p $DATA_DIR
|
||||
sacrebleu -t wmt19 -l $PAIR --echo src > $DATA_DIR/val.source
|
||||
sacrebleu -t wmt19 -l $PAIR --echo ref > $DATA_DIR/val.target
|
||||
|
||||
MODEL_PATH=allenai/wmt19-de-en-6-6-base
|
||||
echo $PAIR $MODEL_PATH
|
||||
PYTHONPATH="src:examples/seq2seq" python examples/seq2seq/run_eval_search.py $MODEL_PATH $DATA_DIR/val.source $SAVE_DIR/test_translations.txt --reference_path $DATA_DIR/val.target --score_path $SAVE_DIR/test_bleu.json --bs $BS --task translation --search="num_beams=5:10:15 length_penalty=0.6:0.7:0.8:0.9:1.0:1.1"
|
||||
|
||||
MODEL_PATH=allenai/wmt19-de-en-6-6-big
|
||||
echo $PAIR $MODEL_PATH
|
||||
PYTHONPATH="src:examples/seq2seq" python examples/seq2seq/run_eval_search.py $MODEL_PATH $DATA_DIR/val.source $SAVE_DIR/test_translations.txt --reference_path $DATA_DIR/val.target --score_path $SAVE_DIR/test_bleu.json --bs $BS --task translation --search="num_beams=5:10:15 length_penalty=0.6:0.7:0.8:0.9:1.0:1.1"
|
||||
@@ -1,148 +0,0 @@
|
||||
#/usr/bin/env bash
|
||||
|
||||
# this script evals the following fsmt models
|
||||
# it covers:
|
||||
# - facebook/wmt19-ru-en
|
||||
# - facebook/wmt19-en-ru
|
||||
# - facebook/wmt19-de-en
|
||||
# - facebook/wmt19-en-de
|
||||
|
||||
|
||||
# this script needs to be run from the top level of the transformers repo
|
||||
if [ ! -d "src/transformers" ]; then
|
||||
echo "Error: This script needs to be run from the top of the transformers repo"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
|
||||
# In these scripts you may have to lower BS if you get CUDA OOM (or increase it if you have a large GPU)
|
||||
|
||||
### a short estimate version for quick testing ###
|
||||
|
||||
export PAIR=en-ru
|
||||
export DATA_DIR=data/$PAIR
|
||||
export SAVE_DIR=data/$PAIR
|
||||
export BS=8
|
||||
export NUM_BEAMS=8
|
||||
mkdir -p $DATA_DIR
|
||||
sacrebleu -t wmt19 -l $PAIR --echo src | head -10 > $DATA_DIR/val.source
|
||||
sacrebleu -t wmt19 -l $PAIR --echo ref | head -10 > $DATA_DIR/val.target
|
||||
echo $PAIR
|
||||
PYTHONPATH="src:examples/seq2seq" python examples/seq2seq/run_eval.py facebook/wmt19-$PAIR $DATA_DIR/val.source $SAVE_DIR/test_translations.txt --reference_path $DATA_DIR/val.target --score_path $SAVE_DIR/test_bleu.json --bs $BS --task translation --num_beams $NUM_BEAMS
|
||||
|
||||
|
||||
|
||||
### Normal eval ###
|
||||
|
||||
# ru-en
|
||||
|
||||
export PAIR=ru-en
|
||||
export DATA_DIR=data/$PAIR
|
||||
export SAVE_DIR=data/$PAIR
|
||||
export BS=8
|
||||
export NUM_BEAMS=50
|
||||
mkdir -p $DATA_DIR
|
||||
sacrebleu -t wmt19 -l $PAIR --echo src > $DATA_DIR/val.source
|
||||
sacrebleu -t wmt19 -l $PAIR --echo ref > $DATA_DIR/val.target
|
||||
PYTHONPATH="src:examples/seq2seq" python examples/seq2seq/run_eval.py facebook/wmt19-$PAIR $DATA_DIR/val.source $SAVE_DIR/test_translations.txt --reference_path $DATA_DIR/val.target --score_path $SAVE_DIR/test_bleu.json --bs $BS --task translation --num_beams $NUM_BEAMS
|
||||
|
||||
|
||||
# (target BLEU: 41.3 http://matrix.statmt.org/matrix/output/1907?run_id=6937)
|
||||
|
||||
|
||||
# en-ru
|
||||
|
||||
export PAIR=en-ru
|
||||
export DATA_DIR=data/$PAIR
|
||||
export SAVE_DIR=data/$PAIR
|
||||
export BS=8
|
||||
export NUM_BEAMS=50
|
||||
mkdir -p $DATA_DIR
|
||||
sacrebleu -t wmt19 -l $PAIR --echo src > $DATA_DIR/val.source
|
||||
sacrebleu -t wmt19 -l $PAIR --echo ref > $DATA_DIR/val.target
|
||||
echo $PAIR
|
||||
PYTHONPATH="src:examples/seq2seq" python examples/seq2seq/run_eval.py facebook/wmt19-$PAIR $DATA_DIR/val.source $SAVE_DIR/test_translations.txt --reference_path $DATA_DIR/val.target --score_path $SAVE_DIR/test_bleu.json --bs $BS --task translation --num_beams $NUM_BEAMS
|
||||
|
||||
# (target BLEU: 36.4 http://matrix.statmt.org/matrix/output/1914?score_id=37605)
|
||||
|
||||
|
||||
|
||||
# en-de
|
||||
|
||||
export PAIR=en-de
|
||||
export DATA_DIR=data/$PAIR
|
||||
export SAVE_DIR=data/$PAIR
|
||||
export BS=8
|
||||
mkdir -p $DATA_DIR
|
||||
sacrebleu -t wmt19 -l $PAIR --echo src > $DATA_DIR/val.source
|
||||
sacrebleu -t wmt19 -l $PAIR --echo ref > $DATA_DIR/val.target
|
||||
echo $PAIR
|
||||
PYTHONPATH="src:examples/seq2seq" python examples/seq2seq/run_eval.py facebook/wmt19-$PAIR $DATA_DIR/val.source $SAVE_DIR/test_translations.txt --reference_path $DATA_DIR/val.target --score_path $SAVE_DIR/test_bleu.json --bs $BS --task translation --num_beams $NUM_BEAMS
|
||||
|
||||
# (target BLEU: 43.1 http://matrix.statmt.org/matrix/output/1909?run_id=6862)
|
||||
|
||||
|
||||
# de-en
|
||||
|
||||
export PAIR=de-en
|
||||
export DATA_DIR=data/$PAIR
|
||||
export SAVE_DIR=data/$PAIR
|
||||
export BS=8
|
||||
export NUM_BEAMS=50
|
||||
mkdir -p $DATA_DIR
|
||||
sacrebleu -t wmt19 -l $PAIR --echo src > $DATA_DIR/val.source
|
||||
sacrebleu -t wmt19 -l $PAIR --echo ref > $DATA_DIR/val.target
|
||||
echo $PAIR
|
||||
PYTHONPATH="src:examples/seq2seq" python examples/seq2seq/run_eval.py facebook/wmt19-$PAIR $DATA_DIR/val.source $SAVE_DIR/test_translations.txt --reference_path $DATA_DIR/val.target --score_path $SAVE_DIR/test_bleu.json --bs $BS --task translation --num_beams $NUM_BEAMS
|
||||
|
||||
# (target BLEU: 42.3 http://matrix.statmt.org/matrix/output/1902?run_id=6750)
|
||||
|
||||
|
||||
### Searching hparams eval ###
|
||||
|
||||
# en-ru
|
||||
|
||||
export PAIR=ru-en
|
||||
export DATA_DIR=data/$PAIR
|
||||
export SAVE_DIR=data/$PAIR
|
||||
export BS=32
|
||||
mkdir -p $DATA_DIR
|
||||
sacrebleu -t wmt19 -l $PAIR --echo src > $DATA_DIR/val.source
|
||||
sacrebleu -t wmt19 -l $PAIR --echo ref > $DATA_DIR/val.target
|
||||
CUDA_VISIBLE_DEVICES="0" PYTHONPATH="src:examples/seq2seq" python examples/seq2seq/run_eval_search.py facebook/wmt19-$PAIR $DATA_DIR/val.source $SAVE_DIR/test_translations.txt --reference_path $DATA_DIR/val.target --score_path $SAVE_DIR/test_bleu.json --bs $BS --task translation --search="num_beams=5 length_penalty=0.6:0.7:0.8:0.9:1.0:1.1"
|
||||
|
||||
|
||||
# en-ru
|
||||
|
||||
export PAIR=en-ru
|
||||
export DATA_DIR=data/$PAIR
|
||||
export SAVE_DIR=data/$PAIR
|
||||
export BS=16
|
||||
mkdir -p $DATA_DIR
|
||||
mkdir -p $DATA_DIR
|
||||
sacrebleu -t wmt19 -l $PAIR --echo src > $DATA_DIR/val.source
|
||||
sacrebleu -t wmt19 -l $PAIR --echo ref > $DATA_DIR/val.target
|
||||
CUDA_VISIBLE_DEVICES="0" PYTHONPATH="src:examples/seq2seq" python examples/seq2seq/run_eval_search.py facebook/wmt19-$PAIR $DATA_DIR/val.source $SAVE_DIR/test_translations.txt --reference_path $DATA_DIR/val.target --score_path $SAVE_DIR/test_bleu.json --bs $BS --task translation --search="num_beams=5:8:11:15 length_penalty=0.6:0.7:0.8:0.9:1.0:1.1 early_stopping=true:false"
|
||||
|
||||
# en-de
|
||||
|
||||
export PAIR=en-de
|
||||
export DATA_DIR=data/$PAIR
|
||||
export SAVE_DIR=data/$PAIR
|
||||
export BS=16
|
||||
mkdir -p $DATA_DIR
|
||||
sacrebleu -t wmt19 -l $PAIR --echo src > $DATA_DIR/val.source
|
||||
sacrebleu -t wmt19 -l $PAIR --echo ref > $DATA_DIR/val.target
|
||||
CUDA_VISIBLE_DEVICES="1" PYTHONPATH="src:examples/seq2seq" python examples/seq2seq/run_eval_search.py facebook/wmt19-$PAIR $DATA_DIR/val.source $SAVE_DIR/test_translations.txt --reference_path $DATA_DIR/val.target --score_path $SAVE_DIR/test_bleu.json --bs $BS --task translation --search="num_beams=5:8:11:15 length_penalty=0.6:0.7:0.8:0.9:1.0:1.1 early_stopping=true:false"
|
||||
|
||||
# de-en
|
||||
|
||||
export PAIR=de-en
|
||||
export DATA_DIR=data/$PAIR
|
||||
export SAVE_DIR=data/$PAIR
|
||||
export BS=16
|
||||
mkdir -p $DATA_DIR
|
||||
mkdir -p $DATA_DIR
|
||||
sacrebleu -t wmt19 -l $PAIR --echo src > $DATA_DIR/val.source
|
||||
sacrebleu -t wmt19 -l $PAIR --echo ref > $DATA_DIR/val.target
|
||||
CUDA_VISIBLE_DEVICES="1" PYTHONPATH="src:examples/seq2seq" python examples/seq2seq/run_eval_search.py facebook/wmt19-$PAIR $DATA_DIR/val.source $SAVE_DIR/test_translations.txt --reference_path $DATA_DIR/val.target --score_path $SAVE_DIR/test_bleu.json --bs $BS --task translation --search="num_beams=5:8:11:15 length_penalty=0.6:0.7:0.8:0.9:1.0:1.1 early_stopping=true:false"
|
||||
@@ -1,134 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Usage:
|
||||
# ./gen-card-allenai-wmt16.py
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
def write_model_card(model_card_dir, src_lang, tgt_lang, model_name):
|
||||
|
||||
texts = {
|
||||
"en": "Machine learning is great, isn't it?",
|
||||
"ru": "Машинное обучение - это здорово, не так ли?",
|
||||
"de": "Maschinelles Lernen ist großartig, nicht wahr?",
|
||||
}
|
||||
|
||||
# BLUE scores as follows:
|
||||
# "pair": [fairseq, transformers]
|
||||
scores = {
|
||||
"wmt16-en-de-dist-12-1": [28.3, 27.52],
|
||||
"wmt16-en-de-dist-6-1": [27.4, 27.11],
|
||||
"wmt16-en-de-12-1": [26.9, 25.75],
|
||||
}
|
||||
pair = f"{src_lang}-{tgt_lang}"
|
||||
|
||||
readme = f"""
|
||||
---
|
||||
|
||||
language: {src_lang}, {tgt_lang}
|
||||
thumbnail:
|
||||
tags:
|
||||
- translation
|
||||
- wmt16
|
||||
- allenai
|
||||
license: Apache 2.0
|
||||
datasets:
|
||||
- http://www.statmt.org/wmt16/ ([test-set](http://matrix.statmt.org/test_sets/newstest2016.tgz?1504722372))
|
||||
|
||||
metrics:
|
||||
- http://www.statmt.org/wmt16/metrics-task.html
|
||||
---
|
||||
|
||||
# FSMT
|
||||
|
||||
## Model description
|
||||
|
||||
This is a ported version of fairseq-based [wmt16 transformer](https://github.com/jungokasai/deep-shallow/) for {src_lang}-{tgt_lang}.
|
||||
|
||||
For more details, please, see [Deep Encoder, Shallow Decoder: Reevaluating the Speed-Quality Tradeoff in Machine Translation](https://arxiv.org/abs/2006.10369).
|
||||
|
||||
All 3 models are available:
|
||||
|
||||
* [wmt16-en-de-dist-12-1](https://huggingface.co/allenai/wmt16-en-de-dist-12-1)
|
||||
* [wmt16-en-de-dist-6-1](https://huggingface.co/allenai/wmt16-en-de-dist-6-1)
|
||||
* [wmt16-en-de-12-1](https://huggingface.co/allenai/wmt16-en-de-12-1)
|
||||
|
||||
```
|
||||
@misc{{kasai2020deep,
|
||||
title={{Deep Encoder, Shallow Decoder: Reevaluating the Speed-Quality Tradeoff in Machine Translation}},
|
||||
author={{Jungo Kasai and Nikolaos Pappas and Hao Peng and James Cross and Noah A. Smith}},
|
||||
year={{2020}},
|
||||
eprint={{2006.10369}},
|
||||
archivePrefix={{arXiv}},
|
||||
primaryClass={{cs.CL}}
|
||||
}}
|
||||
```
|
||||
|
||||
## Intended uses & limitations
|
||||
|
||||
#### How to use
|
||||
|
||||
```python
|
||||
from transformers.tokenization_fsmt import FSMTTokenizer
|
||||
from transformers.modeling_fsmt import FSMTForConditionalGeneration
|
||||
mname = "allenai/{model_name}"
|
||||
tokenizer = FSMTTokenizer.from_pretrained(mname)
|
||||
model = FSMTForConditionalGeneration.from_pretrained(mname)
|
||||
|
||||
input = "{texts[src_lang]}"
|
||||
input_ids = tokenizer.encode(input, return_tensors="pt")
|
||||
outputs = model.generate(input_ids)
|
||||
decoded = tokenizer.decode(outputs[0], skip_special_tokens=True)
|
||||
print(decoded) # {texts[tgt_lang]}
|
||||
|
||||
```
|
||||
|
||||
#### Limitations and bias
|
||||
|
||||
|
||||
## Training data
|
||||
|
||||
Pretrained weights were left identical to the original model released by allenai. For more details, please, see the [paper](https://arxiv.org/abs/2006.10369).
|
||||
|
||||
## Eval results
|
||||
|
||||
Here are the BLEU scores:
|
||||
|
||||
model | fairseq | transformers
|
||||
-------|---------|----------
|
||||
{model_name} | {scores[model_name][0]} | {scores[model_name][1]}
|
||||
|
||||
The score is slightly below the score reported in the paper, as the researchers don't use `sacrebleu` and measure the score on tokenized outputs. `transformers` score was measured using `sacrebleu` on detokenized outputs.
|
||||
|
||||
The score was calculated using this code:
|
||||
|
||||
```bash
|
||||
git clone https://github.com/huggingface/transformers
|
||||
cd transformers
|
||||
export PAIR={pair}
|
||||
export DATA_DIR=data/$PAIR
|
||||
export SAVE_DIR=data/$PAIR
|
||||
export BS=8
|
||||
export NUM_BEAMS=5
|
||||
mkdir -p $DATA_DIR
|
||||
sacrebleu -t wmt16 -l $PAIR --echo src > $DATA_DIR/val.source
|
||||
sacrebleu -t wmt16 -l $PAIR --echo ref > $DATA_DIR/val.target
|
||||
echo $PAIR
|
||||
PYTHONPATH="src:examples/seq2seq" python examples/seq2seq/run_eval.py allenai/{model_name} $DATA_DIR/val.source $SAVE_DIR/test_translations.txt --reference_path $DATA_DIR/val.target --score_path $SAVE_DIR/test_bleu.json --bs $BS --task translation --num_beams $NUM_BEAMS
|
||||
```
|
||||
|
||||
"""
|
||||
model_card_dir.mkdir(parents=True, exist_ok=True)
|
||||
path = os.path.join(model_card_dir, "README.md")
|
||||
print(f"Generating {path}")
|
||||
with open(path, "w", encoding="utf-8") as f:
|
||||
f.write(readme)
|
||||
|
||||
# make sure we are under the root of the project
|
||||
repo_dir = Path(__file__).resolve().parent.parent.parent
|
||||
model_cards_dir = repo_dir / "model_cards"
|
||||
|
||||
for model_name in ["wmt16-en-de-dist-12-1", "wmt16-en-de-dist-6-1", "wmt16-en-de-12-1"]:
|
||||
model_card_dir = model_cards_dir / "allenai" / model_name
|
||||
write_model_card(model_card_dir, src_lang="en", tgt_lang="de", model_name=model_name)
|
||||
@@ -1,116 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Usage:
|
||||
# ./gen-card-allenai-wmt19.py
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
def write_model_card(model_card_dir, src_lang, tgt_lang, model_name):
|
||||
|
||||
texts = {
|
||||
"en": "Machine learning is great, isn't it?",
|
||||
"ru": "Машинное обучение - это здорово, не так ли?",
|
||||
"de": "Maschinelles Lernen ist großartig, nicht wahr?",
|
||||
}
|
||||
|
||||
# BLUE scores as follows:
|
||||
# "pair": [fairseq, transformers]
|
||||
scores = {
|
||||
"wmt19-de-en-6-6-base": [0, 38.37],
|
||||
"wmt19-de-en-6-6-big": [0, 39.90],
|
||||
}
|
||||
pair = f"{src_lang}-{tgt_lang}"
|
||||
|
||||
readme = f"""
|
||||
---
|
||||
|
||||
language: {src_lang}, {tgt_lang}
|
||||
thumbnail:
|
||||
tags:
|
||||
- translation
|
||||
- wmt19
|
||||
- allenai
|
||||
license: Apache 2.0
|
||||
datasets:
|
||||
- http://www.statmt.org/wmt19/ ([test-set](http://matrix.statmt.org/test_sets/newstest2019.tgz?1556572561))
|
||||
metrics:
|
||||
- http://www.statmt.org/wmt19/metrics-task.html
|
||||
---
|
||||
|
||||
# FSMT
|
||||
|
||||
## Model description
|
||||
|
||||
This is a ported version of fairseq-based wmt19 transformer created by [jungokasai]](https://github.com/jungokasai/) @ allenai for {src_lang}-{tgt_lang}.
|
||||
|
||||
2 models are available:
|
||||
|
||||
* [wmt19-de-en-6-6-big](https://huggingface.co/allenai/wmt19-de-en-6-6-big)
|
||||
* [wmt19-de-en-6-6-base](https://huggingface.co/allenai/wmt19-de-en-6-6-base)
|
||||
|
||||
## Intended uses & limitations
|
||||
|
||||
#### How to use
|
||||
|
||||
```python
|
||||
from transformers.tokenization_fsmt import FSMTTokenizer
|
||||
from transformers.modeling_fsmt import FSMTForConditionalGeneration
|
||||
mname = "allenai/{model_name}"
|
||||
tokenizer = FSMTTokenizer.from_pretrained(mname)
|
||||
model = FSMTForConditionalGeneration.from_pretrained(mname)
|
||||
|
||||
input = "{texts[src_lang]}"
|
||||
input_ids = tokenizer.encode(input, return_tensors="pt")
|
||||
outputs = model.generate(input_ids)
|
||||
decoded = tokenizer.decode(outputs[0], skip_special_tokens=True)
|
||||
print(decoded) # {texts[tgt_lang]}
|
||||
|
||||
```
|
||||
|
||||
#### Limitations and bias
|
||||
|
||||
|
||||
## Training data
|
||||
|
||||
Pretrained weights were left identical to the original model released by the researcher.
|
||||
|
||||
## Eval results
|
||||
|
||||
Here are the BLEU scores:
|
||||
|
||||
model | transformers
|
||||
-------|---------|----------
|
||||
{model_name} | {scores[model_name][1]}
|
||||
|
||||
The score was calculated using this code:
|
||||
|
||||
```bash
|
||||
git clone https://github.com/huggingface/transformers
|
||||
cd transformers
|
||||
export PAIR={pair}
|
||||
export DATA_DIR=data/$PAIR
|
||||
export SAVE_DIR=data/$PAIR
|
||||
export BS=8
|
||||
export NUM_BEAMS=5
|
||||
mkdir -p $DATA_DIR
|
||||
sacrebleu -t wmt19 -l $PAIR --echo src > $DATA_DIR/val.source
|
||||
sacrebleu -t wmt19 -l $PAIR --echo ref > $DATA_DIR/val.target
|
||||
echo $PAIR
|
||||
PYTHONPATH="src:examples/seq2seq" python examples/seq2seq/run_eval.py allenai/{model_name} $DATA_DIR/val.source $SAVE_DIR/test_translations.txt --reference_path $DATA_DIR/val.target --score_path $SAVE_DIR/test_bleu.json --bs $BS --task translation --num_beams $NUM_BEAMS
|
||||
```
|
||||
|
||||
"""
|
||||
model_card_dir.mkdir(parents=True, exist_ok=True)
|
||||
path = os.path.join(model_card_dir, "README.md")
|
||||
print(f"Generating {path}")
|
||||
with open(path, "w", encoding="utf-8") as f:
|
||||
f.write(readme)
|
||||
|
||||
# make sure we are under the root of the project
|
||||
repo_dir = Path(__file__).resolve().parent.parent.parent
|
||||
model_cards_dir = repo_dir / "model_cards"
|
||||
|
||||
for model_name in ["wmt19-de-en-6-6-base", "wmt19-de-en-6-6-big"]:
|
||||
model_card_dir = model_cards_dir / "allenai" / model_name
|
||||
write_model_card(model_card_dir, src_lang="de", tgt_lang="en", model_name=model_name)
|
||||
@@ -1,135 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Usage:
|
||||
# ./gen-card-facebook-wmt19.py
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
def write_model_card(model_card_dir, src_lang, tgt_lang):
|
||||
|
||||
texts = {
|
||||
"en": "Machine learning is great, isn't it?",
|
||||
"ru": "Машинное обучение - это здорово, не так ли?",
|
||||
"de": "Maschinelles Lernen ist großartig, oder?",
|
||||
}
|
||||
|
||||
# BLUE scores as follows:
|
||||
# "pair": [fairseq, transformers]
|
||||
scores = {
|
||||
"ru-en": ["[41.3](http://matrix.statmt.org/matrix/output/1907?run_id=6937)", "39.20"],
|
||||
"en-ru": ["[36.4](http://matrix.statmt.org/matrix/output/1914?run_id=6724)", "33.47"],
|
||||
"en-de": ["[43.1](http://matrix.statmt.org/matrix/output/1909?run_id=6862)", "42.83"],
|
||||
"de-en": ["[42.3](http://matrix.statmt.org/matrix/output/1902?run_id=6750)", "41.35"],
|
||||
}
|
||||
pair = f"{src_lang}-{tgt_lang}"
|
||||
|
||||
readme = f"""
|
||||
---
|
||||
|
||||
<!-- This file has been auto-generated by src/transformers/convert_fsmt_original_pytorch_checkpoint_to_pytorch.py - DO NOT EDIT or your changes will be lost -->
|
||||
|
||||
language: {src_lang}, {tgt_lang}
|
||||
thumbnail:
|
||||
tags:
|
||||
- translation
|
||||
- wmt19
|
||||
license: Apache 2.0
|
||||
datasets:
|
||||
- http://www.statmt.org/wmt19/ ([test-set](http://matrix.statmt.org/test_sets/newstest2019.tgz?1556572561))
|
||||
metrics:
|
||||
- http://www.statmt.org/wmt19/metrics-task.html
|
||||
---
|
||||
|
||||
# FSMT
|
||||
|
||||
## Model description
|
||||
|
||||
This is a ported version of [fairseq wmt19 transformer](https://github.com/pytorch/fairseq/blob/master/examples/wmt19/README.md) for {src_lang}-{tgt_lang}.
|
||||
|
||||
For more details, please see, [Facebook FAIR's WMT19 News Translation Task Submission](https://arxiv.org/abs/1907.06616).
|
||||
|
||||
The abbreviation FSMT stands for FairSeqMachineTranslation
|
||||
|
||||
All four models are available:
|
||||
|
||||
* [wmt19-en-ru](https://huggingface.co/facebook/wmt19-en-ru)
|
||||
* [wmt19-ru-en](https://huggingface.co/facebook/wmt19-ru-en)
|
||||
* [wmt19-en-de](https://huggingface.co/facebook/wmt19-en-de)
|
||||
* [wmt19-de-en](https://huggingface.co/facebook/wmt19-de-en)
|
||||
|
||||
## Intended uses & limitations
|
||||
|
||||
#### How to use
|
||||
|
||||
```python
|
||||
from transformers.tokenization_fsmt import FSMTTokenizer
|
||||
from transformers.modeling_fsmt import FSMTForConditionalGeneration
|
||||
mname = "facebook/wmt19-{src_lang}-{tgt_lang}"
|
||||
tokenizer = FSMTTokenizer.from_pretrained(mname)
|
||||
model = FSMTForConditionalGeneration.from_pretrained(mname)
|
||||
|
||||
input = "{texts[src_lang]}
|
||||
input_ids = tokenizer.encode(input, return_tensors="pt")
|
||||
outputs = model.generate(input_ids)
|
||||
decoded = tokenizer.decode(outputs[0], skip_special_tokens=True)
|
||||
print(decoded) # {texts[tgt_lang]}
|
||||
|
||||
```
|
||||
|
||||
#### Limitations and bias
|
||||
|
||||
- The original (and this ported model) doesn't seem to handle well inputs with repeated sub-phrases, [content gets truncated](https://discuss.huggingface.co/t/issues-with-translating-inputs-containing-repeated-phrases/981)
|
||||
|
||||
## Training data
|
||||
|
||||
Pretrained weights were left identical to the original model released by fairseq. For more details, please, see the [paper](https://arxiv.org/abs/1907.06616).
|
||||
|
||||
## Eval results
|
||||
|
||||
pair | fairseq | transformers
|
||||
-------|---------|----------
|
||||
{pair} | {scores[pair][0]} | {scores[pair][1]}
|
||||
|
||||
The score is slightly below the score reported by `fairseq`, since `transformers`` currently doesn't support:
|
||||
- model ensemble, therefore the best performing checkpoint was ported (``model4.pt``).
|
||||
- re-ranking
|
||||
|
||||
The score was calculated using this code:
|
||||
|
||||
```bash
|
||||
git clone https://github.com/huggingface/transformers
|
||||
cd transformers
|
||||
export PAIR={pair}
|
||||
export DATA_DIR=data/$PAIR
|
||||
export SAVE_DIR=data/$PAIR
|
||||
export BS=8
|
||||
export NUM_BEAMS=15
|
||||
mkdir -p $DATA_DIR
|
||||
sacrebleu -t wmt19 -l $PAIR --echo src > $DATA_DIR/val.source
|
||||
sacrebleu -t wmt19 -l $PAIR --echo ref > $DATA_DIR/val.target
|
||||
echo $PAIR
|
||||
PYTHONPATH="src:examples/seq2seq" python examples/seq2seq/run_eval.py facebook/wmt19-$PAIR $DATA_DIR/val.source $SAVE_DIR/test_translations.txt --reference_path $DATA_DIR/val.target --score_path $SAVE_DIR/test_bleu.json --bs $BS --task translation --num_beams $NUM_BEAMS
|
||||
```
|
||||
note: fairseq reports using a beam of 50, so you should get a slightly higher score if re-run with `--num_beams 50`.
|
||||
|
||||
|
||||
## TODO
|
||||
|
||||
- port model ensemble (fairseq uses 4 model checkpoints)
|
||||
|
||||
"""
|
||||
os.makedirs(model_card_dir, exist_ok=True)
|
||||
path = os.path.join(model_card_dir, "README.md")
|
||||
print(f"Generating {path}")
|
||||
with open(path, "w", encoding="utf-8") as f:
|
||||
f.write(readme)
|
||||
|
||||
# make sure we are under the root of the project
|
||||
repo_dir = Path(__file__).resolve().parent.parent.parent
|
||||
model_cards_dir = repo_dir / "model_cards"
|
||||
|
||||
for model_name in ["wmt19-ru-en", "wmt19-en-ru", "wmt19-en-de", "wmt19-de-en"]:
|
||||
base, src_lang, tgt_lang = model_name.split("-")
|
||||
model_card_dir = model_cards_dir / "facebook" / model_name
|
||||
write_model_card(model_card_dir, src_lang=src_lang, tgt_lang=tgt_lang)
|
||||
@@ -70,26 +70,26 @@ extras["sklearn"] = ["scikit-learn"]
|
||||
|
||||
# keras2onnx and onnxconverter-common version is specific through a commit until 1.7.0 lands on pypi
|
||||
extras["tf"] = [
|
||||
"tensorflow>=2.0",
|
||||
"tensorflow",
|
||||
"onnxconverter-common",
|
||||
"keras2onnx"
|
||||
# "onnxconverter-common @ git+git://github.com/microsoft/onnxconverter-common.git@f64ca15989b6dc95a1f3507ff6e4c395ba12dff5#egg=onnxconverter-common",
|
||||
# "keras2onnx @ git+git://github.com/onnx/keras-onnx.git@cbdc75cb950b16db7f0a67be96a278f8d2953b48#egg=keras2onnx",
|
||||
]
|
||||
extras["tf-cpu"] = [
|
||||
"tensorflow-cpu>=2.0",
|
||||
"tensorflow-cpu",
|
||||
"onnxconverter-common",
|
||||
"keras2onnx"
|
||||
# "onnxconverter-common @ git+git://github.com/microsoft/onnxconverter-common.git@f64ca15989b6dc95a1f3507ff6e4c395ba12dff5#egg=onnxconverter-common",
|
||||
# "keras2onnx @ git+git://github.com/onnx/keras-onnx.git@cbdc75cb950b16db7f0a67be96a278f8d2953b48#egg=keras2onnx",
|
||||
]
|
||||
extras["torch"] = ["torch>=1.0"]
|
||||
extras["torch"] = ["torch"]
|
||||
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"]
|
||||
@@ -98,7 +98,7 @@ extras["dev"] = extras["testing"] + extras["quality"] + extras["ja"] + ["scikit-
|
||||
setup(
|
||||
name="transformers",
|
||||
version="3.1.0",
|
||||
author="Thomas Wolf, Lysandre Debut, Victor Sanh, Julien Chaumond, Sam Shleifer, Patrick von Platen, Sylvain Gugger, Google AI Language Team Authors, Open AI team Authors, Facebook AI Authors, Carnegie Mellon University Authors",
|
||||
author="Thomas Wolf, Lysandre Debut, Victor Sanh, Julien Chaumond, Sam Shleifer, Patrick von Platen, Google AI Language Team Authors, Open AI team Authors, Facebook AI Authors, Carnegie Mellon University Authors",
|
||||
author_email="thomas@huggingface.co",
|
||||
description="State-of-the-art Natural Language Processing for TensorFlow 2.0 and PyTorch",
|
||||
long_description=open("README.md", "r", encoding="utf-8").read(),
|
||||
|
||||
@@ -84,6 +84,7 @@ from .file_utils import (
|
||||
cached_path,
|
||||
is_apex_available,
|
||||
is_datasets_available,
|
||||
is_faiss_available,
|
||||
is_psutil_available,
|
||||
is_py3nvml_available,
|
||||
is_tf_available,
|
||||
@@ -712,6 +713,13 @@ if is_tf_available():
|
||||
from .trainer_tf import TFTrainer
|
||||
|
||||
|
||||
if is_torch_available() and is_datasets_available() and is_faiss_available() and is_psutil_available():
|
||||
from .configuration_rag import RagConfig
|
||||
from .modeling_rag import RagModel, RagSequenceForGeneration, RagTokenForGeneration
|
||||
from .retrieval_rag import RagRetriever
|
||||
from .tokenization_rag import RagTokenizer
|
||||
|
||||
|
||||
if not is_tf_available() and not is_torch_available():
|
||||
logger.warning(
|
||||
"Neither PyTorch nor TensorFlow >= 2.0 have been found."
|
||||
|
||||
@@ -1,65 +0,0 @@
|
||||
import math
|
||||
|
||||
import tensorflow as tf
|
||||
|
||||
|
||||
def gelu(x):
|
||||
"""Gaussian Error Linear Unit.
|
||||
Original Implementation of the gelu activation function in Google Bert repo when initially created.
|
||||
For information: OpenAI GPT's gelu is slightly different (and gives slightly different results):
|
||||
0.5 * x * (1 + torch.tanh(math.sqrt(2 / math.pi) * (x + 0.044715 * torch.pow(x, 3))))
|
||||
Also see https://arxiv.org/abs/1606.08415
|
||||
"""
|
||||
x = tf.convert_to_tensor(x)
|
||||
cdf = 0.5 * (1.0 + tf.math.erf(x / tf.math.sqrt(2.0)))
|
||||
|
||||
return x * cdf
|
||||
|
||||
|
||||
def gelu_new(x):
|
||||
"""Gaussian Error Linear Unit.
|
||||
This is a smoother version of the GELU.
|
||||
Original paper: https://arxiv.org/abs/1606.08415
|
||||
Args:
|
||||
x: float Tensor to perform activation.
|
||||
Returns:
|
||||
`x` with the GELU activation applied.
|
||||
"""
|
||||
x = tf.convert_to_tensor(x)
|
||||
pi = tf.cast(math.pi, x.dtype)
|
||||
coeff = tf.cast(0.044715, x.dtype)
|
||||
cdf = 0.5 * (1.0 + tf.tanh(tf.sqrt(2.0 / pi) * (x + coeff * tf.pow(x, 3))))
|
||||
|
||||
return x * cdf
|
||||
|
||||
|
||||
def mish(x):
|
||||
x = tf.convert_to_tensor(x)
|
||||
|
||||
return x * tf.tanh(tf.math.softplus(x))
|
||||
|
||||
|
||||
def gelu_fast(x):
|
||||
x = tf.convert_to_tensor(x)
|
||||
coeff1 = tf.cast(7978845608, x.dtype)
|
||||
coeff2 = tf.cast(0.044715, x.dtype)
|
||||
|
||||
return 0.5 * x * (1.0 + tf.tanh(x * coeff2 * (1.0 + coeff1 * x * x)))
|
||||
|
||||
|
||||
ACT2FN = {
|
||||
"gelu": tf.keras.layers.Activation(gelu),
|
||||
"relu": tf.keras.activations.relu,
|
||||
"swish": tf.keras.activations.swish,
|
||||
"gelu_new": tf.keras.layers.Activation(gelu_new),
|
||||
"mish": tf.keras.layers.Activation(mish),
|
||||
"tanh": tf.keras.activations.tanh,
|
||||
"gelu_fast": tf.keras.layers.Activation(gelu_fast),
|
||||
}
|
||||
|
||||
|
||||
def get_tf_activation(activation_string):
|
||||
if activation_string in ACT2FN:
|
||||
return ACT2FN[activation_string]
|
||||
else:
|
||||
raise KeyError("function {} not found in ACT2FN mapping {}".format(activation_string, list(ACT2FN.keys())))
|
||||
@@ -24,6 +24,7 @@ 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
|
||||
from .configuration_dpr import DPR_PRETRAINED_CONFIG_ARCHIVE_MAP, DPRConfig
|
||||
from .configuration_electra import ELECTRA_PRETRAINED_CONFIG_ARCHIVE_MAP, ElectraConfig
|
||||
from .configuration_encoder_decoder import EncoderDecoderConfig
|
||||
from .configuration_flaubert import FLAUBERT_PRETRAINED_CONFIG_ARCHIVE_MAP, FlaubertConfig
|
||||
@@ -71,6 +72,7 @@ ALL_PRETRAINED_CONFIG_ARCHIVE_MAP = dict(
|
||||
RETRIBERT_PRETRAINED_CONFIG_ARCHIVE_MAP,
|
||||
FUNNEL_PRETRAINED_CONFIG_ARCHIVE_MAP,
|
||||
LXMERT_PRETRAINED_CONFIG_ARCHIVE_MAP,
|
||||
DPR_PRETRAINED_CONFIG_ARCHIVE_MAP,
|
||||
]
|
||||
for key, value, in pretrained_map.items()
|
||||
)
|
||||
@@ -105,6 +107,7 @@ CONFIG_MAPPING = OrderedDict(
|
||||
("encoder-decoder", EncoderDecoderConfig),
|
||||
("funnel", FunnelConfig),
|
||||
("lxmert", LxmertConfig),
|
||||
("dpr", DPRConfig),
|
||||
]
|
||||
)
|
||||
|
||||
@@ -137,6 +140,7 @@ MODEL_NAMES_MAPPING = OrderedDict(
|
||||
("encoder-decoder", "Encoder decoder"),
|
||||
("funnel", "Funnel Transformer"),
|
||||
("lxmert", "LXMERT"),
|
||||
("dpr", "DPR"),
|
||||
]
|
||||
)
|
||||
|
||||
@@ -197,7 +201,9 @@ class AutoConfig:
|
||||
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.
|
||||
|
||||
This class cannot be instantiated directly using ``__init__()`` (throws an error).
|
||||
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.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
@@ -220,77 +226,58 @@ class AutoConfig:
|
||||
@classmethod
|
||||
@replace_list_option_in_docstrings()
|
||||
def from_pretrained(cls, pretrained_model_name_or_path, **kwargs):
|
||||
r"""
|
||||
Instantiate one of the configuration classes of the library from a pretrained model configuration.
|
||||
r""" Instantiates one of the configuration classes of the library
|
||||
from a pre-trained model configuration.
|
||||
|
||||
The configuration class to instantiate is selected based on the :obj:`model_type` property of the config
|
||||
object that is loaded, or when it's missing, by falling back to using pattern matching on
|
||||
:obj:`pretrained_model_name_or_path`:
|
||||
The configuration class to instantiate is selected
|
||||
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
|
||||
|
||||
Args:
|
||||
pretrained_model_name_or_path (:obj:`str`):
|
||||
Can be either:
|
||||
pretrained_model_name_or_path (:obj:`string`):
|
||||
Is either: \
|
||||
- a string with the `shortcut name` of a pre-trained model configuration to load from cache or download, e.g.: ``bert-base-uncased``.
|
||||
- a string with the `identifier name` of a pre-trained model configuration that was user-uploaded to our S3, e.g.: ``dbmdz/bert-base-german-cased``.
|
||||
- a path to a `directory` containing a configuration file saved using the :func:`~transformers.PretrainedConfig.save_pretrained` method, e.g.: ``./my_model_directory/``.
|
||||
- a path or url to a saved configuration JSON `file`, e.g.: ``./my_model_directory/configuration.json``.
|
||||
|
||||
- A string with the `shortcut name` of a pretrained model configuration to load from cache or
|
||||
download, e.g., ``bert-base-uncased``.
|
||||
- A string with the `identifier name` of a pretrained model configuration that was user-uploaded to
|
||||
our S3, e.g., ``dbmdz/bert-base-german-cased``.
|
||||
- A path to a `directory` containing a configuration file saved using the
|
||||
:meth:`~transformers.PretrainedConfig.save_pretrained` method, or the
|
||||
:meth:`~transformers.PretrainedModel.save_pretrained` method, e.g., ``./my_model_directory/``.
|
||||
- A path or url to a saved configuration JSON `file`, e.g.,
|
||||
``./my_model_directory/configuration.json``.
|
||||
cache_dir (:obj:`str`, `optional`):
|
||||
Path to a directory in which a downloaded pretrained model configuration should be cached if the
|
||||
standard cache should not be used.
|
||||
force_download (:obj:`bool`, `optional`, defaults to :obj:`False`):
|
||||
Whether or not to force the (re-)download the model weights and configuration files and override the
|
||||
cached versions if they exist.
|
||||
resume_download (:obj:`bool`, `optional`, defaults to :obj:`False`):
|
||||
Whether or not to delete incompletely received files. Will attempt to resume the download if such a
|
||||
file exists.
|
||||
proxies (:obj:`Dict[str, str]`, `optional`):
|
||||
A dictionary of proxy servers to use by protocol or endpoint, e.g.,
|
||||
:obj:`{'http': 'foo.bar:3128', 'http://hostname': 'foo.bar:4012'}`. The proxies are used on each
|
||||
request.
|
||||
return_unused_kwargs (:obj:`bool`, `optional`, defaults to :obj:`False`):
|
||||
If :obj:`False`, then this function returns just the final configuration object.
|
||||
cache_dir (:obj:`string`, optional, defaults to `None`):
|
||||
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 (:obj:`boolean`, optional, defaults to `False`):
|
||||
Force to (re-)download the model weights and configuration files and override the cached versions if they exist.
|
||||
|
||||
resume_download (:obj:`boolean`, optional, defaults to `False`):
|
||||
Do not delete incompletely received file. Attempt to resume the download if such a file exists.
|
||||
|
||||
proxies (:obj:`Dict[str, str]`, optional, defaults to `None`):
|
||||
A dictionary of proxy servers to use by protocol or endpoint, e.g.: :obj:`{'http': 'foo.bar:3128', 'http://hostname': 'foo.bar:4012'}`.
|
||||
The proxies are used on each request. See `the requests documentation <https://requests.readthedocs.io/en/master/user/advanced/#proxies>`__ for usage.
|
||||
|
||||
return_unused_kwargs (:obj:`boolean`, optional, defaults to `False`):
|
||||
- If False, then this function returns just the final configuration object.
|
||||
- If True, then this functions returns a tuple `(config, unused_kwargs)` where `unused_kwargs` is a dictionary consisting of the key/value pairs whose keys are not configuration attributes: ie the part of kwargs which has not been used to update `config` and is otherwise ignored.
|
||||
|
||||
kwargs (:obj:`Dict[str, any]`, optional, defaults to `{}`): key/value pairs with which to update the configuration object after loading.
|
||||
- The values in kwargs of any keys which are configuration attributes will be used to override the loaded values.
|
||||
- Behavior concerning key/value pairs whose keys are *not* configuration attributes is controlled by the `return_unused_kwargs` keyword parameter.
|
||||
|
||||
If :obj:`True`, then this functions returns a :obj:`Tuple(config, unused_kwargs)` where `unused_kwargs`
|
||||
is a dictionary consisting of the key/value pairs whose keys are not configuration attributes: i.e.,
|
||||
the part of ``kwargs`` which has not been used to update ``config`` and is otherwise ignored.
|
||||
kwargs(additional keyword arguments, `optional`):
|
||||
The values in kwargs of any keys which are configuration attributes will be used to override the loaded
|
||||
values. Behavior concerning key/value pairs whose keys are *not* configuration attributes is
|
||||
controlled by the ``return_unused_kwargs`` keyword parameter.
|
||||
|
||||
Examples::
|
||||
|
||||
>>> from transformers import AutoConfig
|
||||
config = AutoConfig.from_pretrained('bert-base-uncased') # Download configuration from S3 and cache.
|
||||
config = AutoConfig.from_pretrained('./test/bert_saved_model/') # E.g. config (or model) was saved using `save_pretrained('./test/saved_model/')`
|
||||
config = AutoConfig.from_pretrained('./test/bert_saved_model/my_configuration.json')
|
||||
config = AutoConfig.from_pretrained('bert-base-uncased', output_attentions=True, foo=False)
|
||||
assert config.output_attentions == True
|
||||
config, unused_kwargs = AutoConfig.from_pretrained('bert-base-uncased', output_attentions=True,
|
||||
foo=False, return_unused_kwargs=True)
|
||||
assert config.output_attentions == True
|
||||
assert unused_kwargs == {'foo': False}
|
||||
|
||||
>>> # Download configuration from S3 and cache.
|
||||
>>> config = AutoConfig.from_pretrained('bert-base-uncased')
|
||||
|
||||
>>> # Download configuration from S3 (user-uploaded) and cache.
|
||||
>>> config = AutoConfig.from_pretrained('dbmdz/bert-base-german-cased')
|
||||
|
||||
>>> # If configuration file is in a directory (e.g., was saved using `save_pretrained('./test/saved_model/')`).
|
||||
>>> config = AutoConfig.from_pretrained('./test/bert_saved_model/')
|
||||
|
||||
>>> # Load a specific configuration file.
|
||||
>>> config = AutoConfig.from_pretrained('./test/bert_saved_model/my_configuration.json')
|
||||
|
||||
>>> # Change some config attributes when loading a pretrained config.
|
||||
>>> config = AutoConfig.from_pretrained('bert-base-uncased', output_attentions=True, foo=False)
|
||||
>>> config.output_attentions
|
||||
True
|
||||
>>> config, unused_kwargs = AutoConfig.from_pretrained('bert-base-uncased', output_attentions=True, foo=False, return_unused_kwargs=True)
|
||||
>>> config.output_attentions
|
||||
True
|
||||
>>> config.unused_kwargs
|
||||
{'foo': False}
|
||||
"""
|
||||
config_dict, _ = PretrainedConfig.get_config_dict(pretrained_model_name_or_path, **kwargs)
|
||||
|
||||
|
||||
@@ -14,7 +14,7 @@
|
||||
# limitations under the License.
|
||||
""" DPR model configuration """
|
||||
|
||||
from .configuration_bert import BertConfig
|
||||
from .configuration_utils import PretrainedConfig
|
||||
from .utils import logging
|
||||
|
||||
|
||||
@@ -27,7 +27,7 @@ DPR_PRETRAINED_CONFIG_ARCHIVE_MAP = {
|
||||
}
|
||||
|
||||
|
||||
class DPRConfig(BertConfig):
|
||||
class DPRConfig(PretrainedConfig):
|
||||
r"""
|
||||
:class:`~transformers.DPRConfig` is the configuration class to store the configuration of a
|
||||
`DPRModel`.
|
||||
@@ -36,12 +36,73 @@ class DPRConfig(BertConfig):
|
||||
It is used to instantiate the components of the DPR model.
|
||||
|
||||
Args:
|
||||
vocab_size (:obj:`int`, optional, defaults to 30522):
|
||||
Vocabulary size of the DPR model. Defines the different tokens that
|
||||
can be represented by the `inputs_ids` passed to the forward method of :class:`~transformers.BertModel`.
|
||||
hidden_size (:obj:`int`, optional, defaults to 768):
|
||||
Dimensionality of the encoder layers and the pooler layer.
|
||||
num_hidden_layers (:obj:`int`, optional, defaults to 12):
|
||||
Number of hidden layers in the Transformer encoder.
|
||||
num_attention_heads (:obj:`int`, optional, defaults to 12):
|
||||
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 "gelu"):
|
||||
The non-linear activation function (function or string) in the encoder and pooler.
|
||||
If string, "gelu", "relu", "swish" and "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).
|
||||
type_vocab_size (:obj:`int`, optional, defaults to 2):
|
||||
The vocabulary size of the `token_type_ids` passed into :class:`~transformers.BertModel`.
|
||||
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 True, use gradient checkpointing to save memory at the expense of slower backward pass.
|
||||
projection_dim (:obj:`int`, optional, defaults to 0):
|
||||
Dimension of the projection for the context and question encoders.
|
||||
If it is set to zero (default), then no projection is done.
|
||||
"""
|
||||
model_type = "dpr"
|
||||
|
||||
def __init__(self, projection_dim: int = 0, **kwargs): # projection of the encoders, 0 for no projection
|
||||
super().__init__(**kwargs)
|
||||
def __init__(
|
||||
self,
|
||||
vocab_size=30522,
|
||||
hidden_size=768,
|
||||
num_hidden_layers=12,
|
||||
num_attention_heads=12,
|
||||
intermediate_size=3072,
|
||||
hidden_act="gelu",
|
||||
hidden_dropout_prob=0.1,
|
||||
attention_probs_dropout_prob=0.1,
|
||||
max_position_embeddings=512,
|
||||
type_vocab_size=2,
|
||||
initializer_range=0.02,
|
||||
layer_norm_eps=1e-12,
|
||||
pad_token_id=0,
|
||||
gradient_checkpointing=False,
|
||||
projection_dim: int = 0,
|
||||
**kwargs
|
||||
):
|
||||
super().__init__(pad_token_id=pad_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.type_vocab_size = type_vocab_size
|
||||
self.initializer_range = initializer_range
|
||||
self.layer_norm_eps = layer_norm_eps
|
||||
self.gradient_checkpointing = gradient_checkpointing
|
||||
self.projection_dim = projection_dim
|
||||
|
||||
@@ -49,7 +49,7 @@ DEFAULTS = dict(
|
||||
# Config values that vary between checkpoints: for testing and conversion
|
||||
task_specific_params = {
|
||||
# These are task specific params for pegasus-large and normal params for finetuned checkpoints
|
||||
"summarization_xsum": {"length_penalty": 0.6, "max_length": 64, "max_position_embeddings": 512},
|
||||
"summarization_xsum": {"length_penalty": 0.8, "max_length": 64, "max_position_embeddings": 512},
|
||||
"summarization_cnn_dailymail": {"length_penalty": 0.8, "max_length": 128, "max_position_embeddings": 1024},
|
||||
"summarization_newsroom": {"length_penalty": 0.8, "max_length": 128, "max_position_embeddings": 512},
|
||||
"summarization_wikihow": {"length_penalty": 0.6, "max_length": 256, "max_position_embeddings": 512},
|
||||
|
||||
@@ -0,0 +1,202 @@
|
||||
# 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 """
|
||||
|
||||
import copy
|
||||
|
||||
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:`~nlp.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``)
|
||||
The path to the serialized faiss index on disk.
|
||||
passages_path: (:obj:`str`, optional, defaults to ``None``):
|
||||
A path to text passages compatible with the faiss index. Required if using :class:`~transformers.retrieval_rag.LegacyIndex`
|
||||
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.
|
||||
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.
|
||||
|
||||
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.
|
||||
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`.
|
||||
|
||||
"""
|
||||
|
||||
|
||||
@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,
|
||||
dataset="wiki_dpr",
|
||||
dataset_split="train",
|
||||
index_name="compressed",
|
||||
index_path=None,
|
||||
passages_path=None,
|
||||
use_dummy_dataset=False,
|
||||
reduce_loss=False,
|
||||
label_smoothing=0.0,
|
||||
deduplicate=True, # defaults to True
|
||||
num_doc_return_sequences=1, # defaults to 1
|
||||
num_doc_beams=1, # defaults to 1
|
||||
exclude_bos_score=False,
|
||||
do_marginalize=False,
|
||||
output_retrieved=False,
|
||||
**kwargs
|
||||
):
|
||||
super().__init__(**kwargs)
|
||||
assert (
|
||||
"question_encoder" in kwargs and "generator" in kwargs
|
||||
), "Config has to be initialized with question_encoder and generator config"
|
||||
question_encoder_config = kwargs.pop("question_encoder")
|
||||
question_encoder_model_type = question_encoder_config.pop("model_type")
|
||||
decoder_config = kwargs.pop("generator")
|
||||
decoder_model_type = decoder_config.pop("model_type")
|
||||
|
||||
from .configuration_auto import AutoConfig
|
||||
|
||||
self.question_encoder = AutoConfig.for_model(question_encoder_model_type, **question_encoder_config)
|
||||
self.generator = AutoConfig.for_model(decoder_model_type, **decoder_config)
|
||||
|
||||
self.reduce_loss = reduce_loss
|
||||
self.label_smoothing = label_smoothing
|
||||
self.exclude_bos_score = exclude_bos_score
|
||||
self.do_marginalize = do_marginalize
|
||||
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.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.passages_path = passages_path
|
||||
self.index_path = index_path
|
||||
self.use_dummy_dataset = use_dummy_dataset
|
||||
|
||||
self.output_retrieved = output_retrieved
|
||||
|
||||
self.deduplicate = deduplicate
|
||||
self.num_doc_return_sequences = num_doc_return_sequences
|
||||
self.num_doc_beams = num_doc_beams
|
||||
|
||||
@classmethod
|
||||
def from_question_encoder_generator_configs(
|
||||
cls, question_encoder_config: PretrainedConfig, generator_config: PretrainedConfig, **kwargs
|
||||
) -> PretrainedConfig:
|
||||
r"""
|
||||
Instantiate a :class:`~transformers.EncoderDecoderConfig` (or a derived class) from a pre-trained encoder model configuration and decoder model configuration.
|
||||
|
||||
Returns:
|
||||
:class:`EncoderDecoderConfig`: An instance of a configuration object
|
||||
"""
|
||||
return cls(question_encoder=question_encoder_config.to_dict(), generator=generator_config.to_dict(), **kwargs)
|
||||
|
||||
def to_dict(self):
|
||||
"""
|
||||
Serializes this instance to a Python dictionary. Override the default `to_dict()` from `PretrainedConfig`.
|
||||
|
||||
Returns:
|
||||
:obj:`Dict[str, any]`: Dictionary of all the attributes that make up this configuration instance,
|
||||
"""
|
||||
output = copy.deepcopy(self.__dict__)
|
||||
output["question_encoder"] = self.question_encoder.to_dict()
|
||||
output["generator"] = self.generator.to_dict()
|
||||
output["model_type"] = self.__class__.model_type
|
||||
return output
|
||||
@@ -62,6 +62,10 @@ class TransfoXLConfig(PretrainedConfig):
|
||||
Apply LayerNorm to the input instead of the output
|
||||
n_layer (:obj:`int`, optional, defaults to 18):
|
||||
Number of hidden layers in the Transformer encoder.
|
||||
tgt_len (:obj:`int`, optional, defaults to 128):
|
||||
Number of tokens to predict
|
||||
ext_len (:obj:`int`, optional, defaults to 0):
|
||||
Length of the extended context
|
||||
mem_len (:obj:`int`, optional, defaults to 1600):
|
||||
Length of the retained previous heads
|
||||
clamp_len (:obj:`int`, optional, defaults to 1000):
|
||||
@@ -121,6 +125,8 @@ class TransfoXLConfig(PretrainedConfig):
|
||||
div_val=4,
|
||||
pre_lnorm=False,
|
||||
n_layer=18,
|
||||
tgt_len=128,
|
||||
ext_len=0,
|
||||
mem_len=1600,
|
||||
clamp_len=1000,
|
||||
same_length=True,
|
||||
@@ -162,6 +168,8 @@ class TransfoXLConfig(PretrainedConfig):
|
||||
self.pre_lnorm = pre_lnorm
|
||||
self.n_layer = n_layer
|
||||
self.n_head = n_head
|
||||
self.tgt_len = tgt_len
|
||||
self.ext_len = ext_len
|
||||
self.mem_len = mem_len
|
||||
self.same_length = same_length
|
||||
self.attn_type = attn_type
|
||||
@@ -179,9 +187,7 @@ class TransfoXLConfig(PretrainedConfig):
|
||||
|
||||
@property
|
||||
def max_position_embeddings(self):
|
||||
# Message copied from Transformer-XL documentation
|
||||
logger.info(f"The model {self.model_type} is one of the few models that has no sequence length limit.")
|
||||
return -1
|
||||
return self.tgt_len + self.ext_len + self.mem_len
|
||||
|
||||
@property
|
||||
def n_token(self): # Backward compatibility
|
||||
|
||||
@@ -337,9 +337,7 @@ class PretrainedConfig(object):
|
||||
elif os.path.isfile(pretrained_model_name_or_path) or is_remote_url(pretrained_model_name_or_path):
|
||||
config_file = pretrained_model_name_or_path
|
||||
else:
|
||||
config_file = hf_bucket_url(
|
||||
pretrained_model_name_or_path, filename=CONFIG_NAME, use_cdn=False, mirror=None
|
||||
)
|
||||
config_file = hf_bucket_url(pretrained_model_name_or_path, filename=CONFIG_NAME, use_cdn=False)
|
||||
|
||||
try:
|
||||
# Load from URL or cache if already cached
|
||||
|
||||
@@ -47,8 +47,8 @@ PATTERNS = [
|
||||
|
||||
def rename_state_dict_key(k):
|
||||
|
||||
for pegasus_name, hf_name in PATTERNS:
|
||||
k = k.replace(pegasus_name, hf_name)
|
||||
for pegasus_name, bart_name in PATTERNS:
|
||||
k = k.replace(pegasus_name, bart_name)
|
||||
return k
|
||||
|
||||
|
||||
@@ -57,12 +57,13 @@ def rename_state_dict_key(k):
|
||||
# TODO(SS): one constant
|
||||
|
||||
|
||||
def convert_pegasus(tf_weights: dict, cfg_updates: dict) -> PegasusForConditionalGeneration:
|
||||
def convert_pegasus_to_bart(tf_weights: dict, cfg_updates: dict) -> PegasusForConditionalGeneration:
|
||||
cfg_kwargs = DEFAULTS.copy()
|
||||
cfg_kwargs.update(cfg_updates)
|
||||
cfg = PegasusConfig(**cfg_kwargs)
|
||||
torch_model = PegasusForConditionalGeneration(cfg)
|
||||
sd = torch_model.model.state_dict()
|
||||
|
||||
cfg = PegasusConfig(**cfg_updates)
|
||||
bart = PegasusForConditionalGeneration(cfg)
|
||||
sd = bart.model.state_dict()
|
||||
mapping = {}
|
||||
for k, v in tf_weights.items():
|
||||
new_k = rename_state_dict_key(k)
|
||||
@@ -79,13 +80,13 @@ def convert_pegasus(tf_weights: dict, cfg_updates: dict) -> PegasusForConditiona
|
||||
mapping["decoder.embed_tokens.weight"] = mapping["shared.weight"]
|
||||
empty_biases = {k: torch.zeros_like(v) for k, v in sd.items() if k.endswith("bias") and k not in mapping}
|
||||
mapping.update(**empty_biases)
|
||||
missing, extra = torch_model.model.load_state_dict(mapping, strict=False)
|
||||
missing, extra = bart.model.load_state_dict(mapping, strict=False)
|
||||
unexpected_missing = [
|
||||
k for k in missing if k not in ["encoder.embed_positions.weight", "decoder.embed_positions.weight"]
|
||||
]
|
||||
assert unexpected_missing == [], f"no matches found for the following torch keys {unexpected_missing}"
|
||||
assert extra == [], f"no matches found for the following tf keys {extra}"
|
||||
return torch_model
|
||||
return bart
|
||||
|
||||
|
||||
def get_tf_weights_as_numpy(path="./ckpt/aeslc/model.ckpt-32000") -> Dict:
|
||||
@@ -114,7 +115,7 @@ def convert_pegasus_ckpt_to_pytorch(ckpt_path: str, save_dir: str):
|
||||
cfg_updates = task_specific_params[f"summarization_{dataset}"]
|
||||
if dataset == "large":
|
||||
cfg_updates["task_specific_params"] = task_specific_params
|
||||
torch_model = convert_pegasus(tf_weights, cfg_updates)
|
||||
torch_model = convert_pegasus_to_bart(tf_weights, cfg_updates)
|
||||
torch_model.save_pretrained(save_dir)
|
||||
sd = torch_model.state_dict()
|
||||
sd.pop("model.decoder.embed_positions.weight")
|
||||
|
||||
@@ -505,11 +505,10 @@ class DataCollatorForNextSentencePrediction:
|
||||
# This should rarely go for more than one iteration for large
|
||||
# corpora. However, just to be careful, we try to make sure that
|
||||
# the random document is not the same as the document
|
||||
# we're processing. Also check to make sure that the random document
|
||||
# is not empty.
|
||||
# we're processing.
|
||||
for _ in range(10):
|
||||
random_document_index = random.randint(0, len(examples) - 1)
|
||||
if random_document_index != doc_index and len(examples[random_document_index]) > 0:
|
||||
if random_document_index != doc_index:
|
||||
break
|
||||
|
||||
random_document = examples[random_document_index]
|
||||
|
||||
@@ -69,6 +69,7 @@ try:
|
||||
import datasets # noqa: F401
|
||||
|
||||
_datasets_available = True
|
||||
logger.debug(f"Succesfully imported datasets version {datasets.__version__}")
|
||||
|
||||
except ImportError:
|
||||
_datasets_available = False
|
||||
@@ -119,6 +120,16 @@ try:
|
||||
except ImportError:
|
||||
_has_apex = False
|
||||
|
||||
|
||||
try:
|
||||
import faiss # noqa: F401
|
||||
|
||||
_faiss_available = True
|
||||
logger.debug(f"Succesfully imported faiss version {faiss.__version__}")
|
||||
except ImportError:
|
||||
_faiss_available = False
|
||||
|
||||
|
||||
default_cache_path = os.path.join(torch_cache_home, "transformers")
|
||||
|
||||
|
||||
@@ -141,10 +152,6 @@ DUMMY_MASK = [[1, 1, 1, 1, 1], [1, 1, 1, 0, 0], [0, 0, 0, 1, 1]]
|
||||
|
||||
S3_BUCKET_PREFIX = "https://s3.amazonaws.com/models.huggingface.co/bert"
|
||||
CLOUDFRONT_DISTRIB_PREFIX = "https://cdn.huggingface.co"
|
||||
PRESET_MIRROR_DICT = {
|
||||
"tuna": "https://mirrors.tuna.tsinghua.edu.cn/hugging-face-models",
|
||||
"bfsu": "https://mirrors.bfsu.edu.cn/hugging-face-models",
|
||||
}
|
||||
|
||||
|
||||
def is_torch_available():
|
||||
@@ -175,6 +182,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 "")
|
||||
@@ -574,7 +585,7 @@ def is_remote_url(url_or_filename):
|
||||
return parsed.scheme in ("http", "https")
|
||||
|
||||
|
||||
def hf_bucket_url(model_id: str, filename: str, use_cdn=True, mirror=None) -> str:
|
||||
def hf_bucket_url(model_id: str, filename: str, use_cdn=True) -> str:
|
||||
"""
|
||||
Resolve a model identifier, and a file name, to a HF-hosted url
|
||||
on either S3 or Cloudfront (a Content Delivery Network, or CDN).
|
||||
@@ -590,13 +601,7 @@ def hf_bucket_url(model_id: str, filename: str, use_cdn=True, mirror=None) -> st
|
||||
are not shared between the two because the cached file's name contains
|
||||
a hash of the url.
|
||||
"""
|
||||
endpoint = (
|
||||
PRESET_MIRROR_DICT.get(mirror, mirror)
|
||||
if mirror
|
||||
else CLOUDFRONT_DISTRIB_PREFIX
|
||||
if use_cdn
|
||||
else S3_BUCKET_PREFIX
|
||||
)
|
||||
endpoint = CLOUDFRONT_DISTRIB_PREFIX if use_cdn else S3_BUCKET_PREFIX
|
||||
legacy_format = "/" not in model_id
|
||||
if legacy_format:
|
||||
return f"{endpoint}/{model_id}-{filename}"
|
||||
|
||||
@@ -96,7 +96,7 @@ class TFGenerationMixin:
|
||||
Whether to stop the beam search when at least ``num_beams`` sentences are finished per batch or not.
|
||||
num_beams (:obj:`int`, `optional`, defaults to 1):
|
||||
Number of beams for beam search. 1 means no beam search.
|
||||
temperature (:obj:`float`, `optional`, defaults to 1.0):
|
||||
temperature (:obj:`float`, `optional`, defaults tp 1.0):
|
||||
The value used to module the next token probabilities.
|
||||
top_k (:obj:`int`, `optional`, defaults to 50):
|
||||
The number of highest probability vocabulary tokens to keep for top-k-filtering.
|
||||
|
||||
@@ -139,9 +139,7 @@ class ModelCard:
|
||||
elif os.path.isfile(pretrained_model_name_or_path) or is_remote_url(pretrained_model_name_or_path):
|
||||
model_card_file = pretrained_model_name_or_path
|
||||
else:
|
||||
model_card_file = hf_bucket_url(
|
||||
pretrained_model_name_or_path, filename=MODEL_CARD_NAME, use_cdn=False, mirror=None
|
||||
)
|
||||
model_card_file = hf_bucket_url(pretrained_model_name_or_path, filename=MODEL_CARD_NAME, use_cdn=False)
|
||||
|
||||
if find_from_standard_name or pretrained_model_name_or_path in ALL_PRETRAINED_CONFIG_ARCHIVE_MAP:
|
||||
model_card_file = model_card_file.replace(CONFIG_NAME, MODEL_CARD_NAME)
|
||||
|
||||
+687
-412
File diff suppressed because it is too large
Load Diff
@@ -272,6 +272,7 @@ class DPRPretrainedContextEncoder(PreTrainedModel):
|
||||
config_class = DPRConfig
|
||||
load_tf_weights = None
|
||||
base_model_prefix = "ctx_encoder"
|
||||
authorized_missing_keys = [r"position_ids"]
|
||||
|
||||
def init_weights(self):
|
||||
self.ctx_encoder.init_weights()
|
||||
@@ -285,6 +286,7 @@ class DPRPretrainedQuestionEncoder(PreTrainedModel):
|
||||
config_class = DPRConfig
|
||||
load_tf_weights = None
|
||||
base_model_prefix = "question_encoder"
|
||||
authorized_missing_keys = [r"position_ids"]
|
||||
|
||||
def init_weights(self):
|
||||
self.question_encoder.init_weights()
|
||||
@@ -298,6 +300,7 @@ class DPRPretrainedReader(PreTrainedModel):
|
||||
config_class = DPRConfig
|
||||
load_tf_weights = None
|
||||
base_model_prefix = "span_predictor"
|
||||
authorized_missing_keys = [r"position_ids"]
|
||||
|
||||
def init_weights(self):
|
||||
self.span_predictor.encoder.init_weights()
|
||||
|
||||
@@ -127,7 +127,7 @@ def load_tf_weights_in_funnel(model, config, tf_checkpoint_path):
|
||||
skipped = False
|
||||
for m_name in name[1:]:
|
||||
if not isinstance(pointer, FunnelPositionwiseFFN) and re.fullmatch(r"layer_\d+", m_name):
|
||||
layer_index = int(re.search(r"layer_(\d+)", m_name).groups()[0])
|
||||
layer_index = int(re.search("layer_(\d+)", m_name).groups()[0])
|
||||
if layer_index < config.num_hidden_layers:
|
||||
block_idx = 0
|
||||
while layer_index >= config.block_sizes[block_idx]:
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -21,7 +21,6 @@ from typing import Optional, Tuple
|
||||
|
||||
import tensorflow as tf
|
||||
|
||||
from .activations_tf import get_tf_activation
|
||||
from .configuration_albert import AlbertConfig
|
||||
from .file_utils import (
|
||||
MULTIPLE_CHOICE_DUMMY_INPUTS,
|
||||
@@ -31,7 +30,7 @@ from .file_utils import (
|
||||
add_start_docstrings_to_callable,
|
||||
replace_return_docstrings,
|
||||
)
|
||||
from .modeling_tf_bert import TFBertSelfAttention
|
||||
from .modeling_tf_bert import ACT2FN, TFBertSelfAttention
|
||||
from .modeling_tf_outputs import (
|
||||
TFBaseModelOutput,
|
||||
TFBaseModelOutputWithPooling,
|
||||
@@ -355,7 +354,7 @@ class TFAlbertLayer(tf.keras.layers.Layer):
|
||||
)
|
||||
|
||||
if isinstance(config.hidden_act, str):
|
||||
self.activation = get_tf_activation(config.hidden_act)
|
||||
self.activation = ACT2FN[config.hidden_act]
|
||||
else:
|
||||
self.activation = config.hidden_act
|
||||
|
||||
@@ -495,7 +494,7 @@ class TFAlbertMLMHead(tf.keras.layers.Layer):
|
||||
config.embedding_size, kernel_initializer=get_initializer(config.initializer_range), name="dense"
|
||||
)
|
||||
if isinstance(config.hidden_act, str):
|
||||
self.activation = get_tf_activation(config.hidden_act)
|
||||
self.activation = ACT2FN[config.hidden_act]
|
||||
else:
|
||||
self.activation = config.hidden_act
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -19,9 +19,9 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import numpy as np
|
||||
import tensorflow as tf
|
||||
|
||||
from .activations_tf import get_tf_activation
|
||||
from .configuration_bert import BertConfig
|
||||
from .file_utils import (
|
||||
MULTIPLE_CHOICE_DUMMY_INPUTS,
|
||||
@@ -88,6 +88,44 @@ TF_BERT_PRETRAINED_MODEL_ARCHIVE_LIST = [
|
||||
]
|
||||
|
||||
|
||||
def gelu(x):
|
||||
"""Gaussian Error Linear Unit.
|
||||
Original Implementation of the gelu activation function in Google Bert repo when initially created.
|
||||
For information: OpenAI GPT's gelu is slightly different (and gives slightly different results):
|
||||
0.5 * x * (1 + torch.tanh(math.sqrt(2 / math.pi) * (x + 0.044715 * torch.pow(x, 3))))
|
||||
Also see https://arxiv.org/abs/1606.08415
|
||||
"""
|
||||
cdf = 0.5 * (1.0 + tf.math.erf(x / tf.math.sqrt(2.0)))
|
||||
|
||||
return x * cdf
|
||||
|
||||
|
||||
def gelu_new(x):
|
||||
"""Gaussian Error Linear Unit.
|
||||
This is a smoother version of the RELU.
|
||||
Original paper: https://arxiv.org/abs/1606.08415
|
||||
Args:
|
||||
x: float Tensor to perform activation.
|
||||
Returns:
|
||||
`x` with the GELU activation applied.
|
||||
"""
|
||||
cdf = 0.5 * (1.0 + tf.tanh((np.sqrt(2 / np.pi) * (x + 0.044715 * tf.pow(x, 3)))))
|
||||
|
||||
return x * cdf
|
||||
|
||||
|
||||
def swish(x):
|
||||
return x * tf.sigmoid(x)
|
||||
|
||||
|
||||
ACT2FN = {
|
||||
"gelu": tf.keras.layers.Activation(gelu),
|
||||
"relu": tf.keras.activations.relu,
|
||||
"swish": tf.keras.layers.Activation(swish),
|
||||
"gelu_new": tf.keras.layers.Activation(gelu_new),
|
||||
}
|
||||
|
||||
|
||||
class TFBertEmbeddings(tf.keras.layers.Layer):
|
||||
"""Construct the embeddings from word, position and token_type embeddings."""
|
||||
|
||||
@@ -314,7 +352,7 @@ class TFBertIntermediate(tf.keras.layers.Layer):
|
||||
)
|
||||
|
||||
if isinstance(config.hidden_act, str):
|
||||
self.intermediate_act_fn = get_tf_activation(config.hidden_act)
|
||||
self.intermediate_act_fn = ACT2FN[config.hidden_act]
|
||||
else:
|
||||
self.intermediate_act_fn = config.hidden_act
|
||||
|
||||
@@ -429,7 +467,7 @@ class TFBertPredictionHeadTransform(tf.keras.layers.Layer):
|
||||
)
|
||||
|
||||
if isinstance(config.hidden_act, str):
|
||||
self.transform_act_fn = get_tf_activation(config.hidden_act)
|
||||
self.transform_act_fn = ACT2FN[config.hidden_act]
|
||||
else:
|
||||
self.transform_act_fn = config.hidden_act
|
||||
|
||||
|
||||
@@ -18,9 +18,9 @@
|
||||
|
||||
import math
|
||||
|
||||
import numpy as np
|
||||
import tensorflow as tf
|
||||
|
||||
from .activations_tf import get_tf_activation
|
||||
from .configuration_distilbert import DistilBertConfig
|
||||
from .file_utils import (
|
||||
MULTIPLE_CHOICE_DUMMY_INPUTS,
|
||||
@@ -68,6 +68,31 @@ TF_DISTILBERT_PRETRAINED_MODEL_ARCHIVE_LIST = [
|
||||
]
|
||||
|
||||
|
||||
# UTILS AND BUILDING BLOCKS OF THE ARCHITECTURE #
|
||||
def gelu(x):
|
||||
"""Gaussian Error Linear Unit.
|
||||
Original Implementation of the gelu activation function in Google Bert repo when initially created.
|
||||
For information: OpenAI GPT's gelu is slightly different (and gives slightly different results):
|
||||
0.5 * x * (1 + torch.tanh(math.sqrt(2 / math.pi) * (x + 0.044715 * torch.pow(x, 3))))
|
||||
Also see https://arxiv.org/abs/1606.08415
|
||||
"""
|
||||
cdf = 0.5 * (1.0 + tf.math.erf(x / tf.cast(tf.math.sqrt(2.0), dtype=x.dtype)))
|
||||
return x * cdf
|
||||
|
||||
|
||||
def gelu_new(x):
|
||||
"""Gaussian Error Linear Unit.
|
||||
This is a smoother version of the RELU.
|
||||
Original paper: https://arxiv.org/abs/1606.08415
|
||||
Args:
|
||||
x: float Tensor to perform activation.
|
||||
Returns:
|
||||
`x` with the GELU activation applied.
|
||||
"""
|
||||
cdf = 0.5 * (1.0 + tf.tanh((np.sqrt(2 / np.pi) * (x + 0.044715 * tf.pow(x, 3)))))
|
||||
return x * cdf
|
||||
|
||||
|
||||
class TFEmbeddings(tf.keras.layers.Layer):
|
||||
def __init__(self, config, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
@@ -273,7 +298,9 @@ class TFFFN(tf.keras.layers.Layer):
|
||||
assert config.activation in ["relu", "gelu"], "activation ({}) must be in ['relu', 'gelu']".format(
|
||||
config.activation
|
||||
)
|
||||
self.activation = get_tf_activation(config.activation)
|
||||
self.activation = (
|
||||
tf.keras.layers.Activation(gelu) if config.activation == "gelu" else tf.keras.activations.relu
|
||||
)
|
||||
|
||||
def call(self, input, training=False):
|
||||
x = self.lin1(input)
|
||||
@@ -624,7 +651,7 @@ class TFDistilBertForMaskedLM(TFDistilBertPreTrainedModel, TFMaskedLanguageModel
|
||||
self.vocab_transform = tf.keras.layers.Dense(
|
||||
config.dim, kernel_initializer=get_initializer(config.initializer_range), name="vocab_transform"
|
||||
)
|
||||
self.act = get_tf_activation("gelu")
|
||||
self.act = tf.keras.layers.Activation(gelu)
|
||||
self.vocab_layer_norm = tf.keras.layers.LayerNormalization(epsilon=1e-12, name="vocab_layer_norm")
|
||||
self.vocab_projector = TFDistilBertLMHead(config, self.distilbert.embeddings, name="vocab_projector")
|
||||
|
||||
|
||||
@@ -5,7 +5,6 @@ import tensorflow as tf
|
||||
|
||||
from transformers import ElectraConfig
|
||||
|
||||
from .activations_tf import get_tf_activation
|
||||
from .file_utils import (
|
||||
MULTIPLE_CHOICE_DUMMY_INPUTS,
|
||||
ModelOutput,
|
||||
@@ -14,7 +13,7 @@ from .file_utils import (
|
||||
add_start_docstrings_to_callable,
|
||||
replace_return_docstrings,
|
||||
)
|
||||
from .modeling_tf_bert import TFBertEncoder, TFBertPreTrainedModel
|
||||
from .modeling_tf_bert import ACT2FN, TFBertEncoder, TFBertPreTrainedModel
|
||||
from .modeling_tf_outputs import (
|
||||
TFBaseModelOutput,
|
||||
TFMaskedLMOutput,
|
||||
@@ -174,7 +173,7 @@ class TFElectraDiscriminatorPredictions(tf.keras.layers.Layer):
|
||||
|
||||
def call(self, discriminator_hidden_states, training=False):
|
||||
hidden_states = self.dense(discriminator_hidden_states)
|
||||
hidden_states = get_tf_activation(self.config.hidden_act)(hidden_states)
|
||||
hidden_states = ACT2FN[self.config.hidden_act](hidden_states)
|
||||
logits = tf.squeeze(self.dense_prediction(hidden_states))
|
||||
|
||||
return logits
|
||||
@@ -189,7 +188,7 @@ class TFElectraGeneratorPredictions(tf.keras.layers.Layer):
|
||||
|
||||
def call(self, generator_hidden_states, training=False):
|
||||
hidden_states = self.dense(generator_hidden_states)
|
||||
hidden_states = get_tf_activation("gelu")(hidden_states)
|
||||
hidden_states = ACT2FN["gelu"](hidden_states)
|
||||
hidden_states = self.LayerNorm(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
@@ -568,7 +567,7 @@ class TFElectraForMaskedLM(TFElectraPreTrainedModel, TFMaskedLanguageModelingLos
|
||||
self.electra = TFElectraMainLayer(config, name="electra")
|
||||
self.generator_predictions = TFElectraGeneratorPredictions(config, name="generator_predictions")
|
||||
if isinstance(config.hidden_act, str):
|
||||
self.activation = get_tf_activation(config.hidden_act)
|
||||
self.activation = ACT2FN[config.hidden_act]
|
||||
else:
|
||||
self.activation = config.hidden_act
|
||||
self.generator_lm_head = TFElectraMaskedLMHead(config, self.electra.embeddings, name="generator_lm_head")
|
||||
@@ -659,7 +658,7 @@ class TFElectraClassificationHead(tf.keras.layers.Layer):
|
||||
x = inputs[:, 0, :] # take <s> token (equiv. to [CLS])
|
||||
x = self.dropout(x)
|
||||
x = self.dense(x)
|
||||
x = get_tf_activation("gelu")(x) # although BERT uses tanh here, it seems Electra authors used gelu here
|
||||
x = ACT2FN["gelu"](x) # although BERT uses tanh here, it seems Electra authors used gelu here
|
||||
x = self.dropout(x)
|
||||
x = self.out_proj(x)
|
||||
|
||||
|
||||
@@ -19,7 +19,6 @@ from typing import Optional, Tuple
|
||||
|
||||
import tensorflow as tf
|
||||
|
||||
from .activations_tf import get_tf_activation
|
||||
from .configuration_funnel import FunnelConfig
|
||||
from .file_utils import (
|
||||
MULTIPLE_CHOICE_DUMMY_INPUTS,
|
||||
@@ -29,6 +28,7 @@ from .file_utils import (
|
||||
add_start_docstrings_to_callable,
|
||||
replace_return_docstrings,
|
||||
)
|
||||
from .modeling_tf_bert import ACT2FN
|
||||
from .modeling_tf_outputs import (
|
||||
TFBaseModelOutput,
|
||||
TFMaskedLMOutput,
|
||||
@@ -578,7 +578,7 @@ class TFFunnelPositionwiseFFN(tf.keras.layers.Layer):
|
||||
super().__init__(**kwargs)
|
||||
initializer = get_initializer(config.initializer_range)
|
||||
self.linear_1 = tf.keras.layers.Dense(config.d_inner, kernel_initializer=initializer, name="linear_1")
|
||||
self.activation_function = get_tf_activation(config.hidden_act)
|
||||
self.activation_function = ACT2FN[config.hidden_act]
|
||||
self.activation_dropout = tf.keras.layers.Dropout(config.activation_dropout)
|
||||
self.linear_2 = tf.keras.layers.Dense(config.d_model, kernel_initializer=initializer, name="linear_2")
|
||||
self.dropout = tf.keras.layers.Dropout(config.hidden_dropout)
|
||||
@@ -966,7 +966,7 @@ class TFFunnelDiscriminatorPredictions(tf.keras.layers.Layer):
|
||||
super().__init__(**kwargs)
|
||||
initializer = get_initializer(config.initializer_range)
|
||||
self.dense = tf.keras.layers.Dense(config.d_model, kernel_initializer=initializer, name="dense")
|
||||
self.activation_function = get_tf_activation(config.hidden_act)
|
||||
self.activation_function = ACT2FN[config.hidden_act]
|
||||
self.dense_prediction = tf.keras.layers.Dense(1, kernel_initializer=initializer, name="dense_prediction")
|
||||
|
||||
def call(self, discriminator_hidden_states):
|
||||
|
||||
@@ -19,9 +19,9 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
import numpy as np
|
||||
import tensorflow as tf
|
||||
|
||||
from .activations_tf import get_tf_activation
|
||||
from .configuration_gpt2 import GPT2Config
|
||||
from .file_utils import (
|
||||
ModelOutput,
|
||||
@@ -60,6 +60,19 @@ TF_GPT2_PRETRAINED_MODEL_ARCHIVE_LIST = [
|
||||
]
|
||||
|
||||
|
||||
def gelu(x):
|
||||
"""Gaussian Error Linear Unit.
|
||||
This is a smoother version of the RELU.
|
||||
Original paper: https://arxiv.org/abs/1606.08415
|
||||
Args:
|
||||
x: float Tensor to perform activation.
|
||||
Returns:
|
||||
`x` with the GELU activation applied.
|
||||
"""
|
||||
cdf = 0.5 * (1.0 + tf.tanh((np.sqrt(2 / np.pi) * (x + 0.044715 * tf.pow(x, 3)))))
|
||||
return x * cdf
|
||||
|
||||
|
||||
class TFAttention(tf.keras.layers.Layer):
|
||||
def __init__(self, nx, n_ctx, config, scale=False, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
@@ -167,7 +180,7 @@ class TFMLP(tf.keras.layers.Layer):
|
||||
nx = config.n_embd
|
||||
self.c_fc = TFConv1D(n_state, nx, initializer_range=config.initializer_range, name="c_fc")
|
||||
self.c_proj = TFConv1D(nx, n_state, initializer_range=config.initializer_range, name="c_proj")
|
||||
self.act = get_tf_activation("gelu")
|
||||
self.act = gelu
|
||||
self.dropout = tf.keras.layers.Dropout(config.resid_pdrop)
|
||||
|
||||
def call(self, x, training=False):
|
||||
|
||||
@@ -21,11 +21,11 @@ import logging
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict, Optional, Tuple
|
||||
|
||||
import numpy as np
|
||||
import tensorflow as tf
|
||||
|
||||
from transformers import BatchEncoding
|
||||
|
||||
from .activations_tf import get_tf_activation
|
||||
from .configuration_lxmert import LxmertConfig
|
||||
from .file_utils import (
|
||||
ModelOutput,
|
||||
@@ -48,6 +48,42 @@ TF_LXMERT_PRETRAINED_MODEL_ARCHIVE_LIST = [
|
||||
]
|
||||
|
||||
|
||||
def gelu(x):
|
||||
"""Gaussian Error Linear Unit.
|
||||
Original Implementation of the gelu activation function in Google Bert repo when initially created.
|
||||
For information: OpenAI GPT's gelu is slightly different (and gives slightly different results):
|
||||
0.5 * x * (1 + torch.tanh(math.sqrt(2 / math.pi) * (x + 0.044715 * torch.pow(x, 3))))
|
||||
Also see https://arxiv.org/abs/1606.08415
|
||||
"""
|
||||
cdf = 0.5 * (1.0 + tf.math.erf(x / tf.math.sqrt(2.0)))
|
||||
return x * cdf
|
||||
|
||||
|
||||
def gelu_new(x):
|
||||
"""Gaussian Error Linear Unit.
|
||||
This is a smoother version of the RELU.
|
||||
Original paper: https://arxiv.org/abs/1606.08415
|
||||
Args:
|
||||
x: float Tensor to perform activation.
|
||||
Returns:
|
||||
`x` with the GELU activation applied.
|
||||
"""
|
||||
cdf = 0.5 * (1.0 + tf.tanh((np.sqrt(2 / np.pi) * (x + 0.044715 * tf.pow(x, 3)))))
|
||||
return x * cdf
|
||||
|
||||
|
||||
def swish(x):
|
||||
return x * tf.sigmoid(x)
|
||||
|
||||
|
||||
ACT2FN = {
|
||||
"gelu": tf.keras.layers.Activation(gelu),
|
||||
"relu": tf.keras.activations.relu,
|
||||
"swish": tf.keras.layers.Activation(swish),
|
||||
"gelu_new": tf.keras.layers.Activation(gelu_new),
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class TFLxmertModelOutput(ModelOutput):
|
||||
"""
|
||||
@@ -368,7 +404,7 @@ class TFLxmertIntermediate(tf.keras.layers.Layer):
|
||||
name="dense",
|
||||
)
|
||||
if isinstance(config.hidden_act, str):
|
||||
self.intermediate_act_fn = get_tf_activation(config.hidden_act)
|
||||
self.intermediate_act_fn = ACT2FN[config.hidden_act]
|
||||
else:
|
||||
self.intermediate_act_fn = config.hidden_act
|
||||
|
||||
@@ -976,7 +1012,7 @@ class TFLxmertPredictionHeadTransform(tf.keras.layers.Layer):
|
||||
name="dense",
|
||||
)
|
||||
if isinstance(config.hidden_act, str):
|
||||
self.transform_act_fn = get_tf_activation(config.hidden_act)
|
||||
self.transform_act_fn = ACT2FN[config.hidden_act]
|
||||
else:
|
||||
self.transform_act_fn = config.hidden_act
|
||||
self.LayerNorm = tf.keras.layers.LayerNormalization(epsilon=config.layer_norm_eps, name="LayerNorm")
|
||||
@@ -1046,7 +1082,7 @@ class TFLxmertVisualAnswerHead(tf.keras.layers.Layer):
|
||||
kernel_initializer=get_initializer(config.initializer_range),
|
||||
name="logit_fc_._0",
|
||||
)
|
||||
self.activation = get_tf_activation("gelu")
|
||||
self.activation = tf.keras.layers.Activation(gelu)
|
||||
self.layer_norm = tf.keras.layers.LayerNormalization(epsilon=config.layer_norm_eps, name="logit_fc_._2")
|
||||
self.dense_1 = tf.keras.layers.Dense(
|
||||
num_labels,
|
||||
|
||||
@@ -22,7 +22,6 @@ from typing import Optional, Tuple
|
||||
import tensorflow as tf
|
||||
|
||||
from . import MobileBertConfig
|
||||
from .activations_tf import get_tf_activation
|
||||
from .file_utils import (
|
||||
MULTIPLE_CHOICE_DUMMY_INPUTS,
|
||||
ModelOutput,
|
||||
@@ -31,7 +30,7 @@ from .file_utils import (
|
||||
add_start_docstrings_to_callable,
|
||||
replace_return_docstrings,
|
||||
)
|
||||
from .modeling_tf_bert import TFBertIntermediate
|
||||
from .modeling_tf_bert import TFBertIntermediate, gelu, gelu_new, swish
|
||||
from .modeling_tf_outputs import (
|
||||
TFBaseModelOutput,
|
||||
TFBaseModelOutputWithPooling,
|
||||
@@ -68,6 +67,10 @@ TF_MOBILEBERT_PRETRAINED_MODEL_ARCHIVE_LIST = [
|
||||
]
|
||||
|
||||
|
||||
def mish(x):
|
||||
return x * tf.tanh(tf.math.softplus(x))
|
||||
|
||||
|
||||
class TFLayerNorm(tf.keras.layers.LayerNormalization):
|
||||
def __init__(self, feat_size, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
@@ -86,6 +89,12 @@ class TFNoNorm(tf.keras.layers.Layer):
|
||||
return inputs * self.weight + self.bias
|
||||
|
||||
|
||||
ACT2FN = {
|
||||
"gelu": tf.keras.layers.Activation(gelu),
|
||||
"relu": tf.keras.activations.relu,
|
||||
"swish": tf.keras.layers.Activation(swish),
|
||||
"gelu_new": tf.keras.layers.Activation(gelu_new),
|
||||
}
|
||||
NORM2FN = {"layer_norm": TFLayerNorm, "no_norm": TFNoNorm}
|
||||
|
||||
|
||||
@@ -612,7 +621,7 @@ class TFMobileBertPredictionHeadTransform(tf.keras.layers.Layer):
|
||||
config.hidden_size, kernel_initializer=get_initializer(config.initializer_range), name="dense"
|
||||
)
|
||||
if isinstance(config.hidden_act, str):
|
||||
self.transform_act_fn = get_tf_activation(config.hidden_act)
|
||||
self.transform_act_fn = ACT2FN[config.hidden_act]
|
||||
else:
|
||||
self.transform_act_fn = config.hidden_act
|
||||
self.LayerNorm = NORM2FN["layer_norm"](config.hidden_size, epsilon=config.layer_norm_eps, name="LayerNorm")
|
||||
|
||||
@@ -19,9 +19,9 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import numpy as np
|
||||
import tensorflow as tf
|
||||
|
||||
from .activations_tf import get_tf_activation
|
||||
from .configuration_openai import OpenAIGPTConfig
|
||||
from .file_utils import (
|
||||
ModelOutput,
|
||||
@@ -56,6 +56,30 @@ TF_OPENAI_GPT_PRETRAINED_MODEL_ARCHIVE_LIST = [
|
||||
]
|
||||
|
||||
|
||||
def gelu(x):
|
||||
"""Gaussian Error Linear Unit.
|
||||
This is a smoother version of the RELU.
|
||||
Original paper: https://arxiv.org/abs/1606.08415
|
||||
Args:
|
||||
x: float Tensor to perform activation.
|
||||
Returns:
|
||||
`x` with the GELU activation applied.
|
||||
"""
|
||||
cdf = 0.5 * (1.0 + tf.tanh((np.sqrt(2 / np.pi) * (x + 0.044715 * tf.pow(x, 3)))))
|
||||
return x * cdf
|
||||
|
||||
|
||||
def swish(x):
|
||||
return x * tf.math.sigmoid(x)
|
||||
|
||||
|
||||
ACT_FNS = {
|
||||
"gelu": tf.keras.layers.Activation(gelu),
|
||||
"relu": tf.keras.activations.relu,
|
||||
"swish": tf.keras.layers.Activation(swish),
|
||||
}
|
||||
|
||||
|
||||
class TFAttention(tf.keras.layers.Layer):
|
||||
def __init__(self, nx, n_ctx, config, scale=False, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
@@ -155,7 +179,7 @@ class TFMLP(tf.keras.layers.Layer):
|
||||
nx = config.n_embd
|
||||
self.c_fc = TFConv1D(n_state, nx, initializer_range=config.initializer_range, name="c_fc")
|
||||
self.c_proj = TFConv1D(nx, n_state, initializer_range=config.initializer_range, name="c_proj")
|
||||
self.act = get_tf_activation("gelu")
|
||||
self.act = gelu
|
||||
self.dropout = tf.keras.layers.Dropout(config.resid_pdrop)
|
||||
|
||||
def call(self, x, training=False):
|
||||
|
||||
@@ -18,7 +18,6 @@
|
||||
|
||||
import tensorflow as tf
|
||||
|
||||
from .activations_tf import get_tf_activation
|
||||
from .configuration_roberta import RobertaConfig
|
||||
from .file_utils import (
|
||||
MULTIPLE_CHOICE_DUMMY_INPUTS,
|
||||
@@ -26,7 +25,7 @@ from .file_utils import (
|
||||
add_start_docstrings,
|
||||
add_start_docstrings_to_callable,
|
||||
)
|
||||
from .modeling_tf_bert import TFBertEmbeddings, TFBertMainLayer
|
||||
from .modeling_tf_bert import TFBertEmbeddings, TFBertMainLayer, gelu
|
||||
from .modeling_tf_outputs import (
|
||||
TFBaseModelOutputWithPooling,
|
||||
TFMaskedLMOutput,
|
||||
@@ -238,7 +237,7 @@ class TFRobertaLMHead(tf.keras.layers.Layer):
|
||||
config.hidden_size, kernel_initializer=get_initializer(config.initializer_range), name="dense"
|
||||
)
|
||||
self.layer_norm = tf.keras.layers.LayerNormalization(epsilon=config.layer_norm_eps, name="layer_norm")
|
||||
self.act = get_tf_activation("gelu")
|
||||
self.act = tf.keras.layers.Activation(gelu)
|
||||
|
||||
# The output weights are the same as the input embeddings, but there is
|
||||
# an output-only bias for each token.
|
||||
|
||||
@@ -15,7 +15,8 @@
|
||||
# limitations under the License.
|
||||
""" TF 2.0 Transformer XL model.
|
||||
"""
|
||||
import warnings
|
||||
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
@@ -106,7 +107,10 @@ class TFRelPartialLearnableMultiHeadAttn(tf.keras.layers.Layer):
|
||||
d_model,
|
||||
d_head,
|
||||
dropout,
|
||||
dropatt=0.0,
|
||||
dropatt=0,
|
||||
tgt_len=None,
|
||||
ext_len=None,
|
||||
mem_len=None,
|
||||
pre_lnorm=False,
|
||||
r_r_bias=None,
|
||||
r_w_bias=None,
|
||||
@@ -257,6 +261,9 @@ class TFRelPartialLearnableDecoderLayer(tf.keras.layers.Layer):
|
||||
d_head,
|
||||
d_inner,
|
||||
dropout,
|
||||
tgt_len=None,
|
||||
ext_len=None,
|
||||
mem_len=None,
|
||||
dropatt=0.0,
|
||||
pre_lnorm=False,
|
||||
r_w_bias=None,
|
||||
@@ -273,6 +280,9 @@ class TFRelPartialLearnableDecoderLayer(tf.keras.layers.Layer):
|
||||
d_model,
|
||||
d_head,
|
||||
dropout,
|
||||
tgt_len=tgt_len,
|
||||
ext_len=ext_len,
|
||||
mem_len=mem_len,
|
||||
dropatt=dropatt,
|
||||
pre_lnorm=pre_lnorm,
|
||||
r_w_bias=r_w_bias,
|
||||
@@ -404,7 +414,12 @@ class TFTransfoXLMainLayer(tf.keras.layers.Layer):
|
||||
self.drop = tf.keras.layers.Dropout(config.dropout)
|
||||
|
||||
self.n_layer = config.n_layer
|
||||
|
||||
self.tgt_len = config.tgt_len
|
||||
self.mem_len = config.mem_len
|
||||
self.ext_len = config.ext_len
|
||||
self.max_klen = config.tgt_len + config.ext_len + config.mem_len
|
||||
|
||||
self.attn_type = config.attn_type
|
||||
|
||||
self.layers = []
|
||||
@@ -417,6 +432,9 @@ class TFTransfoXLMainLayer(tf.keras.layers.Layer):
|
||||
config.d_head,
|
||||
config.d_inner,
|
||||
config.dropout,
|
||||
tgt_len=config.tgt_len,
|
||||
ext_len=config.ext_len,
|
||||
mem_len=config.mem_len,
|
||||
dropatt=config.dropatt,
|
||||
pre_lnorm=config.pre_lnorm,
|
||||
r_w_bias=None if self.untie_r else self.r_w_bias,
|
||||
@@ -460,8 +478,10 @@ class TFTransfoXLMainLayer(tf.keras.layers.Layer):
|
||||
def backward_compatible(self):
|
||||
self.sample_softmax = -1
|
||||
|
||||
def reset_memory_length(self, mem_len):
|
||||
def reset_length(self, tgt_len, ext_len, mem_len):
|
||||
self.tgt_len = tgt_len
|
||||
self.mem_len = mem_len
|
||||
self.ext_len = ext_len
|
||||
|
||||
def _prune_heads(self, heads):
|
||||
raise NotImplementedError
|
||||
@@ -486,8 +506,12 @@ class TFTransfoXLMainLayer(tf.keras.layers.Layer):
|
||||
assert len(hids) == len(mems), "len(hids) != len(mems)"
|
||||
|
||||
# There are `mlen + qlen` steps that can be cached into mems
|
||||
# For the next step, the last `ext_len` of the `qlen` tokens
|
||||
# will be used as the extended context. Hence, we only cache
|
||||
# the tokens from `mlen + qlen - self.ext_len - self.mem_len`
|
||||
# to `mlen + qlen - self.ext_len`.
|
||||
new_mems = []
|
||||
end_idx = mlen + max(0, qlen)
|
||||
end_idx = mlen + max(0, qlen - 0 - self.ext_len)
|
||||
beg_idx = max(0, end_idx - self.mem_len)
|
||||
for i in range(len(hids)):
|
||||
|
||||
@@ -843,14 +867,7 @@ class TFTransfoXLLMHeadModel(TFTransfoXLPreTrainedModel):
|
||||
return None
|
||||
|
||||
def reset_length(self, tgt_len, ext_len, mem_len):
|
||||
warnings.warn(
|
||||
"The method `reset_length` is deprecated and will be removed in a future version, use `reset_memory_length` instead.",
|
||||
FutureWarning,
|
||||
)
|
||||
self.transformer.reset_memory_length(mem_len)
|
||||
|
||||
def reset_memory_length(self, mem_len):
|
||||
self.transformer.reset_memory_length(mem_len)
|
||||
self.transformer.reset_length(tgt_len, ext_len, mem_len)
|
||||
|
||||
def init_mems(self, bsz):
|
||||
return self.transformer.init_mems(bsz)
|
||||
|
||||
@@ -484,10 +484,6 @@ class TFPreTrainedModel(tf.keras.Model, TFModelUtilsMixin, TFGenerationMixin):
|
||||
use_cdn(:obj:`bool`, `optional`, defaults to :obj:`True`):
|
||||
Whether or not to use Cloudfront (a Content Delivery Network, or CDN) when searching for the model on
|
||||
our S3 (faster). Should be set to :obj:`False` for checkpoints larger than 20GB.
|
||||
mirror(:obj:`str`, `optional`, defaults to :obj:`None`):
|
||||
Mirror source to accelerate downloads in China. If you are from China and have an accessibility problem,
|
||||
you can set this option to resolve it. Note that we do not guarantee the timeliness or safety. Please
|
||||
refer to the mirror site for more information.
|
||||
kwargs (remaining dictionary of keyword arguments, `optional`):
|
||||
Can be used to update the configuration object (after it being loaded) and initiate the model (e.g.,
|
||||
:obj:`output_attentions=True`). Behaves differently depending on whether a ``config`` is provided or
|
||||
@@ -526,7 +522,6 @@ class TFPreTrainedModel(tf.keras.Model, TFModelUtilsMixin, TFGenerationMixin):
|
||||
output_loading_info = kwargs.pop("output_loading_info", False)
|
||||
local_files_only = kwargs.pop("local_files_only", False)
|
||||
use_cdn = kwargs.pop("use_cdn", True)
|
||||
mirror = kwargs.pop("mirror", None)
|
||||
|
||||
# Load config if we don't provide a configuration
|
||||
if not isinstance(config, PretrainedConfig):
|
||||
@@ -569,7 +564,6 @@ class TFPreTrainedModel(tf.keras.Model, TFModelUtilsMixin, TFGenerationMixin):
|
||||
pretrained_model_name_or_path,
|
||||
filename=(WEIGHTS_NAME if from_pt else TF2_WEIGHTS_NAME),
|
||||
use_cdn=use_cdn,
|
||||
mirror=mirror,
|
||||
)
|
||||
|
||||
try:
|
||||
@@ -705,7 +699,7 @@ class TFConv1D(tf.keras.layers.Layer):
|
||||
|
||||
|
||||
class TFSharedEmbeddings(tf.keras.layers.Layer):
|
||||
r"""
|
||||
"""
|
||||
Construct shared token embeddings.
|
||||
|
||||
The weights of the embedding layer is usually shared with the weights of the linear decoder when doing
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user