Compare commits
168
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
34eb3fcd55 | ||
|
|
1ab6644b8b | ||
|
|
8a1b6c2439 | ||
|
|
18e736f5e3 | ||
|
|
0352ce44b2 | ||
|
|
ba149a6ef1 | ||
|
|
c0bbf1178e | ||
|
|
e136853ca4 | ||
|
|
50d93bb051 | ||
|
|
ec73340b2f | ||
|
|
72f17d24c5 | ||
|
|
5668d35e36 | ||
|
|
fc1af2377d | ||
|
|
e9ddee9208 | ||
|
|
7ff0aea1d9 | ||
|
|
9ff0d016e8 | ||
|
|
3a1a2ee9bf | ||
|
|
e129f240e5 | ||
|
|
3c0b8e3b36 | ||
|
|
1a3b40a1e0 | ||
|
|
57fed3ff3a | ||
|
|
11bc12f98a | ||
|
|
5b2ace0c98 | ||
|
|
ca3c39e3a4 | ||
|
|
4d65a5dbe0 | ||
|
|
77cd835a06 | ||
|
|
e5efa0bb6a | ||
|
|
371f1ba53e | ||
|
|
4f635ae946 | ||
|
|
f0483a09c6 | ||
|
|
642e51b9c5 | ||
|
|
69b4921c12 | ||
|
|
2715b25989 | ||
|
|
588bdf747a | ||
|
|
d9d340998a | ||
|
|
eb4cd09143 | ||
|
|
961af20683 | ||
|
|
751e95cc88 | ||
|
|
6ff219825d | ||
|
|
19496e6c1b | ||
|
|
8e61b730f4 | ||
|
|
c03bc9932e | ||
|
|
7b42e87835 | ||
|
|
12659ba87a | ||
|
|
90f1210148 | ||
|
|
c986d7fd28 | ||
|
|
9cfd4dc0ea | ||
|
|
99bc11119a | ||
|
|
f7c8b6018e | ||
|
|
ab98343c08 | ||
|
|
6eab35737d | ||
|
|
bcd9cf6813 | ||
|
|
b18f28d2e5 | ||
|
|
537a8febca | ||
|
|
6c7355bd5f | ||
|
|
9379a7720c | ||
|
|
3f7883b350 | ||
|
|
ce8fa251b7 | ||
|
|
907331ffe2 | ||
|
|
851a371c2c | ||
|
|
3ad5f3401b | ||
|
|
ea57d77150 | ||
|
|
b33cad7edd | ||
|
|
964ce608d5 | ||
|
|
3447c2aeb4 | ||
|
|
4dd0300e8c | ||
|
|
2702c37634 | ||
|
|
1f6071f76b | ||
|
|
00a73d3809 | ||
|
|
9f4bfa55ea | ||
|
|
3bcff4fcdd | ||
|
|
5dbc2e9d09 | ||
|
|
f5169f3907 | ||
|
|
6ac9174299 | ||
|
|
d286fc9c4c | ||
|
|
d6bfa8c823 | ||
|
|
d53c539176 | ||
|
|
3791efa561 | ||
|
|
687ae67cc7 | ||
|
|
06d391e94d | ||
|
|
376ff07c85 | ||
|
|
77d2d3dfe3 | ||
|
|
b0a57d443a | ||
|
|
5c6ec437be | ||
|
|
213b49c3e9 | ||
|
|
0b2e2bf599 | ||
|
|
9399b1ff95 | ||
|
|
f7d347df0e | ||
|
|
ed8e6baf7c | ||
|
|
3ac98d0384 | ||
|
|
302b03596d | ||
|
|
9e91f69583 | ||
|
|
3d1394429a | ||
|
|
c332e6a7d8 | ||
|
|
5c78acd5ed | ||
|
|
2f6fdb375b | ||
|
|
c57c4b824f | ||
|
|
4a818f3ff0 | ||
|
|
833f706880 | ||
|
|
83493999be | ||
|
|
bdbbe24462 | ||
|
|
a85f5a4a27 | ||
|
|
a25426f680 | ||
|
|
e24acdef49 | ||
|
|
068a9eea01 | ||
|
|
26e3bdba2b | ||
|
|
875b5b0359 | ||
|
|
14aadf410c | ||
|
|
c8463845be | ||
|
|
30e26004a0 | ||
|
|
d074f58194 | ||
|
|
c857cf163d | ||
|
|
a08cdfd577 | ||
|
|
c3a40f2e1c | ||
|
|
d802181c95 | ||
|
|
9e9c50ae3a | ||
|
|
9bf8df4bbc | ||
|
|
da04344f13 | ||
|
|
a82b85475d | ||
|
|
da0c4f8bb4 | ||
|
|
df5bf64bb2 | ||
|
|
fb1f0b4209 | ||
|
|
e50807900f | ||
|
|
07e1a3b0e6 | ||
|
|
6a6527e702 | ||
|
|
79e6c39d29 | ||
|
|
08d3f716bd | ||
|
|
3e933766c9 | ||
|
|
724dc87981 | ||
|
|
84dc65ea09 | ||
|
|
9d7908939e | ||
|
|
08e630e1be | ||
|
|
bb988e901d | ||
|
|
b77daa402d | ||
|
|
f8cb16c541 | ||
|
|
6367bf8d99 | ||
|
|
1fbf338417 | ||
|
|
7a5a9d9018 | ||
|
|
9ac911271b | ||
|
|
63a799c3d0 | ||
|
|
766a65b193 | ||
|
|
1a2db58dd5 | ||
|
|
a26f8d28a6 | ||
|
|
05a262cff3 | ||
|
|
9738920b6f | ||
|
|
36c8039e93 | ||
|
|
f6de270f68 | ||
|
|
7f31ebfd34 | ||
|
|
958ca055fb | ||
|
|
6d93aa4221 | ||
|
|
6d2f767b12 | ||
|
|
3eb866c1e2 | ||
|
|
0a0cbe48f1 | ||
|
|
673646553e | ||
|
|
83ffc33ae5 | ||
|
|
1053e64318 | ||
|
|
335e794c86 | ||
|
|
1b1b9e7968 | ||
|
|
0472762c8a | ||
|
|
b55db43c42 | ||
|
|
9894cd388b | ||
|
|
1e1556a57d | ||
|
|
237652b61d | ||
|
|
6c7e29c01c | ||
|
|
9f6c49fbc2 | ||
|
|
775e03b7b6 | ||
|
|
c7a7a5a01f | ||
|
|
c71834634f |
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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,
|
||||
|
||||
@@ -11,6 +11,7 @@ git-python==1.0.3
|
||||
faiss-cpu
|
||||
streamlit
|
||||
elasticsearch
|
||||
nltk
|
||||
pandas
|
||||
datasets
|
||||
fire
|
||||
|
||||
@@ -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)
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -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))
|
||||
@@ -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)
|
||||
@@ -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
@@ -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.
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"),
|
||||
|
||||
@@ -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"
|
||||
@@ -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 = {
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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)
|
||||
@@ -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),
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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)),
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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,
|
||||
|
||||
@@ -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__}"
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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.",
|
||||
)
|
||||
@@ -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):
|
||||
|
||||
@@ -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]]
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user