Compare commits

...
168 Commits
Author SHA1 Message Date
Sam Shleifer 34eb3fcd55 Fix automodel test 2020-09-28 01:28:17 -04:00
Sam Shleifer 1ab6644b8b Remove unneeded bart changes 2020-09-28 01:21:09 -04:00
Sam Shleifer 8a1b6c2439 skip 2 flaky tests 2020-09-28 01:13:03 -04:00
Sam Shleifer 18e736f5e3 common tests passing on cpu 2020-09-28 01:09:57 -04:00
Sam Shleifer 0352ce44b2 leading space 2020-09-28 01:02:46 -04:00
Sam Shleifer ba149a6ef1 style 2020-09-28 01:01:21 -04:00
Sam Shleifer c0bbf1178e replace layernorm_variant with another boolean 2020-09-28 01:00:16 -04:00
Sam Shleifer e136853ca4 2 failing 2020-09-27 22:30:53 -04:00
Sam Shleifer 50d93bb051 tests passing 2020-09-27 22:13:30 -04:00
Sam Shleifer ec73340b2f fix tokenizer 2020-09-27 22:05:39 -04:00
Sam Shleifer 72f17d24c5 save an if 2020-09-27 21:53:22 -04:00
Sam Shleifer 5668d35e36 Failing 2020-09-27 19:03:35 -04:00
Sam Shleifer fc1af2377d merged master 2020-09-27 16:32:24 -04:00
Sam Shleifer e9ddee9208 Progress 2020-09-27 16:32:09 -04:00
Sam Shleifer 7ff0aea1d9 test cleanup 2020-09-25 10:13:26 -04:00
Sam Shleifer 9ff0d016e8 boom boom 2020-09-24 22:37:17 -04:00
Sam Shleifer 3a1a2ee9bf Fix mbart with config.norm_embed_before 2020-09-24 22:21:23 -04:00
Sam Shleifer e129f240e5 boom boom 2020-09-24 21:52:26 -04:00
Sam Shleifer 3c0b8e3b36 Fix tests 2020-09-24 21:43:46 -04:00
Sam Shleifer 1a3b40a1e0 add assert 2020-09-24 21:23:36 -04:00
Sam Shleifer 57fed3ff3a cleaner assert 2020-09-24 21:20:13 -04:00
Sam Shleifer 11bc12f98a fix pegasus by adding 3rd layernorm 2020-09-24 20:48:42 -04:00
Sam Shleifer 5b2ace0c98 boom boom 2020-09-24 20:02:55 -04:00
Sam Shleifer ca3c39e3a4 Mbart passing 2020-09-24 20:00:59 -04:00
Sam Shleifer 4d65a5dbe0 Merge branch 'master' into blenderbot 2020-09-24 19:38:35 -04:00
Sam Shleifer 77cd835a06 remove normalize before 2020-09-24 19:33:40 -04:00
Sam Shleifer e5efa0bb6a still breaking pegasus 2020-09-24 19:08:07 -04:00
Sam Shleifer 371f1ba53e Return dict 2020-09-24 18:06:20 -04:00
Sam Shleifer 4f635ae946 style 2020-09-23 18:43:19 -04:00
Sam Shleifer f0483a09c6 Merge master 2020-09-23 18:42:56 -04:00
Sam Shleifer 642e51b9c5 2 failures 2020-09-19 18:55:50 -04:00
Sam Shleifer 69b4921c12 style 2020-09-19 18:53:59 -04:00
Sam Shleifer 2715b25989 add semi-broken tokenizer test 2020-09-18 00:42:32 -04:00
Sam Shleifer 588bdf747a closer 2020-09-18 00:34:28 -04:00
Sam Shleifer d9d340998a Delete some copy pasted code 2020-09-18 00:14:24 -04:00
Sam Shleifer eb4cd09143 cleanup 2020-09-17 23:59:34 -04:00
Sam Shleifer 961af20683 Refactor to 1 file 2020-09-17 22:54:58 -04:00
Sam Shleifer 751e95cc88 Merge branch 'master' into blenderbot 2020-09-17 21:44:07 -04:00
Sam Shleifer 6ff219825d Merged master 2020-09-17 11:32:12 -04:00
Mariama Drame 19496e6c1b fix typos 2020-08-14 16:52:58 +02:00
Mariama Drame 8e61b730f4 fix typos 2020-08-14 16:49:52 +02:00
Mariama Drame c03bc9932e fix typos 2020-08-14 16:48:32 +02:00
Mariama Drame 7b42e87835 fix typos 2020-08-14 16:39:49 +02:00
Mariama Drame 12659ba87a fix typo 2020-08-14 16:29:22 +02:00
Mariama Drame 90f1210148 fix typo 2020-08-14 15:50:18 +02:00
Mariama Drame c986d7fd28 fix typos 2020-08-14 15:45:02 +02:00
Mariama Drame 9cfd4dc0ea fix typos 2020-08-14 14:22:13 +02:00
Mariama Drame 99bc11119a fix typos 2020-08-14 12:45:56 +02:00
Mariama Drame f7c8b6018e fix typos 2020-08-14 12:43:09 +02:00
Mariama Drame ab98343c08 fix typos 2020-08-14 12:27:19 +02:00
Mariama Drame 6eab35737d fix typos 2020-08-14 12:13:25 +02:00
Mariama Drame bcd9cf6813 fix typos 2020-08-14 12:04:49 +02:00
Mariama Drame b18f28d2e5 fix typos 2020-08-14 11:15:51 +02:00
Mariama Drame 537a8febca fix typos 2020-08-14 11:06:17 +02:00
Mariama Drame 6c7355bd5f fix typos 2020-08-14 10:57:28 +02:00
Mariama Drame 9379a7720c fix typos 2020-08-14 10:45:11 +02:00
Mariama Drame 3f7883b350 fix typos 2020-08-14 10:30:51 +02:00
Mariama Drame ce8fa251b7 fix typos 2020-08-14 10:09:13 +02:00
mariamabarham 907331ffe2 Merge branch 'master' into blenderbot 2020-08-14 09:58:58 +02:00
Mariama Drame 851a371c2c fix typos 2020-08-14 09:58:09 +02:00
Mariama Drame 3ad5f3401b fix typos 2020-08-14 09:50:16 +02:00
Mariama Drame ea57d77150 fix typos 2020-08-14 09:31:34 +02:00
Mariama Drame b33cad7edd fix typos 2020-08-14 09:21:04 +02:00
Mariama Drame 964ce608d5 fix typos 2020-08-14 08:58:33 +02:00
Mariama Drame 3447c2aeb4 fix typos 2020-08-13 23:17:52 +02:00
Mariama Drame 4dd0300e8c fix typos 2020-08-13 21:55:04 +02:00
Mariama Drame 2702c37634 Merge branch 'blenderbot' of https://github.com/huggingface/transformers into blenderbot 2020-08-13 18:54:40 +02:00
Mariama Drame 1f6071f76b fix typos 2020-08-13 18:50:44 +02:00
mariamabarham 00a73d3809 Merge branch 'master' into blenderbot 2020-08-13 18:36:35 +02:00
Mariama Drame 9f4bfa55ea fix typos 2020-08-13 18:35:15 +02:00
Mariama Drame 3bcff4fcdd fix typos 2020-08-13 18:32:12 +02:00
Mariama Drame 5dbc2e9d09 fix typos 2020-08-13 17:36:23 +02:00
Mariama Drame f5169f3907 fix typos 2020-08-13 17:16:02 +02:00
Mariama Drame 6ac9174299 fix typos 2020-08-13 17:15:45 +02:00
Mariama Drame d286fc9c4c fix some tests 2020-08-13 16:27:09 +02:00
Mariama Drame d6bfa8c823 fix code style 2020-08-13 12:00:36 +02:00
Mariama Drame d53c539176 fix typos 2020-08-13 11:53:37 +02:00
Mariama Drame 3791efa561 remove trailling whitespace 2020-08-13 11:34:52 +02:00
Mariama Drame 687ae67cc7 fix make html 2020-08-13 11:31:26 +02:00
Mariama Drame 06d391e94d fix make html 2020-08-13 11:27:25 +02:00
Mariama Drame 376ff07c85 fix make html 2020-08-13 11:14:16 +02:00
Mariama Drame 77d2d3dfe3 fix make html 2020-08-13 10:28:02 +02:00
Mariama Drame b0a57d443a fix code style 2020-08-13 09:38:20 +02:00
Mariama Drame 5c6ec437be add blenderbot doc 2020-08-13 09:11:52 +02:00
Mariama Drame 213b49c3e9 fix code style 2020-08-12 23:33:17 +02:00
Mariama Drame 0b2e2bf599 fix code style 2020-08-12 23:26:02 +02:00
Mariama Drame 9399b1ff95 fix code style 2020-08-12 23:15:43 +02:00
Mariama Drame f7d347df0e fix code style 2020-08-12 22:32:47 +02:00
Mariama Drame ed8e6baf7c fix code style 2020-08-12 22:08:32 +02:00
Mariama Drame 3ac98d0384 fix code style 2020-08-12 21:55:31 +02:00
Mariama Drame 302b03596d fix code style 2020-08-12 21:36:38 +02:00
Mariama Drame 9e91f69583 fix code style 2020-08-12 21:13:12 +02:00
Mariama Drame 3d1394429a fix code style 2020-08-12 18:57:47 +02:00
Mariama Drame c332e6a7d8 fix code style 2020-08-12 18:50:54 +02:00
Mariama Drame 5c78acd5ed fix code style 2020-08-12 18:47:12 +02:00
Mariama Drame 2f6fdb375b fix code quality 2020-08-12 18:41:29 +02:00
Mariama Drame c57c4b824f reformat line length 2020-08-12 18:30:38 +02:00
Mariama Drame 4a818f3ff0 fix typo 2020-08-12 18:12:04 +02:00
Mariama Drame 833f706880 remove test_parity_blenderbot 2020-08-12 18:03:09 +02:00
Mariama Drame 83493999be fix circle CI line length 2020-08-12 17:52:49 +02:00
Mariama Drame bdbbe24462 fix typos 2020-08-12 17:46:21 +02:00
Mariama Drame a85f5a4a27 fix configuration args 2020-08-12 17:44:46 +02:00
Mariama Drame a25426f680 fix configuration args 2020-08-12 17:38:49 +02:00
Mariama Drame e24acdef49 change return_tuple to return_dict 2020-08-12 17:22:21 +02:00
Mariama Drame 068a9eea01 add unittest 2020-08-12 17:15:49 +02:00
Mariama Drame 26e3bdba2b update blenderbot generator model to inherit on Bart 2020-08-12 17:12:34 +02:00
Mariama Drame 875b5b0359 first draft for blenderbot generator model 2020-08-12 17:07:20 +02:00
Mariama Drame 14aadf410c change return_tuple to return_dict 2020-08-12 16:50:42 +02:00
Mariama Drame c8463845be fix some typo and move pretrained weight to facebook/blenderbot 2020-08-12 15:31:25 +02:00
Mariama Drame 30e26004a0 fix typos 2020-08-11 11:59:02 +02:00
Mariama Drame d074f58194 remove parlai dependencies and fix some typos 2020-08-11 09:40:45 +02:00
Mariama Drame c857cf163d add forward tests 2020-08-10 16:21:40 +02:00
Mariama Drame a08cdfd577 keeping return_tuple for now 2020-08-10 15:36:42 +02:00
Mariama Drame c3a40f2e1c remove return_tuple and some Bart depnedences 2020-08-10 15:22:36 +02:00
Mariama Drame d802181c95 working test 2020-08-05 16:06:32 +02:00
Mariama Drame 9e9c50ae3a update moceling_bart to remove an extra normalize step when variant is prelayernorm 2020-07-27 14:02:25 +02:00
Mariama Drame 9bf8df4bbc update moceling_bart to remove an extra normalize step when variant is prelayernorm 2020-07-27 13:55:05 +02:00
Mariama Drame da04344f13 Merge remote-tracking branch 'origin/master' into blenderbot 2020-07-23 10:45:16 +02:00
Mariama Drame a82b85475d some updates 2020-07-23 10:41:20 +02:00
Mariama Drame da0c4f8bb4 remove prints 2020-07-16 16:57:12 +02:00
Mariama Drame df5bf64bb2 update small model tokenizer 2020-07-16 16:54:21 +02:00
Sam Shleifer fb1f0b4209 Merge branch 'master' into blenderbot 2020-07-14 06:14:05 -04:00
Sam Shleifer e50807900f add failing slow tests 2020-07-13 12:46:34 -04:00
Sam Shleifer 07e1a3b0e6 Simplify config 2020-07-13 12:10:11 -04:00
Sam Shleifer 6a6527e702 Merge branch 'master' into blenderbot 2020-07-13 11:57:21 -04:00
Sam Shleifer 79e6c39d29 Fix tok 2020-07-13 11:40:30 -04:00
Sam Shleifer 08d3f716bd remove bad imports 2020-07-13 10:24:40 -04:00
Sam Shleifer 3e933766c9 Fix is_valid_mbart 2020-07-13 10:20:50 -04:00
Mariama Drame 724dc87981 add tokenizer test 2020-07-13 13:24:36 +02:00
Sam Shleifer 84dc65ea09 fix mbart 2020-07-12 23:21:15 -04:00
Sam Shleifer 9d7908939e Working 90m model 2020-07-12 23:06:20 -04:00
Sam Shleifer 08e630e1be passing 2020-07-12 23:05:31 -04:00
Sam Shleifer bb988e901d Merge branch 'master' into blenderbot 2020-07-12 22:58:19 -04:00
Sam Shleifer b77daa402d merged sylvain change 2020-07-11 01:54:56 -04:00
Sam Shleifer f8cb16c541 Merge branch 'master' into blenderbot 2020-07-11 01:42:55 -04:00
Sam Shleifer 6367bf8d99 Update integration tests 2020-07-10 11:39:19 -04:00
Mariama Drame 1fbf338417 test blenderbot small tokenizer 2020-07-10 15:13:03 +02:00
Mariama Drame 7a5a9d9018 Merge remote-tracking branch 'origin/master' into blenderbot 2020-07-10 14:20:29 +02:00
Mariama Drame 9ac911271b fixe some files typo 2020-07-10 12:37:43 +02:00
Mariama Drame 63a799c3d0 add tokenizer for 90M model 2020-07-10 12:28:00 +02:00
Sam Shleifer 766a65b193 More changes 2020-07-10 03:27:51 -04:00
Sam Shleifer 1a2db58dd5 Merge branch 'master' into blenderbot 2020-07-10 03:07:25 -04:00
sshleifer a26f8d28a6 Merge branch 'master' into blenderbot 2020-07-08 10:30:45 -04:00
Mariama Drame 05a262cff3 add unittest 2020-07-08 15:58:32 +02:00
Mariama Drame 9738920b6f Merge remote-tracking branch 'origin/master' into blenderbot 2020-07-01 18:53:39 +02:00
Mariama Drame 36c8039e93 some updates 2020-07-01 18:50:15 +02:00
sshleifer f6de270f68 Fix config 2020-06-29 14:14:57 -04:00
sshleifer 7f31ebfd34 19 passing, 3 failing 2020-06-29 14:02:28 -04:00
sshleifer 958ca055fb Fix some tests 2020-06-29 13:47:43 -04:00
sshleifer 6d93aa4221 Merge branch 'master' into blenderbot 2020-06-29 13:38:42 -04:00
sshleifer 6d2f767b12 pre-merge 2020-06-29 13:38:29 -04:00
sshleifer 3eb866c1e2 Merge branch 'blenderbot' of github.com:huggingface/transformers into blenderbot 2020-06-29 13:17:02 -04:00
Mariama Drame 0a0cbe48f1 fix config file 2020-06-29 13:16:44 -04:00
Mariama Drame 673646553e add unittest 2020-06-29 13:16:44 -04:00
Mariama Drame 83ffc33ae5 update blenderbot generator model to inherit on Bart 2020-06-29 13:16:44 -04:00
Mariama Drame 1053e64318 fix typos 2020-06-29 13:16:44 -04:00
Mariama Drame 335e794c86 first draft for blenderbot generator model 2020-06-29 13:16:44 -04:00
Mariama Drame 1b1b9e7968 fix config file 2020-06-24 10:36:18 +02:00
Mariama Drame 0472762c8a add unittest 2020-06-24 10:36:18 +02:00
Mariama Drame b55db43c42 update blenderbot generator model to inherit on Bart 2020-06-24 10:36:18 +02:00
Mariama Drame 9894cd388b fix typos 2020-06-24 10:36:18 +02:00
Mariama Drame 1e1556a57d first draft for blenderbot generator model 2020-06-24 10:36:18 +02:00
Mariama Drame 237652b61d Merge remote-tracking branch 'origin/master' into blenderbot 2020-06-23 17:51:31 +02:00
Mariama Drame 6c7e29c01c fix config file 2020-06-17 15:10:54 +02:00
Mariama Drame 9f6c49fbc2 add unittest 2020-06-17 14:26:36 +02:00
Mariama Drame 775e03b7b6 update blenderbot generator model to inherit on Bart 2020-06-11 10:37:08 +02:00
Mariama Drame c7a7a5a01f fix typos 2020-06-05 21:18:10 +02:00
Mariama Drame c71834634f first draft for blenderbot generator model 2020-06-05 20:42:31 +02:00
32 changed files with 1104 additions and 41 deletions
+5 -3
View File
@@ -137,12 +137,13 @@ conversion utilities for the following models:
27. `Bert For Sequence Generation <https://tfhub.dev/s?module-type=text-generation&subtype=module,placeholder>`_ (from Google) released with the paper
`Leveraging Pre-trained Checkpoints for Sequence Generation Tasks
<https://arxiv.org/abs/1907.12461>`_ by Sascha Rothe, Shashi Narayan, Aliaksei Severyn.
28. `LayoutLM <https://github.com/microsoft/unilm/tree/master/layoutlm>`_ (from Microsoft Research Asia) released with the paper
28. `Blenderbot <https://github.com/facebookresearch/ParlAI>`_ (from Facebook AI Research) released with the paper `Recipes for building an open-domain chatbot
<https://arxiv.org/abs/2004.13637>`_ by Stephen Roller, Emily Dinan, Naman Goyal, Da Ju, Mary Williamson, Yinhan Liu, Jing Xu, Myle Ott, Kurt Shuster, Eric M. Smith, Y-Lan Boureau, Jason Weston
29. `LayoutLM <https://github.com/microsoft/unilm/tree/master/layoutlm>`_ (from Microsoft Research Asia) released with the paper
`LayoutLM: Pre-training of Text and Layout for Document Image Understanding
<https://arxiv.org/abs/1912.13318>`_ by Yiheng Xu, Minghao Li, Lei Cui, Shaohan Huang, Furu Wei, Ming Zhou.
29. `Other community models <https://huggingface.co/models>`_, contributed by the `community
30. `Other community models <https://huggingface.co/models>`_, contributed by the `community
<https://huggingface.co/users>`_.
.. toctree::
:maxdepth: 2
:caption: Get started
@@ -225,6 +226,7 @@ conversion utilities for the following models:
model_doc/mobilebert
model_doc/dpr
model_doc/pegasus
model_doc/blenderbot
model_doc/mbart
model_doc/fsmt
model_doc/funnel
+48
View File
@@ -0,0 +1,48 @@
Blenderbot
----------------------------------------------------
**DISCLAIMER:** If you see something strange,
file a `Github Issue <https://github.com/huggingface/transformers/issues/new?assignees=&labels=&template=bug-report.md&title>`
Overview
~~~~~~~~~~~~~~~~~~~~~
The Blender chatbot model was `proposed in Recipes for building an open-domain chatbot <https://arxiv.org/pdf/2004.13637.pdf>`_ Stephen Roller, Emily Dinan, Naman Goyal, Da Ju, Mary Williamson, Yinhan Liu, Jing Xu, Myle Ott, Kurt Shuster, Eric M. Smith, Y-Lan Boureau, Jason Weston on 30 Apr 2020.
Here the abstract,
Building open-domain chatbots is a challenging area for machine learning research. While prior work has shown that scaling neural models in the number of parameters and the size of the data they are trained on gives improved results, we show that other ingredients are important for a high-performing chatbot. Good conversation requires a number of skills that an expert conversationalist blends in a seamless way: providing engaging talking points and listening to their partners, and displaying knowledge, empathy and personality appropriately, while maintaining a consistent persona. We show that large scale models can learn these skills when given appropriate training data and choice of generation strategy. We build variants of these recipes with 90M, 2.7B and 9.4B parameter models, and make our models and code publicly available. Human evaluations show our best models are superior to existing approaches in multi-turn dialogue in terms of engagingness and humanness measurements. We then discuss the limitations of this work by analyzing failure cases of our models.
The The Authors' code can be found `here <https://github.com/facebookresearch/ParlAI>`_
Implementation Notes
~~~~~~~~~~~~~~~~~~~~
Blenderbot uses a standard seq2seq model transformer <https://arxiv.org/pdf/1706.03762.pdf> based architecture
BlenderbotConfig
~~~~~~~~~~~~~~~~~~~~~
.. autoclass:: transformers.BlenderbotConfig
:members:
BlenderbotTokenizer
~~~~~~~~~~~~~~~~~~~~~
.. autoclass:: transformers.BlenderbotTokenizer
:members: build_inputs_with_special_tokens
BlenderbotSmallTokenizer
~~~~~~~~~~~~~~~~~~~~~~~~~~
.. autoclass:: transformers.BlenderbotSmallTokenizer
:members: bpe, convert_tokens_to_string, save_vocabulary
BlenderbotForConditionalGeneration
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
.. autoclass:: transformers.BlenderbotForConditionalGeneration
:members: generate, forward
+1 -1
View File
@@ -96,7 +96,7 @@ As of Aug 10, 2020, they are:
pad_token_id=0,
eos_token_id=1,
is_encoder_decoder=True,
normalize_before=True,
variant='prelayernorm',
scale_embedding=True,
normalize_embedding=False,
add_final_layer_norm=True,
+1
View File
@@ -11,6 +11,7 @@ git-python==1.0.3
faiss-cpu
streamlit
elasticsearch
nltk
pandas
datasets
fire
+17
View File
@@ -0,0 +1,17 @@
import fire
from utils import calculate_rouge, save_json
def calculate_rouge_path(pred_path, tgt_path, save_path=None, **kwargs):
"""Kwargs will be passed to calculate_rouge"""
pred_lns = [x.strip() for x in open(pred_path).readlines()]
tgt_lns = [x.strip() for x in open(tgt_path).readlines()][: len(pred_lns)]
metrics = calculate_rouge(pred_lns, tgt_lns, **kwargs)
if save_path is not None:
save_json(metrics, save_path)
return metrics # these print nicely
if __name__ == "__main__":
fire.Fire(calculate_rouge_path)
+2 -1
View File
@@ -7,13 +7,14 @@ import sys
from collections import OrderedDict
from run_eval import datetime_now, run_generate
from utils import ROUGE_KEYS
# 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"],
"summarization": ["rouge1", "rouge2", "rougeL"],
"summarization": ROUGE_KEYS,
}
+21
View File
@@ -0,0 +1,21 @@
import re
try:
import nltk
NLTK_AVAILABLE = True
except (ImportError, ModuleNotFoundError):
NLTK_AVAILABLE = False
if NLTK_AVAILABLE:
try:
nltk.download("punkt", quiet=True)
except FileExistsError: # multiprocessing race condition
pass
def add_newline_to_end_of_each_sentence(x: str) -> str:
re.sub("<n>", "", x) # remove pegasus newline char
assert NLTK_AVAILABLE, "nltk must be installed to separate newlines betwee sentences. (pip install nltk)"
return "\n".join(nltk.sent_tokenize(x))
+80
View File
@@ -0,0 +1,80 @@
from collections import defaultdict
from pathlib import Path
import pandas as pd
from rouge_cli import calculate_rouge_path
from utils import calculate_rouge
PRED = [
'Prosecutor: "No videos were used in the crash investigation" German papers say they saw a cell phone video of the final seconds on board Flight 9525. The Germanwings co-pilot says he had a "previous episode of severe depression" German airline confirms it knew of Andreas Lubitz\'s depression years before he took control.',
"The Palestinian Authority officially becomes the 123rd member of the International Criminal Court. The formal accession was marked with a ceremony at The Hague, in the Netherlands. The Palestinians signed the ICC's founding Rome Statute in January. Israel and the United States opposed the Palestinians' efforts to join the body.",
"Amnesty International releases its annual report on the death penalty. The report catalogs the use of state-sanctioned killing as a punitive measure across the globe. At least 607 people were executed around the world in 2014, compared to 778 in 2013. The U.S. remains one of the worst offenders for imposing capital punishment.",
]
TGT = [
'Marseille prosecutor says "so far no videos were used in the crash investigation" despite media reports . Journalists at Bild and Paris Match are "very confident" the video clip is real, an editor says . Andreas Lubitz had informed his Lufthansa training school of an episode of severe depression, airline says .',
"Membership gives the ICC jurisdiction over alleged crimes committed in Palestinian territories since last June . Israel and the United States opposed the move, which could open the door to war crimes investigations against Israelis .",
"Amnesty's annual death penalty report catalogs encouraging signs, but setbacks in numbers of those sentenced to death . Organization claims that governments around the world are using the threat of terrorism to advance executions . The number of executions worldwide has gone down by almost 22% compared with 2013, but death sentences up by 28% .",
]
def test_disaggregated_scores_are_determinstic():
no_aggregation = calculate_rouge(PRED, TGT, bootstrap_aggregation=False, rouge_keys=["rouge2", "rougeL"])
assert isinstance(no_aggregation, defaultdict)
no_aggregation_just_r2 = calculate_rouge(PRED, TGT, bootstrap_aggregation=False, rouge_keys=["rouge2"])
assert (
pd.DataFrame(no_aggregation["rouge2"]).fmeasure.mean()
== pd.DataFrame(no_aggregation_just_r2["rouge2"]).fmeasure.mean()
)
def test_newline_cnn_improvement():
k = "rougeLsum"
score = calculate_rouge(PRED, TGT, newline_sep=True, rouge_keys=[k])[k]
score_no_sep = calculate_rouge(PRED, TGT, newline_sep=False, rouge_keys=[k])[k]
assert score > score_no_sep
def test_newline_irrelevant_for_other_metrics():
k = ["rouge1", "rouge2", "rougeL"]
score_sep = calculate_rouge(PRED, TGT, newline_sep=True, rouge_keys=k)
score_no_sep = calculate_rouge(PRED, TGT, newline_sep=False, rouge_keys=k)
assert score_sep == score_no_sep
def test_single_sent_scores_dont_depend_on_newline_sep():
pred = [
"Her older sister, Margot Frank, died in 1945, a month earlier than previously thought.",
'Marseille prosecutor says "so far no videos were used in the crash investigation" despite media reports .',
]
tgt = [
"Margot Frank, died in 1945, a month earlier than previously thought.",
'Prosecutor: "No videos were used in the crash investigation" German papers say they saw a cell phone video of the final seconds on board Flight 9525.',
]
assert calculate_rouge(pred, tgt, newline_sep=True) == calculate_rouge(pred, tgt, newline_sep=False)
def test_pegasus_newline():
pred = [
"""" "a person who has such a video needs to immediately give it to the investigators," prosecutor says .<n> "it is a very disturbing scene," editor-in-chief of bild online tells "erin burnett: outfront" """
]
tgt = [
""" Marseille prosecutor says "so far no videos were used in the crash investigation" despite media reports . Journalists at Bild and Paris Match are "very confident" the video clip is real, an editor says . Andreas Lubitz had informed his Lufthansa training school of an episode of severe depression, airline says ."""
]
prev_score = calculate_rouge(pred, tgt, rouge_keys=["rougeLsum"], newline_sep=False)["rougeLsum"]
new_score = calculate_rouge(pred, tgt, rouge_keys=["rougeLsum"])["rougeLsum"]
assert new_score > prev_score
def test_rouge_cli():
data_dir = Path("examples/seq2seq/test_data/wmt_en_ro")
metrics = calculate_rouge_path(data_dir.joinpath("test.source"), data_dir.joinpath("test.target"))
assert isinstance(metrics, dict)
metrics_default_dict = calculate_rouge_path(
data_dir.joinpath("test.source"), data_dir.joinpath("test.target"), bootstrap_aggregation=False
)
assert isinstance(metrics_default_dict, defaultdict)
+2 -2
View File
@@ -20,7 +20,7 @@ from run_eval_search import run_search
from transformers import AutoConfig, AutoModelForSeq2SeqLM
from transformers.hf_api import HfApi
from transformers.testing_utils import CaptureStderr, CaptureStdout, require_multigpu, require_torch_and_cuda, slow
from utils import label_smoothed_nll_loss, lmap, load_json
from utils import ROUGE_KEYS, label_smoothed_nll_loss, lmap, load_json
logging.basicConfig(level=logging.DEBUG)
@@ -365,7 +365,7 @@ def test_run_eval_search(model):
if "translation" in task:
expected_strings.append("bleu")
else:
expected_strings.extend(["rouge1", "rouge2", "rougeL"])
expected_strings.extend(ROUGE_KEYS)
for w in expected_strings:
assert w in cs.out
for w in un_expected_strings:
+53 -11
View File
@@ -18,6 +18,7 @@ from sacrebleu import corpus_bleu
from torch import nn
from torch.utils.data import Dataset, Sampler
from sentence_splitter import add_newline_to_end_of_each_sentence
from transformers import BartTokenizer
from transformers.file_utils import cached_property
@@ -378,19 +379,63 @@ def get_git_info():
return repo_infos
ROUGE_KEYS = ["rouge1", "rouge2", "rougeL"]
ROUGE_KEYS = ["rouge1", "rouge2", "rougeL", "rougeLsum"]
def calculate_rouge(output_lns: List[str], reference_lns: List[str], use_stemmer=True) -> Dict:
scorer = rouge_scorer.RougeScorer(ROUGE_KEYS, use_stemmer=use_stemmer)
def extract_rouge_mid_statistics(dct):
new_dict = {}
for k1, v1 in dct.items():
mid = v1.mid
new_dict[k1] = {stat: round(getattr(mid, stat), 4) for stat in ["precision", "recall", "fmeasure"]}
return new_dict
def calculate_rouge(
pred_lns: List[str],
tgt_lns: List[str],
use_stemmer=True,
rouge_keys=ROUGE_KEYS,
return_precision_and_recall=False,
bootstrap_aggregation=True,
newline_sep=True,
) -> Dict:
"""Calculate rouge using rouge_scorer package.
Args:
pred_lns: list of summaries generated by model
tgt_lns: list of groundtruth summaries (e.g. contents of val.target)
use_stemmer: Bool indicating whether Porter stemmer should be used to
strip word suffixes to improve matching.
rouge_keys: which metrics to compute, defaults to rouge1, rouge2, rougeL, rougeLsum
return_precision_and_recall: (False) whether to also return precision and recall.
bootstrap_aggregation: whether to do the typical bootstrap resampling of scores. Defaults to True, if False
this function returns a collections.defaultdict[metric: list of values for each observation for each subscore]``
newline_sep:(default=True) whether to add newline between sentences. This is essential for calculation rougeL
on multi sentence summaries (CNN/DM dataset).
Returns:
Dict[score: value] if aggregate else defaultdict(list) keyed by rouge_keys
"""
scorer = rouge_scorer.RougeScorer(rouge_keys, use_stemmer=use_stemmer)
aggregator = scoring.BootstrapAggregator()
for reference_ln, output_ln in zip(reference_lns, output_lns):
scores = scorer.score(reference_ln, output_ln)
for pred, tgt in zip(tgt_lns, pred_lns):
# rougeLsum expects "\n" separated sentences within a summary
if newline_sep:
pred = add_newline_to_end_of_each_sentence(pred)
tgt = add_newline_to_end_of_each_sentence(tgt)
scores = scorer.score(pred, tgt)
aggregator.add_scores(scores)
result = aggregator.aggregate()
return {k: round(v.mid.fmeasure * 100, 4) for k, v in result.items()}
if bootstrap_aggregation:
result = aggregator.aggregate()
if return_precision_and_recall:
return extract_rouge_mid_statistics(result) # here we return dict
else:
return {k: round(v.mid.fmeasure * 100, 4) for k, v in result.items()}
else:
return aggregator._scores # here we return defaultdict(list)
# Utilities for freezing parameters and checking whether they are frozen
@@ -423,9 +468,6 @@ def assert_not_all_frozen(model):
assert any(model_grads), f"none of {npars} weights require grad"
# 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.
-1
View File
@@ -1677,7 +1677,6 @@
" 'label2id': {'contradiction': 0, 'entailment': 2, 'neutral': 1},\n",
" 'max_position_embeddings': 1024,\n",
" 'model_type': 'bart',\n",
" 'normalize_before': False,\n",
" 'normalize_embedding': True,\n",
" 'num_hidden_layers': 12,\n",
" 'output_past': False,\n",
+3
View File
@@ -33,6 +33,7 @@ from .configuration_auto import ALL_PRETRAINED_CONFIG_ARCHIVE_MAP, CONFIG_MAPPIN
from .configuration_bart import BartConfig
from .configuration_bert import BERT_PRETRAINED_CONFIG_ARCHIVE_MAP, BertConfig
from .configuration_bert_generation import BertGenerationConfig
from .configuration_blenderbot import BLENDERBOT_PRETRAINED_CONFIG_ARCHIVE_MAP, BlenderbotConfig
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
@@ -154,6 +155,7 @@ from .tokenization_bert import BasicTokenizer, BertTokenizer, BertTokenizerFast,
from .tokenization_bert_generation import BertGenerationTokenizer
from .tokenization_bert_japanese import BertJapaneseTokenizer, CharacterTokenizer, MecabTokenizer
from .tokenization_bertweet import BertweetTokenizer
from .tokenization_blenderbot import BlenderbotSmallTokenizer, BlenderbotTokenizer
from .tokenization_camembert import CamembertTokenizer
from .tokenization_ctrl import CTRLTokenizer
from .tokenization_distilbert import DistilBertTokenizer, DistilBertTokenizerFast
@@ -299,6 +301,7 @@ if is_torch_available():
BertGenerationEncoder,
load_tf_weights_in_bert_generation,
)
from .modeling_blenderbot import BLENDERBOT_PRETRAINED_MODEL_ARCHIVE_LIST, BlenderbotForConditionalGeneration
from .modeling_camembert import (
CAMEMBERT_PRETRAINED_MODEL_ARCHIVE_LIST,
CamembertForCausalLM,
+4
View File
@@ -21,6 +21,7 @@ from .configuration_albert import ALBERT_PRETRAINED_CONFIG_ARCHIVE_MAP, AlbertCo
from .configuration_bart import BART_PRETRAINED_CONFIG_ARCHIVE_MAP, BartConfig
from .configuration_bert import BERT_PRETRAINED_CONFIG_ARCHIVE_MAP, BertConfig
from .configuration_bert_generation import BertGenerationConfig
from .configuration_blenderbot import BLENDERBOT_PRETRAINED_CONFIG_ARCHIVE_MAP, BlenderbotConfig
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
@@ -56,6 +57,7 @@ ALL_PRETRAINED_CONFIG_ARCHIVE_MAP = dict(
for pretrained_map in [
BERT_PRETRAINED_CONFIG_ARCHIVE_MAP,
BART_PRETRAINED_CONFIG_ARCHIVE_MAP,
BLENDERBOT_PRETRAINED_CONFIG_ARCHIVE_MAP,
MBART_PRETRAINED_CONFIG_ARCHIVE_MAP,
OPENAI_GPT_PRETRAINED_CONFIG_ARCHIVE_MAP,
TRANSFO_XL_PRETRAINED_CONFIG_ARCHIVE_MAP,
@@ -97,6 +99,7 @@ CONFIG_MAPPING = OrderedDict(
("marian", MarianConfig),
("mbart", MBartConfig),
("bart", BartConfig),
("blenderbot", BlenderbotConfig),
("reformer", ReformerConfig),
("longformer", LongformerConfig),
("roberta", RobertaConfig),
@@ -130,6 +133,7 @@ MODEL_NAMES_MAPPING = OrderedDict(
("camembert", "CamemBERT"),
("xlm-roberta", "XLM-RoBERTa"),
("pegasus", "Pegasus"),
("blenderbot", "Blenderbot"),
("marian", "Marian"),
("mbart", "mBART"),
("bart", "BART"),
+4 -2
View File
@@ -137,6 +137,7 @@ class BartConfig(PretrainedConfig):
normalize_embedding=True,
static_position_embeddings=False,
add_bias_logits=False,
do_blenderbot_90_layernorm=False,
force_bos_token_to_be_generated=False,
**common_kwargs
):
@@ -174,7 +175,7 @@ class BartConfig(PretrainedConfig):
self.max_position_embeddings = max_position_embeddings
self.init_std = init_std # Normal(0, this parameter)
self.activation_function = activation_function
self.do_blenderbot_90_layernorm = do_blenderbot_90_layernorm
# Params introduced for Mbart
self.scale_embedding = scale_embedding # scale factor will be sqrt(d_model) if True
self.normalize_embedding = normalize_embedding # True for mbart, False otherwise
@@ -194,7 +195,8 @@ class BartConfig(PretrainedConfig):
self.classif_dropout = classifier_dropout
# pos embedding offset
self.extra_pos_embeddings = self.pad_token_id + 1
self.extra_pos_embeddings = extra_pos_embeddings
# bart has a hack that offsets positional embeddings by 2, other models don't do do this
self.force_bos_token_to_be_generated = force_bos_token_to_be_generated
@@ -0,0 +1,68 @@
#!/usr/bin/env python3
# coding=utf-8
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the MIT license found in the;
# 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.
# LICENSE file in the root directory of this source tree.
from .configuration_bart import BartConfig
BLENDERBOT_PRETRAINED_CONFIG_ARCHIVE_MAP = {
"facebook/blenderbot-3B": "https://cdn.huggingface.co/facebook/blenderbot-3B/config.json",
"facebook/blenderbot-90M": "https://cdn.huggingface.co/facebook/blenderbot-/config.json",
}
class BlenderbotConfig(BartConfig):
"""
This is the configuration class to store the configuration of a :class:`~transformers.BlenderbotForConditionalGeneration`.
Instantiating a configuration with the defaults will yield a similar configuration to that of
the `blenderbot <https://huggingface.co/blenderbot>`__ architecture.
Configuration objects inherit from :class:`~transformers.BartConfig` and can be used
to control the model outputs. Read the documentation from :class:`~transformers.BartConfig`
for more information. The
Args:
d_model: (:obj:`int`, default to 2560), dimension of the embeddings vector
encoder_layers: (:obj:`int`, default to 2), number of layers in the encoder
encoder_ffn_size: (:obj:`int`, default to 10240), size of hidden layers in the FFN in the encoder
decoder_layers: (:obj:`int`, default to 24), number of layers in the decoder
decoder_ffn_size: (:obj:`int`, default to 10240), size of hidden layers in the FFN in the decoder
dropout: (:obj:`float`, default to 0.1), embedding dropout
activation_dropout: (:obj:`float`, default to 0.0), dropout after activation function
encoder_layerdrop: (:obj:`float`, default to 0.0,
decoder_layerdrop: (:obj:`float`, default to 0.0),
encoder_attention_heads:(:obj:`int`, default to 32), number of multi heads attention in the encoder
decoder_attention_heads:(:obj:`int`, default to 32), number of multi heads attention in the encoder
max_positions_embeddings:(:obj:`int`, default to 128), size of the position embeddings
activation: (:obj:`string`, default to 'gelu'), activation function to use
attention_dropout: (:obj:`float`, default to 0.0), multi head attention dropout
relu_dropout: (:obj:`float`, default to 0.0), relu dropout
vocab_size: (:obj:`int`, default to 8008), the size of the vocabulary
layernorm_variant: (obj: str, default to "prelayernorm") defines when to apply a layernorm
init_std: (obj: float, default to 0.02): The standard deviation of the truncated_normal_initializer for initializing all weight matrices.
is_encoder_decoder: (obj:`boolean`, default to True)
pad_token_id: (obj:`int`, default to 1): token id used to pad sequences.
bos_token_id: (obj:`int`, default to 0): begginning of sequence token id.
eos_token_id: (obj:`int`, default to 2): end of sequence token id.
add_final_layer_norm: (obj:`boolean`, default to False): if set to true a final Layernorm is added
scale_embedding: (obj:`boolean`, default to False): Scale embeddings by diving by sqrt(d_model)
normalize_embedding: (obj:`boolean`, default to False): apply Layernorm to the embedding layer output
static_position_embeddings: (:obj:`boolean`, default to False): if set to True positional embeddings are learnt otherwise use sinusoidal
Attributes:
pretrained_config_archive_map (Dict[str, str]): A dictionary containing all the available pre-trained checkpoints.
"""
model_type = "blenderbot"
+1 -1
View File
@@ -38,13 +38,13 @@ DEFAULTS = dict(
pad_token_id=0,
eos_token_id=1,
is_encoder_decoder=True,
normalize_before=True,
scale_embedding=True,
normalize_embedding=False,
add_final_layer_norm=True,
static_position_embeddings=True,
num_beams=8,
activation_function="relu",
layernorm_variant="prelayernorm",
)
# Config values that vary between checkpoints: for testing and conversion
task_specific_params = {
+1 -1
View File
@@ -54,7 +54,7 @@ RAG_CONFIG_DOC = r"""
A path to text passages compatible with the faiss index. Required if using
:class:`~transformers.retrieval_rag.LegacyIndex`
use_dummy_dataset (:obj:`bool`, `optional`, defaults to ``False``)
Whether to load a "dummy" variant of the dataset specified by :obj:`dataset`.
Whether to load a "dummy" layernorm_variant of the dataset specified by :obj:`dataset`.
label_smoothing (:obj:`float`, `optional`, defaults to 0.0):
Only relevant if ``return_loss`` is set to :obj:`True`. Controls the ``epsilon`` parameter value for label
smoothing in the loss calculation. If set to 0, no label smoothing is performed.
+114
View File
@@ -0,0 +1,114 @@
# coding=utf-8
# Copyright 2020 The HuggingFace Inc. team.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Convert Blenderbot checkpoint."""
import argparse
import logging
import torch
from transformers import BartConfig, BartForConditionalGeneration
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
PATTERNS = [
["attention", "attn"],
["encoder_attention", "encoder_attn"],
["q_lin", "q_proj"],
["k_lin", "k_proj"],
["v_lin", "v_proj"],
["out_lin", "out_proj"],
["norm_embeddings", "layernorm_embedding"],
["position_embeddings", "embed_positions"],
["embeddings", "embed_tokens"],
["ffn.lin", "fc"],
]
def rename_state_dict_key(k):
if k == "embeddings.weight":
return "shared.weight"
for parlai_name, hf_name in PATTERNS:
k = k.replace(parlai_name, hf_name)
if k.startswith("encoder"):
k = k.replace(".attn", ".self_attn")
k = k.replace("norm1", "self_attn_layer_norm")
k = k.replace("norm2", "final_layer_norm")
elif k.startswith("decoder"):
k = k.replace("norm1", "self_attn_layer_norm")
k = k.replace("norm2", "encoder_attn_layer_norm")
k = k.replace("norm3", "final_layer_norm")
return k
def rename_layernorm_keys(sd):
keys = [
"model.encoder.layernorm_embedding.weight",
"model.encoder.layernorm_embedding.bias",
"model.decoder.layernorm_embedding.weight",
"model.decoder.layernorm_embedding.bias",
]
for k in keys:
v = sd.pop(k)
new_k = k.replace("layernorm_embedding", "layer_norm")
assert new_k not in sd
sd[new_k] = v
IGNORE_KEYS = ["START"]
@torch.no_grad()
def convert_parlai_checkpoint(checkpoint_path, pytorch_dump_folder_path, config_json_path):
"""
Copy/paste/tweak model's weights to our BERT structure.
"""
model = torch.load(checkpoint_path, map_location="cpu")
sd = model["model"]
cfg = BartConfig.from_json_file(config_json_path)
m = BartForConditionalGeneration(cfg)
valid_keys = m.model.state_dict().keys()
failures = []
mapping = {}
for k, v in sd.items():
if k in IGNORE_KEYS:
continue
new_k = rename_state_dict_key(k)
if new_k not in valid_keys:
failures.append([k, new_k])
else:
mapping[new_k] = v
if cfg.layernorm_variant == "prelayernorm":
rename_layernorm_keys(sd)
m.model.load_state_dict(mapping, strict=True)
m.half()
m.save_pretrained(pytorch_dump_folder_path)
if __name__ == "__main__":
parser = argparse.ArgumentParser()
# Required parameters
parser.add_argument("--src_path", type=str, help="like blenderbot-model.bin")
parser.add_argument("--save_dir", default="hf_blenderbot", type=str, help="Where to save converted model.")
parser.add_argument(
"--hf_config_json", default="blenderbot-3b-config.json", type=str, help="Path to config to use"
)
args = parser.parse_args()
convert_parlai_checkpoint(args.src_path, args.save_dir, args.hf_config_json)
+3
View File
@@ -24,6 +24,7 @@ from .configuration_auto import (
BartConfig,
BertConfig,
BertGenerationConfig,
BlenderbotConfig,
CamembertConfig,
CTRLConfig,
DistilBertConfig,
@@ -80,6 +81,7 @@ from .modeling_bert import (
BertModel,
)
from .modeling_bert_generation import BertGenerationDecoder, BertGenerationEncoder
from .modeling_blenderbot import BlenderbotForConditionalGeneration
from .modeling_camembert import (
CamembertForCausalLM,
CamembertForMaskedLM,
@@ -337,6 +339,7 @@ MODEL_FOR_SEQ_TO_SEQ_CAUSAL_LM_MAPPING = OrderedDict(
(PegasusConfig, PegasusForConditionalGeneration),
(MarianConfig, MarianMTModel),
(MBartConfig, MBartForConditionalGeneration),
(BlenderbotConfig, BlenderbotForConditionalGeneration),
(BartConfig, BartForConditionalGeneration),
(FSMTConfig, FSMTForConditionalGeneration),
(EncoderDecoderConfig, EncoderDecoderModel),
+8 -2
View File
@@ -475,6 +475,7 @@ class BartDecoder(nn.Module):
super().__init__()
self.dropout = config.dropout
self.layerdrop = config.decoder_layerdrop
self.do_blenderbot_90_layernorm = config.do_blenderbot_90_layernorm # layernorm variant
self.padding_idx = embed_tokens.padding_idx
self.max_target_positions = config.max_position_embeddings
self.embed_scale = math.sqrt(config.d_model) if config.scale_embedding else 1.0
@@ -554,8 +555,13 @@ class BartDecoder(nn.Module):
positions = positions[:, -1:]
x = self.embed_tokens(input_ids) * self.embed_scale
x += positions
x = self.layernorm_embedding(x)
if self.do_blenderbot_90_layernorm:
x = self.layernorm_embedding(x)
x += positions
else:
x += positions
x = self.layernorm_embedding(x)
x = F.dropout(x, p=self.dropout, training=self.training)
# Convert to Bart output format: (seq_len, BS, model_dim) -> (BS, seq_len, model_dim)
+34
View File
@@ -0,0 +1,34 @@
#!/usr/bin/env python3
# coding=utf-8
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the MIT license found in the;
# 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.
# LICENSE file in the root directory of this source tree.
import torch
from .configuration_blenderbot import BlenderbotConfig
from .modeling_bart import BartForConditionalGeneration
BLENDERBOT_PRETRAINED_MODEL_ARCHIVE_LIST = ["facebook/blenderbot-3B", "facebook/blenderbot-90M"]
class BlenderbotForConditionalGeneration(BartForConditionalGeneration):
config_class = BlenderbotConfig
def adjust_logits_during_generation(self, logits, cur_len, max_length):
logits[:, self.config.bos_token_id] = -torch.finfo(torch.float16).max # near infinity fp16
if cur_len == max_length - 1 and self.config.eos_token_id is not None:
self._force_token_ids_generation(logits, self.config.eos_token_id)
return logits
+1 -1
View File
@@ -44,5 +44,5 @@ class MBartForConditionalGeneration(BartForConditionalGeneration):
>>> translation = tokenizer.batch_decode(translated_tokens, skip_special_tokens=True)[0]
>>> assert translation == "Şeful ONU declară că nu există o soluţie militară în Siria"
"""
model_type = "mbart"
config_class = MBartConfig
+4
View File
@@ -23,6 +23,7 @@ from .configuration_auto import (
BartConfig,
BertConfig,
BertGenerationConfig,
BlenderbotConfig,
CamembertConfig,
CTRLConfig,
DistilBertConfig,
@@ -59,6 +60,7 @@ from .tokenization_bert import BertTokenizer, BertTokenizerFast
from .tokenization_bert_generation import BertGenerationTokenizer
from .tokenization_bert_japanese import BertJapaneseTokenizer
from .tokenization_bertweet import BertweetTokenizer
from .tokenization_blenderbot import BlenderbotTokenizer
from .tokenization_camembert import CamembertTokenizer
from .tokenization_ctrl import CTRLTokenizer
from .tokenization_distilbert import DistilBertTokenizer, DistilBertTokenizerFast
@@ -104,6 +106,8 @@ TOKENIZER_MAPPING = OrderedDict(
(MBartConfig, (MBartTokenizer, None)),
(XLMRobertaConfig, (XLMRobertaTokenizer, None)),
(MarianConfig, (MarianTokenizer, None)),
(BlenderbotConfig, (BlenderbotTokenizer, None)),
(LongformerConfig, (LongformerTokenizer, None)),
(BartConfig, (BartTokenizer, BartTokenizerFast)),
(LongformerConfig, (LongformerTokenizer, LongformerTokenizerFast)),
(RobertaConfig, (BertweetTokenizer, None)),
+238
View File
@@ -0,0 +1,238 @@
import json
import logging
import os
from typing import List
import regex as re
from .tokenization_roberta import RobertaTokenizer
from .tokenization_utils import PreTrainedTokenizer
logger = logging.getLogger(__name__)
VOCAB_FILES_NAMES = {
"vocab_file": "vocab.json",
"merges_file": "merges.txt",
# "tokenizer_config_file": "tokenizer_config.json",
}
CKPT_3B = "facebook/blenderbot-3B"
logger = logging.getLogger(__name__)
class BlenderbotTokenizer(RobertaTokenizer):
vocab_files_names = {
"vocab_file": "vocab.json",
"merges_file": "merges.txt",
"tokenizer_config_file": "tokenizer_config.json",
}
pretrained_vocab_files_map = {
"vocab_file": {CKPT_3B: "https://cdn.huggingface.co/facebook/blenderbot-3B/vocab.json"},
"merges_file": {CKPT_3B: "https://cdn.huggingface.co/facebook/blenderbot-3B/merges.txt"},
"tokenizer_config_file": {CKPT_3B: "https://cdn.huggingface.co/facebook/blenderbot-3B/tokenizer_config.json"},
}
max_model_input_sizes = {"facebook/blenderbot-3B": 128}
def build_inputs_with_special_tokens(self, token_ids_0: List[int], token_ids_1: List[int] = None):
"""
Build model inputs from a sequence or a pair of sequence for sequence classification tasks
by concatenating and adding special tokens.
A RoBERTa sequence has the following format:
- single sequence: `` X </s>``
Args:
token_ids_0 (:obj:`List[int]`):
List of IDs to which the special tokens will be added
Returns:
:obj:`List[int]`: list of `input IDs <../glossary.html#input-ids>`__ with the appropriate special tokens.
"""
return token_ids_0 + [self.eos_token_id]
def get_pairs(word):
"""Return set of symbol pairs in a word.
Word is represented as tuple of symbols (symbols being variable-length strings).
"""
pairs = set()
prev_char = word[0]
for char in word[1:]:
pairs.add((prev_char, char))
prev_char = char
pairs = set(pairs)
return pairs
class BlenderbotSmallTokenizer(PreTrainedTokenizer):
"""
Constructs a Blenderbot-90M tokenizer. Peculiarities:
- Byte-Pair-Encoding
This tokenizer inherits from :class:`~transformers.PreTrainedTokenizer` which contains most of the methods. Users
should refer to the superclass for more information regarding methods.
Args:
vocab_file (:obj:`str`): Path to the vocabulary file.
merges_file (:obj:`str`): Path to the merges file.
bos_token (:obj:`string`, `optional`, defaults to "__start__"): The beginning of sentence token.
eos_token (:obj:`string`, `optional`, defaults to "__end__"): The end of sentence token.
unk_token (:obj:`string`, `optional`, defaults to "<unk>"): The unknown token. A token that is not in the
vocabulary cannot be converted to an ID and is set to be this token instead.
"""
vocab_files_names = {"vocab_file": "vocab.json", "merges_file": "merges.txt"}
pretrained_vocab_files_map = {
"vocab_file": {"facebook/blenderbot-90M": "https://cdn.huggingface.co/facebook/blenderbot-90M/vocab.json"},
"merges_file": {"facebook/blenderbot-90M": "https://cdn.huggingface.co/facebook/blenderbot-90M/merges.txt"},
}
max_model_input_sizes = {"facebook/blenderbot-90M": 512}
def __init__(
self,
vocab_file,
merges_file,
bos_token="__start__",
eos_token="__end__",
unk_token="__unk__",
pad_token="__null",
**kwargs
):
super().__init__(unk_token=unk_token, bos_token=bos_token, eos_token=eos_token, pad_token=pad_token, **kwargs)
with open(vocab_file, encoding="utf-8") as vocab_handle:
self.encoder = json.load(vocab_handle)
self.decoder = {v: k for k, v in self.encoder.items()}
with open(merges_file, encoding="utf-8") as merges_handle:
merges = merges_handle.read().split("\n")[1:-1]
merges = [tuple(merge.split()) for merge in merges]
self.bpe_ranks = dict(zip(merges, range(len(merges))))
self.cache = {}
@property
def vocab_size(self):
return len(self.encoder)
def get_vocab(self):
return dict(self.encoder, **self.added_tokens_encoder)
def bpe(self, token):
if token in self.cache:
return self.cache[token]
token = re.sub("([.,!?()])", r" \1", token)
token = re.sub("(')", r" \1 ", token)
token = re.sub("\s{2,}", " ", token)
if "\n" in token:
token = token.replace("\n", " __newln__")
tokens = token.split(" ")
words = []
for token in tokens:
token = token.lower()
word = tuple(token)
word = tuple(list(word[:-1]) + [word[-1] + "</w>"])
pairs = get_pairs(word)
if not pairs:
words.append(token)
continue
while True:
bigram = min(pairs, key=lambda pair: self.bpe_ranks.get(pair, float("inf")))
if bigram not in self.bpe_ranks:
break
first, second = bigram
new_word = []
i = 0
while i < len(word):
try:
j = word.index(first, i)
new_word.extend(word[i:j])
i = j
except ValueError:
new_word.extend(word[i:])
break
if word[i] == first and i < len(word) - 1 and word[i + 1] == second:
new_word.append(first + second)
i += 2
else:
new_word.append(word[i])
i += 1
new_word = tuple(new_word)
word = new_word
if len(word) == 1:
break
else:
pairs = get_pairs(word)
word = "@@ ".join(word)
word = word[:-4]
self.cache[token] = word
words.append(word)
return " ".join(words)
def _tokenize(self, text):
"""Tokenize a string."""
split_tokens = []
words = re.findall(r"\S+\n?", text)
for token in words:
split_tokens.extend([t for t in self.bpe(token).split(" ")])
return split_tokens
def _convert_token_to_id(self, token):
""" Converts a token (str) in an id using the vocab. """
token = token.lower()
return self.encoder.get(token, self.encoder.get(self.unk_token))
def _convert_id_to_token(self, index):
"""Converts an index (integer) in a token (str) using the vocab."""
return self.decoder.get(index, self.unk_token)
def convert_tokens_to_string(self, tokens):
""" Converts a sequence of tokens (string) in a single string. """
out_string = " ".join(tokens).replace("@@ ", "").strip()
return out_string
def save_vocabulary(self, save_directory):
"""
Save the vocabulary and special tokens file to a directory.
Args:
save_directory (:obj:`str`):
The directory in which to save the vocabulary.
Returns:
:obj:`Tuple(str)`: Paths to the files saved.
"""
if not os.path.isdir(save_directory):
logger.error("Vocabulary path ({}) should be a directory".format(save_directory))
return
vocab_file = os.path.join(save_directory, VOCAB_FILES_NAMES["vocab_file"])
merge_file = os.path.join(save_directory, VOCAB_FILES_NAMES["merges_file"])
with open(vocab_file, "w", encoding="utf-8") as f:
f.write(json.dumps(self.encoder, ensure_ascii=False))
index = 0
with open(merge_file, "w", encoding="utf-8") as writer:
writer.write("#version: 0.2\n")
for bpe_tokens, token_index in sorted(self.bpe_ranks.items(), key=lambda kv: kv[1]):
if index != token_index:
logger.warning(
"Saving vocabulary to {}: BPE merge indices are not consecutive."
" Please check that the tokenizer is not corrupted!".format(merge_file)
)
index = token_index
writer.write(" ".join(bpe_tokens) + "\n")
index += 1
return vocab_file, merge_file
+50
View File
@@ -0,0 +1,50 @@
import json
import os
import unittest
from transformers.testing_utils import slow
from transformers.tokenization_blenderbot import VOCAB_FILES_NAMES, BlenderbotTokenizer, BlenderbotSmallTokenizer
from .test_tokenization_common import TokenizerTesterMixin
class BlenderbotSmallTokenizerTest(TokenizerTesterMixin, unittest.TestCase):
tokenizer_class = BlenderbotSmallTokenizer
def setUp(self):
super().setUp()
# Adapted from Sennrich et al. 2015 and https://github.com/rsennrich/subword-nmt
vocab = ["adapt", "react", "read@@", "ap@@", "t", "__unk__", "__start__", "__end__", "__null__"]
vocab_tokens = dict(zip(vocab, range(len(vocab))))
merges = ["#version: 0.2", "a p", "ap t</w>", "r e", "a d", "ad apt</w>", ""]
self.special_tokens_map = {"bos_token": "__start", "eos_token": "__end__", "pad_token": "__null__", "unk_token": "__unk__"}
self.vocab_file = os.path.join(self.tmpdirname, VOCAB_FILES_NAMES["vocab_file"])
self.merges_file = os.path.join(self.tmpdirname, VOCAB_FILES_NAMES["merges_file"])
with open(self.vocab_file, "w", encoding="utf-8") as fp:
fp.write(json.dumps(vocab_tokens) + "\n")
with open(self.merges_file, "w", encoding="utf-8") as fp:
fp.write("\n".join(merges))
def get_tokenizer(self, **kwargs):
kwargs.update(self.special_tokens_map)
return BlenderbotSmallTokenizer.from_pretrained(self.tmpdirname, **kwargs)
def get_input_output_texts(self, tokenizer):
input_text = "adapt react readapt apt"
output_text = "adapt react readapt apt"
return input_text, output_text
def test_full_blenderbot_small_tokenizer(self):
tokenizer = BlenderbotSmallTokenizer(self.vocab_file, self.merges_file, **self.special_tokens_map)
text = "adapt react readapt apt"
bpe_tokens = ['adapt', 'react', 'read@@', 'ap@@', 't', 'ap@@', 't']
tokens = tokenizer.tokenize(text)
self.assertListEqual(tokens, bpe_tokens)
input_tokens = [tokenizer.bos_token] + tokens + [tokenizer.eos_token]
print(input_tokens)
# input_bpe_tokens = [0, 1, 2, 4, 5, 1, 0, 3, 6]
# self.assertListEqual(tokenizer.convert_tokens_to_ids(input_tokens), input_bpe_tokens)
+1
View File
@@ -152,6 +152,7 @@ class BenchmarkTest(unittest.TestCase):
def test_inference_encoder_decoder_with_configs(self):
MODEL_ID = "sshleifer/tinier_bart"
config = AutoConfig.from_pretrained(MODEL_ID)
config.use_cache = False
benchmark_args = PyTorchBenchmarkArguments(
models=[MODEL_ID],
training=False,
+6 -5
View File
@@ -212,8 +212,9 @@ class AutoModelTest(unittest.TestCase):
mapping = tuple(mapping.items())
for index, (child_config, child_model) in enumerate(mapping[1:]):
for parent_config, parent_model in mapping[: index + 1]:
with self.subTest(
msg="Testing if {} is child of {}".format(child_config.__name__, parent_config.__name__)
):
self.assertFalse(issubclass(child_config, parent_config))
self.assertFalse(issubclass(child_model, parent_model))
assert not issubclass(
child_config, parent_config
), "{child_config.__name__} is child of {parent_config.__name__}"
assert not issubclass(
child_model, parent_model
), "{child_config.__name__} is child of {parent_config.__name__}"
+14 -8
View File
@@ -168,7 +168,7 @@ class BARTModelTest(ModelTesterMixin, unittest.TestCase):
decoder_features_with_passed_mask = model(
decoder_attention_mask=invert_mask(decoder_attn_mask), decoder_input_ids=decoder_input_ids, **inputs_dict
)[0]
_assert_tensors_equal(decoder_features_with_passed_mask, decoder_features_with_created_mask)
assert_tensors_close(decoder_features_with_passed_mask, decoder_features_with_created_mask)
useless_mask = torch.zeros_like(decoder_attn_mask)
decoder_features = model(decoder_attention_mask=useless_mask, **inputs_dict)[0]
self.assertTrue(isinstance(decoder_features, torch.Tensor)) # no hidden states or attentions
@@ -182,7 +182,7 @@ class BARTModelTest(ModelTesterMixin, unittest.TestCase):
decoder_features_with_long_encoder_mask = model(
inputs_dict["input_ids"], attention_mask=inputs_dict["attention_mask"].long()
)[0]
_assert_tensors_equal(decoder_features_with_long_encoder_mask, decoder_features_with_created_mask)
assert_tensors_close(decoder_features_with_long_encoder_mask, decoder_features_with_created_mask)
def test_save_load_strict(self):
config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()
@@ -357,7 +357,7 @@ class BartHeadTests(unittest.TestCase):
]
for ex, desired_result in zip(examples, fairseq_results):
bart_toks = tokenizer.encode(ex, return_tensors="pt")
_assert_tensors_equal(desired_result.long(), bart_toks, prefix=ex)
assert_tensors_close(desired_result.long(), bart_toks, prefix=ex)
def test_generate_fp16(self):
config, input_ids, batch_size = self._get_config_and_data()
@@ -404,16 +404,22 @@ class BartHeadTests(unittest.TestCase):
self.assertTrue(torch.eq(input_new, output_new).all())
def _assert_tensors_equal(a, b, atol=1e-12, prefix=""):
"""If tensors not close, or a and b arent both tensors, raise a nice Assertion error."""
def assert_tensors_close(a, b, atol=1e-12, prefix=""):
"""If tensors not close, or a and b aren't both tensors, raise a nice Assertion error."""
if a is None and b is None:
return True
assert a.shape == b.shape
try:
if torch.allclose(a, b, atol=atol):
return True
raise
except Exception:
msg = "{} != {}".format(a, b)
pct_different = (torch.gt((a - b).abs(), atol)).float().mean().item()
if a.numel() > 100:
msg = f"tensor values are {pct_different:.1%} percent different."
else:
msg = f"{a} != {b}"
if prefix:
msg = prefix + ": " + msg
raise AssertionError(msg)
@@ -489,8 +495,8 @@ class BartModelIntegrationTests(unittest.TestCase):
inputs_dict = prepare_bart_inputs_dict(model.config, input_ids=input_ids_no_pad)
with torch.no_grad():
logits2 = model(**inputs_dict)[0]
_assert_tensors_equal(batched_logits[1], logits2, atol=TOLERANCE)
_assert_tensors_equal(expected_slice, logits_arr, atol=TOLERANCE)
assert_tensors_close(batched_logits[1], logits2, atol=TOLERANCE)
assert_tensors_close(expected_slice, logits_arr, atol=TOLERANCE)
@slow
def test_xsum_summarization_same_as_fairseq(self):
+215
View File
@@ -0,0 +1,215 @@
#!/usr/bin/env python3
# coding=utf-8
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the MIT license found in the;
# 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.
# LICENSE file in the root directory of this source tree.
import unittest
from transformers import is_torch_available
from transformers.file_utils import cached_property
from transformers.testing_utils import require_torch, slow, torch_device
from .test_configuration_common import ConfigTester
from .test_modeling_common import ModelTesterMixin, ids_tensor
if is_torch_available():
import torch
from transformers import BlenderbotConfig, BlenderbotForConditionalGeneration, BlenderbotTokenizer
from transformers.tokenization_blenderbot import BlenderbotSmallTokenizer
def _long_tensor(tok_lst):
return torch.tensor(tok_lst, dtype=torch.long, device=torch_device, requires_grad=False)
TOK_DECODE_KW = dict(skip_special_tokens=True, clean_up_tokenization_spaces=True)
FASTER_GEN_KWARGS = dict(num_beams=1, early_stopping=True, min_length=15, max_length=25)
@require_torch
class BlenderbotModelTester:
# Required attributes
vocab_size = 99
batch_size = 13
seq_length = 7
num_hidden_layers = 2
hidden_size = 16
num_attention_heads = 4
is_training = True
def __init__(self, parent):
torch.manual_seed(0)
self.parent = parent
self.config = BlenderbotConfig(
d_model=self.hidden_size,
dropout=0.0,
activation_function="gelu",
vocab_size=self.vocab_size,
encoder_layers=self.num_hidden_layers,
decoder_layers=self.num_hidden_layers,
encoder_attention_heads=self.num_attention_heads,
decoder_attention_heads=self.num_attention_heads,
attention_dropout=0.0,
encoder_ffn_dim=4,
decoder_ffn_dim=4,
do_blenderbot_90_layernorm=False,
normalize_before=True,
max_position_embeddings=50,
static_position_embeddings=False,
scale_embedding=True,
bos_token_id=0,
eos_token_id=2,
pad_token_id=1,
num_beams=1,
min_length=3,
max_length=10,
)
def prepare_config_and_inputs_for_common(self):
input_ids = ids_tensor([self.batch_size, self.seq_length], self.vocab_size)
attention_mask = ids_tensor([self.batch_size, self.seq_length], vocab_size=2)
inputs_dict = {"input_ids": input_ids, "attention_mask": attention_mask}
return self.config, inputs_dict
@require_torch
class BlenderbotTesterMixin(ModelTesterMixin, unittest.TestCase):
if is_torch_available():
all_generative_model_classes = (BlenderbotForConditionalGeneration,)
all_model_classes = (BlenderbotForConditionalGeneration,)
else:
all_generative_model_classes = ()
all_model_classes = ()
is_encoder_decoder = True
test_head_masking = False
test_pruning = False
test_missing_keys = False
test_torchscript = False
def setUp(self):
self.model_tester = BlenderbotModelTester(self)
self.config_tester = ConfigTester(self, config_class=BlenderbotConfig)
def test_inputs_embeds(self):
pass
def test_initialization_module(self):
config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()
model = BlenderbotForConditionalGeneration(config).model
model.to(torch_device)
model.eval()
enc_embeds = model.encoder.embed_tokens.weight
assert (enc_embeds == model.shared.weight).all().item()
self.assertAlmostEqual(torch.std(enc_embeds).item(), config.init_std, 2)
def test_embed_pos_shape(self):
config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()
model = BlenderbotForConditionalGeneration(config)
expected_shape = (config.max_position_embeddings + config.extra_pos_embeddings, config.d_model)
assert model.model.encoder.embed_positions.weight.shape == expected_shape
model.model.decoder.embed_positions.weight.shape == expected_shape
@unittest.skip("This test is flaky")
def test_feed_forward_chunking(self):
pass
@unittest.skip("This test is flaky")
def test_model_outputs_equivalence(self):
pass
@unittest.skipUnless(torch_device != "cpu", "3B test too slow on CPU.")
@require_torch
class Blenderbot3BIntegrationTests(unittest.TestCase):
ckpt = "facebook/blenderbot-3B"
@cached_property
def model(self):
model = BlenderbotForConditionalGeneration.from_pretrained(self.ckpt).to(torch_device)
if torch_device == "cuda":
model = model.half()
return model
@cached_property
def tokenizer(self):
return BlenderbotTokenizer.from_pretrained(self.ckpt)
@slow
def test_generation_from_short_input_same_as_parlai_3B(self):
src_text = ["Sam"]
model_inputs = self.tokenizer(src_text, return_tensors="pt").to(torch_device)
generated_utterances = self.model.generate(**model_inputs, **FASTER_GEN_KWARGS)
tgt_text = 'Sam is a great name. It means "sun" in Gaelic.'
generated_txt = self.tokenizer.batch_decode(generated_utterances, **TOK_DECODE_KW)
assert generated_txt[0].strip() == tgt_text
@slow
def test_generation_from_long_input_same_as_parlai_3B(self):
src_text = "Social anxiety\nWow, I am never shy. Do you have anxiety?\nYes. I end up sweating and blushing and feel like i'm going to throw up.\nand why is that?"
model_inputs = self.tokenizer([src_text], return_tensors="pt").to(torch_device)
generated_ids = self.model.generate(**model_inputs, **FASTER_GEN_KWARGS)[0]
reply = self.tokenizer.decode(generated_ids, **TOK_DECODE_KW)
assert "I think it's because we are so worried about what people think of us." == reply.strip()
@require_torch
class Blenderbot90MIntegrationTests(unittest.TestCase):
ckpt = "facebook/blenderbot-90M"
@cached_property
def model(self):
model = BlenderbotForConditionalGeneration.from_pretrained(self.ckpt).to(torch_device)
if torch_device == "cuda":
model = model.half()
return model
@cached_property
def tokenizer(self):
return BlenderbotSmallTokenizer.from_pretrained(self.ckpt)
@slow
def test_90_generation_from_long_input(self):
src_text = [
"Social anxiety\nWow, I am never shy. Do you have anxiety?\nYes. I end up sweating and blushing and feel like\
i'm going to throw up.\nand why is that?"
]
model_inputs = self.tokenizer(src_text, return_tensors="pt").to(torch_device)
generated_ids = self.model.generate(**model_inputs)[0]
reply = self.tokenizer.decode(generated_ids, **TOK_DECODE_KW)
assert reply in (
"i don't know. i just feel like i'm going to throw up. it's not fun.",
"i'm not sure. i just feel like i've been feeling like i have to be in a certain place",
)
def test_90_generation_from_short_input(self):
model_inputs = self.tokenizer(["sam"], return_tensors="pt").to(torch_device)
generated_utterances = self.model.generate(**model_inputs)
# generated_txt = self.tokenizer.decode(generated_utterances[0])
# assert generated_txt == "__start__ have you ever heard of sam harris? he's an american singer, songwriter, and actor. __end__"
clean_txt = self.tokenizer.decode(generated_utterances[0], **TOK_DECODE_KW)
assert clean_txt in (
"have you ever been to a sam club? it's a great club in the south.",
"have you ever heard of sam harris? he's an american singer, songwriter, and actor.",
)
+12 -2
View File
@@ -4,7 +4,7 @@ from transformers import is_torch_available
from transformers.file_utils import cached_property
from transformers.testing_utils import require_torch, slow, torch_device
from .test_modeling_bart import TOLERANCE, _assert_tensors_equal, _long_tensor
from .test_modeling_bart import TOLERANCE, _long_tensor, assert_tensors_close
if is_torch_available():
@@ -79,7 +79,17 @@ class MBartEnroIntegrationTest(AbstractSeq2SeqIntegrationTest):
expected_slice = torch.tensor([9.0078, 10.1113, 14.4787], device=logits.device, dtype=logits.dtype)
result_slice = logits[0, 0, :3]
_assert_tensors_equal(expected_slice, result_slice, atol=TOLERANCE)
assert_tensors_close(expected_slice, result_slice, atol=TOLERANCE)
@slow
def test_enro_generate_one(self):
batch: BatchEncoding = self.tokenizer.prepare_seq2seq_batch(
["UN Chief Says There Is No Military Solution in Syria"]
).to(torch_device)
translated_tokens = self.model.generate(**batch)
decoded = self.tokenizer.batch_decode(translated_tokens, skip_special_tokens=True)
self.assertEqual(self.tgt_text[0], decoded[0])
# self.assertEqual(self.tgt_text[1], decoded[1])
@slow
def test_enro_generate(self):
+92
View File
@@ -0,0 +1,92 @@
#!/usr/bin/env python3
# coding=utf-8
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the MIT license found in the;
# 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.
# LICENSE file in the root directory of this source tree.
import json
import os
import unittest
from transformers.file_utils import cached_property
from transformers.tokenization_blenderbot import VOCAB_FILES_NAMES, BlenderbotSmallTokenizer, BlenderbotTokenizer
from .test_tokenization_common import TokenizerTesterMixin
class BlenderbotSmallTokenizerTest(TokenizerTesterMixin, unittest.TestCase):
tokenizer_class = BlenderbotSmallTokenizer
def setUp(self):
super().setUp()
vocab = ["__start__", "adapt", "act", "ap@@", "te", "__end__", "__unk__"]
vocab_tokens = dict(zip(vocab, range(len(vocab))))
merges = ["#version: 0.2", "a p", "t e</w>", "ap t</w>", "a d", "ad apt</w>", "a c", "ac t</w>", ""]
self.special_tokens_map = {"unk_token": "__unk__", "bos_token": "__start__", "eos_token": "__end__"}
self.vocab_file = os.path.join(self.tmpdirname, VOCAB_FILES_NAMES["vocab_file"])
self.merges_file = os.path.join(self.tmpdirname, VOCAB_FILES_NAMES["merges_file"])
with open(self.vocab_file, "w", encoding="utf-8") as fp:
fp.write(json.dumps(vocab_tokens) + "\n")
with open(self.merges_file, "w", encoding="utf-8") as fp:
fp.write("\n".join(merges))
def get_tokenizer(self, **kwargs):
kwargs.update(self.special_tokens_map)
return BlenderbotSmallTokenizer.from_pretrained(self.tmpdirname, **kwargs)
def get_input_output_texts(self, tokenizer):
input_text = "adapt act apte"
output_text = "adapt act apte"
return input_text, output_text
def test_full_blenderbot_small_tokenizer(self):
tokenizer = BlenderbotSmallTokenizer(self.vocab_file, self.merges_file, **self.special_tokens_map)
text = "adapt act apte"
bpe_tokens = ["adapt", "act", "ap@@", "te"]
tokens = tokenizer.tokenize(text)
self.assertListEqual(tokens, bpe_tokens)
input_tokens = [tokenizer.bos_token] + tokens + [tokenizer.eos_token]
input_bpe_tokens = [0, 1, 2, 3, 4, 5]
self.assertListEqual(tokenizer.convert_tokens_to_ids(input_tokens), input_bpe_tokens)
def test_special_tokens_small_tok(self):
tok = BlenderbotSmallTokenizer.from_pretrained("facebook/blenderbot-90M")
assert tok("sam").input_ids == [1384]
src_text = "I am a small frog."
encoded = tok([src_text], padding=False, truncation=False)["input_ids"]
decoded = tok.batch_decode(encoded, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0]
assert src_text != decoded # I wish it did!
assert decoded == "i am a small frog ."
class Blenderbot3BTokenizerTests(unittest.TestCase):
@cached_property
def tokenizer_3b(self):
return BlenderbotTokenizer.from_pretrained("facebook/blenderbot-3B")
def test_encode_decode_cycle(self):
tok = self.tokenizer_3b
src_text = " I am a small frog."
encoded = tok([src_text], padding=False, truncation=False)["input_ids"]
decoded = tok.batch_decode(encoded, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0]
assert src_text == decoded
def test_3B_tokenization_same_as_parlai(self):
assert self.tokenizer_3b.add_prefix_space
assert self.tokenizer_3b([" Sam", "Sam"]).input_ids == [[5502, 2], [5502, 2]]
+1
View File
@@ -568,6 +568,7 @@ class TokenizerTesterMixin:
output = tokenizer(
[seq_2], [seq_1], padding=padding_state, truncation=truncation_state
)
self.assertEqual(len(output["input_ids"][0]), model_max_length)
# Simple