Compare commits
11
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
83f3d108f9 | ||
|
|
778cf7cff6 | ||
|
|
58f60f8a1f | ||
|
|
5fe3a8e1e6 | ||
|
|
4089d8f0d2 | ||
|
|
143333e49b | ||
|
|
26e1a6132c | ||
|
|
d2421ead71 | ||
|
|
dbba0d83f7 | ||
|
|
f6955ebde2 | ||
|
|
7c7b74d85c |
@@ -264,7 +264,7 @@ if is_torch_available():
|
||||
CamembertForQuestionAnswering,
|
||||
CAMEMBERT_PRETRAINED_MODEL_ARCHIVE_MAP,
|
||||
)
|
||||
from .modeling_encoder_decoder import PreTrainedEncoderDecoder
|
||||
from .modeling_encoder_decoder import EncoderDecoderModel
|
||||
from .modeling_t5 import (
|
||||
T5PreTrainedModel,
|
||||
T5Model,
|
||||
|
||||
@@ -957,7 +957,29 @@ class BertForMaskedLM(BertPreTrainedModel):
|
||||
ltr_lm_loss = loss_fct(prediction_scores.view(-1, self.config.vocab_size), lm_labels.view(-1))
|
||||
outputs = (ltr_lm_loss,) + outputs
|
||||
|
||||
return outputs # (masked_lm_loss), (ltr_lm_loss), prediction_scores, (hidden_states), (attentions)
|
||||
return outputs # (ltr_lm_loss), (masked_lm_loss), prediction_scores, (hidden_states), (attentions)
|
||||
|
||||
def prepare_inputs_for_generation(self, input_ids, attention_mask=None, **model_kwargs):
|
||||
input_shape = input_ids.shape
|
||||
effective_batch_size = input_shape[0]
|
||||
|
||||
# if model is used as a decoder in encoder-decoder model decoder attention mask is created on the fly
|
||||
if attention_mask is None:
|
||||
attention_mask = input_ids.new_ones(input_shape)
|
||||
|
||||
# if model is does not use a casaul mask then add a dummy token
|
||||
if self.config.is_decoder is False:
|
||||
assert self.config.pad_token_id is not None, "The PAD token should be defined for generation"
|
||||
attention_mask = torch.cat(
|
||||
[attention_mask, attention_mask.new_zeros((attention_mask.shape[0], 1))], dim=-1
|
||||
)
|
||||
|
||||
dummy_token = torch.full(
|
||||
(effective_batch_size, 1), self.config.pad_token_id, dtype=torch.long, device=input_ids.device
|
||||
)
|
||||
input_ids = torch.cat([input_ids, dummy_token], dim=1)
|
||||
|
||||
return {"input_ids": input_ids, "attention_mask": attention_mask}
|
||||
|
||||
|
||||
@add_start_docstrings(
|
||||
|
||||
@@ -18,35 +18,57 @@
|
||||
import logging
|
||||
import os
|
||||
|
||||
from torch import nn
|
||||
|
||||
from .modeling_auto import AutoModel, AutoModelWithLMHead
|
||||
from .modeling_utils import PreTrainedModel
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class PreTrainedEncoderDecoder(nn.Module):
|
||||
class EncoderDecoderModel(PreTrainedModel):
|
||||
r"""
|
||||
:class:`~transformers.PreTrainedEncoderDecoder` is a generic model class that will be
|
||||
:class:`~transformers.EncoderDecoder` is a generic model class that will be
|
||||
instantiated as a transformer architecture with one of the base model
|
||||
classes of the library as encoder and (optionally) another one as
|
||||
classes of the library as encoder and another one as
|
||||
decoder when created with the `AutoModel.from_pretrained(pretrained_model_name_or_path)`
|
||||
class method.
|
||||
class method for the encoder and `AutoModelWithLMHead.from_pretrained(pretrained_model_name_or_path)` class method for the decoder.
|
||||
"""
|
||||
|
||||
def __init__(self, encoder, decoder):
|
||||
super().__init__()
|
||||
assert encoder is not None, "The encoder has to be defined"
|
||||
assert decoder is not None, "The decoder has to be defined"
|
||||
|
||||
config = self._init_config(encoder.config, decoder.config)
|
||||
config.is_encoder_decoder = True
|
||||
super().__init__(config)
|
||||
|
||||
self.encoder = encoder
|
||||
assert (
|
||||
self.encoder.get_output_embeddings() is None
|
||||
), "The encoder {} should not have a LM Head. Please use a model without LM Head"
|
||||
self.decoder = decoder
|
||||
|
||||
def _init_config(self, encoder_config, decoder_config):
|
||||
# decoder config is used as default config (important for generation)
|
||||
config = decoder_config
|
||||
|
||||
return config
|
||||
|
||||
def get_encoder(self):
|
||||
return self.encoder
|
||||
|
||||
def get_decoder(self):
|
||||
return self.decoder
|
||||
|
||||
def get_input_embeddings(self):
|
||||
return self.encoder.get_input_embeddings()
|
||||
|
||||
def get_output_embeddings(self):
|
||||
return self.decoder.get_output_embeddings()
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(
|
||||
cls,
|
||||
encoder_pretrained_model_name_or_path=None,
|
||||
decoder_pretrained_model_name_or_path=None,
|
||||
*model_args,
|
||||
**kwargs
|
||||
cls, pretrained_model_name_or_path=None, decoder_pretrained_model_name_or_path=None, *model_args, **kwargs
|
||||
):
|
||||
r""" Instantiates an encoder and a decoder from one or two base classes of the library from pre-trained model checkpoints.
|
||||
|
||||
@@ -116,38 +138,31 @@ class PreTrainedEncoderDecoder(nn.Module):
|
||||
# `encoder_`), decoder-specific (prefixed by `decoder_`) and those
|
||||
# that apply to the model as a whole.
|
||||
# We let the specific kwargs override the common ones in case of conflict.
|
||||
kwargs_common = {
|
||||
argument: value
|
||||
for argument, value in kwargs.items()
|
||||
if not argument.startswith("encoder_") and not argument.startswith("decoder_")
|
||||
|
||||
kwargs_encoder = {
|
||||
argument[len("encoder_") :]: value for argument, value in kwargs.items() if argument.startswith("encoder_")
|
||||
}
|
||||
|
||||
kwargs_decoder = {
|
||||
argument[len("decoder_") :]: value for argument, value in kwargs.items() if argument.startswith("decoder_")
|
||||
}
|
||||
kwargs_decoder = kwargs_common.copy()
|
||||
kwargs_encoder = kwargs_common.copy()
|
||||
kwargs_encoder.update(
|
||||
{
|
||||
argument[len("encoder_") :]: value
|
||||
for argument, value in kwargs.items()
|
||||
if argument.startswith("encoder_")
|
||||
}
|
||||
)
|
||||
kwargs_decoder.update(
|
||||
{
|
||||
argument[len("decoder_") :]: value
|
||||
for argument, value in kwargs.items()
|
||||
if argument.startswith("decoder_")
|
||||
}
|
||||
)
|
||||
|
||||
# Load and initialize the encoder and decoder
|
||||
# The distinction between encoder and decoder at the model level is made
|
||||
# by the value of the flag `is_decoder` that we need to set correctly.
|
||||
encoder = kwargs_encoder.pop("model", None)
|
||||
if encoder is None:
|
||||
encoder = AutoModel.from_pretrained(encoder_pretrained_model_name_or_path, *model_args, **kwargs_encoder)
|
||||
assert (
|
||||
pretrained_model_name_or_path is not None
|
||||
), "If `model` is not defined as an argument, a `encoder_pretrained_model_name_or_path` has to be defined"
|
||||
encoder = AutoModel.from_pretrained(pretrained_model_name_or_path, *model_args, **kwargs_encoder)
|
||||
encoder.config.is_decoder = False
|
||||
|
||||
decoder = kwargs_decoder.pop("model", None)
|
||||
if decoder is None:
|
||||
assert (
|
||||
decoder_pretrained_model_name_or_path is not None
|
||||
), "If `decoder_model` is not defined as an argument, a `decoder_pretrained_model_name_or_path` has to be defined"
|
||||
decoder = AutoModelWithLMHead.from_pretrained(decoder_pretrained_model_name_or_path, **kwargs_decoder)
|
||||
decoder.config.is_decoder = True
|
||||
|
||||
@@ -201,18 +216,22 @@ class PreTrainedEncoderDecoder(nn.Module):
|
||||
os.mkdir(os.path.join(save_directory, "decoder"))
|
||||
self.decoder.save_pretrained(os.path.join(save_directory, "decoder"))
|
||||
|
||||
def forward(self, encoder_input_ids, decoder_input_ids, **kwargs):
|
||||
""" The forward pass on a seq2eq depends what we are performing:
|
||||
|
||||
- During training we perform one forward pass through both the encoder
|
||||
and decoder;
|
||||
- During prediction, we perform one forward pass through the encoder,
|
||||
and then perform several forward passes with the encoder's hidden
|
||||
state through the decoder to decode a full sequence.
|
||||
|
||||
Therefore, we skip the forward pass on the encoder if an argument named
|
||||
`encoder_hidden_state` is passed to this function.
|
||||
def forward(
|
||||
self,
|
||||
input_ids=None,
|
||||
inputs_embeds=None,
|
||||
attention_mask=None,
|
||||
head_mask=None,
|
||||
encoder_outputs=None,
|
||||
decoder_input_ids=None,
|
||||
decoder_attention_mask=None,
|
||||
decoder_head_mask=None,
|
||||
decoder_inputs_embeds=None,
|
||||
masked_lm_labels=None,
|
||||
lm_labels=None,
|
||||
):
|
||||
|
||||
"""
|
||||
Params:
|
||||
encoder_input_ids: ``torch.LongTensor`` of shape ``(batch_size, sequence_length)``
|
||||
Indices of encoder input sequence tokens in the vocabulary.
|
||||
@@ -220,17 +239,47 @@ class PreTrainedEncoderDecoder(nn.Module):
|
||||
Indices of decoder input sequence tokens in the vocabulary.
|
||||
kwargs: (`optional`) Remaining dictionary of keyword arguments.
|
||||
"""
|
||||
kwargs_encoder, kwargs_decoder = self.prepare_model_kwargs(**kwargs)
|
||||
|
||||
# Encode if needed (training, first prediction pass)
|
||||
encoder_hidden_states = kwargs_encoder.pop("hidden_states", None)
|
||||
if encoder_hidden_states is None:
|
||||
encoder_outputs = self.encoder(encoder_input_ids, **kwargs_encoder)
|
||||
encoder_hidden_states = encoder_outputs[0]
|
||||
else:
|
||||
encoder_outputs = ()
|
||||
if encoder_outputs is None:
|
||||
encoder_outputs = self.encoder(
|
||||
input_ids=input_ids, attention_mask=attention_mask, inputs_embeds=inputs_embeds, head_mask=head_mask
|
||||
)
|
||||
|
||||
kwargs_decoder["encoder_hidden_states"] = encoder_hidden_states
|
||||
decoder_outputs = self.decoder(decoder_input_ids, **kwargs_decoder)
|
||||
hidden_states = encoder_outputs[0]
|
||||
|
||||
# Decode
|
||||
decoder_outputs = self.decoder(
|
||||
input_ids=decoder_input_ids,
|
||||
inputs_embeds=decoder_inputs_embeds,
|
||||
attention_mask=decoder_attention_mask,
|
||||
encoder_hidden_states=hidden_states,
|
||||
encoder_attention_mask=attention_mask,
|
||||
head_mask=decoder_head_mask,
|
||||
lm_labels=lm_labels,
|
||||
masked_lm_labels=masked_lm_labels,
|
||||
)
|
||||
|
||||
return decoder_outputs + encoder_outputs
|
||||
|
||||
def prepare_inputs_for_generation(self, input_ids, past, attention_mask, **kwargs):
|
||||
assert past is not None, "past has to be defined for encoder_outputs"
|
||||
|
||||
# first step
|
||||
if type(past) is tuple:
|
||||
encoder_outputs = past
|
||||
else:
|
||||
encoder_outputs = (past,)
|
||||
|
||||
decoder_inputs = self.decoder.prepare_inputs_for_generation(input_ids)
|
||||
|
||||
return {
|
||||
"attention_mask": attention_mask,
|
||||
"decoder_attention_mask": decoder_inputs["attention_mask"],
|
||||
"decoder_input_ids": decoder_inputs["input_ids"],
|
||||
"encoder_outputs": encoder_outputs,
|
||||
}
|
||||
|
||||
def _reorder_cache(self, past, beam_idx):
|
||||
# as a default encoder-decoder models do not re-order the past.
|
||||
# TODO(PVP): might have to be updated, e.g. if GPT2 is to be used as a decoder
|
||||
return past
|
||||
|
||||
@@ -1,47 +0,0 @@
|
||||
# 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.
|
||||
""" Classes to support Encoder-Decoder architectures """
|
||||
|
||||
|
||||
def prepare_encoder_decoder_model_kwargs(**kwargs):
|
||||
""" Prepare the encoder and decoder's keyword arguments.
|
||||
|
||||
Keyword arguments come in 3 flavors:
|
||||
- encoder-specific (prefixed by `encoder_`)
|
||||
- decoder-specific (prefixed by `decoder_`)
|
||||
- those that apply to the model as whole.
|
||||
|
||||
We let the specific kwargs override the common ones in case of
|
||||
conflict.
|
||||
"""
|
||||
|
||||
kwargs_common = {
|
||||
argument: value
|
||||
for argument, value in kwargs.items()
|
||||
if not argument.startswith("encoder_") and not argument.startswith("decoder_")
|
||||
}
|
||||
if "input_ids" in kwargs_common:
|
||||
kwargs["encoder_input_ids"] = kwargs_common.pop("input_ids")
|
||||
|
||||
decoder_kwargs = kwargs_common.copy()
|
||||
encoder_kwargs = kwargs_common.copy()
|
||||
encoder_kwargs.update(
|
||||
{argument[len("encoder_") :]: value for argument, value in kwargs.items() if argument.startswith("encoder_")}
|
||||
)
|
||||
decoder_kwargs.update(
|
||||
{argument[len("decoder_") :]: value for argument, value in kwargs.items() if argument.startswith("decoder_")}
|
||||
)
|
||||
decoder_kwargs["encoder_attention_mask"] = encoder_kwargs.get("attention_mask", None)
|
||||
return encoder_kwargs, decoder_kwargs
|
||||
@@ -0,0 +1,248 @@
|
||||
# coding=utf-8
|
||||
# Copyright 2018 Google T5 Authors and HuggingFace Inc. team.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
|
||||
import os
|
||||
import tempfile
|
||||
import unittest
|
||||
|
||||
from transformers import is_torch_available
|
||||
|
||||
# this line reruns all the tests in BertModelTest; not sure whether this can be prevented
|
||||
# for now only run module with pytest tests/test_modeling_encoder_decoder.py::EncoderDecoderModelTest
|
||||
from .test_modeling_bert import BertModelTest
|
||||
from .utils import require_torch, slow
|
||||
|
||||
|
||||
if is_torch_available():
|
||||
from transformers import BertModel, BertForMaskedLM, EncoderDecoderModel
|
||||
|
||||
|
||||
@require_torch
|
||||
class EncoderDecoderModelTest(unittest.TestCase):
|
||||
def prepare_config_and_inputs_bert(self):
|
||||
bert_model_tester = BertModelTest.BertModelTester(self)
|
||||
encoder_config_and_inputs = bert_model_tester.prepare_config_and_inputs()
|
||||
decoder_config_and_inputs = bert_model_tester.prepare_config_and_inputs_for_decoder()
|
||||
(
|
||||
config,
|
||||
input_ids,
|
||||
token_type_ids,
|
||||
input_mask,
|
||||
sequence_labels,
|
||||
token_labels,
|
||||
choice_labels,
|
||||
) = encoder_config_and_inputs
|
||||
(
|
||||
decoder_config,
|
||||
decoder_input_ids,
|
||||
decoder_token_type_ids,
|
||||
decoder_input_mask,
|
||||
decoder_sequence_labels,
|
||||
decoder_token_labels,
|
||||
decoder_choice_labels,
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
) = decoder_config_and_inputs
|
||||
return {
|
||||
"config": config,
|
||||
"input_ids": input_ids,
|
||||
"attention_mask": input_mask,
|
||||
"decoder_config": decoder_config,
|
||||
"decoder_input_ids": decoder_input_ids,
|
||||
"decoder_token_type_ids": decoder_token_type_ids,
|
||||
"decoder_attention_mask": decoder_input_mask,
|
||||
"decoder_sequence_labels": decoder_sequence_labels,
|
||||
"decoder_token_labels": decoder_token_labels,
|
||||
"decoder_choice_labels": decoder_choice_labels,
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"lm_labels": decoder_token_labels,
|
||||
"masked_lm_labels": decoder_token_labels,
|
||||
}
|
||||
|
||||
def create_and_check_bert_encoder_decoder_model(
|
||||
self,
|
||||
config,
|
||||
input_ids,
|
||||
attention_mask,
|
||||
encoder_hidden_states,
|
||||
decoder_config,
|
||||
decoder_input_ids,
|
||||
decoder_attention_mask,
|
||||
**kwargs
|
||||
):
|
||||
encoder_model = BertModel(config)
|
||||
decoder_model = BertForMaskedLM(decoder_config)
|
||||
enc_dec_model = EncoderDecoderModel(encoder_model, decoder_model)
|
||||
outputs_encoder_decoder = enc_dec_model(
|
||||
input_ids=input_ids,
|
||||
decoder_input_ids=decoder_input_ids,
|
||||
attention_mask=attention_mask,
|
||||
decoder_attention_mask=decoder_attention_mask,
|
||||
)
|
||||
|
||||
self.assertEqual(outputs_encoder_decoder[0].shape, (decoder_input_ids.shape + (decoder_config.vocab_size,)))
|
||||
self.assertEqual(outputs_encoder_decoder[1].shape, (input_ids.shape + (config.hidden_size,)))
|
||||
encoder_outputs = (encoder_hidden_states,)
|
||||
outputs_encoder_decoder = enc_dec_model(
|
||||
encoder_outputs=encoder_outputs,
|
||||
decoder_input_ids=decoder_input_ids,
|
||||
attention_mask=attention_mask,
|
||||
decoder_attention_mask=decoder_attention_mask,
|
||||
)
|
||||
|
||||
self.assertEqual(outputs_encoder_decoder[0].shape, (decoder_input_ids.shape + (decoder_config.vocab_size,)))
|
||||
self.assertEqual(outputs_encoder_decoder[1].shape, (input_ids.shape + (config.hidden_size,)))
|
||||
|
||||
def create_and_check_bert_encoder_decoder_model_from_pretrained(
|
||||
self,
|
||||
config,
|
||||
input_ids,
|
||||
attention_mask,
|
||||
encoder_hidden_states,
|
||||
decoder_config,
|
||||
decoder_input_ids,
|
||||
decoder_attention_mask,
|
||||
**kwargs
|
||||
):
|
||||
encoder_model = BertModel(config)
|
||||
decoder_model = BertForMaskedLM(decoder_config)
|
||||
kwargs = {"encoder_model": encoder_model, "decoder_model": decoder_model}
|
||||
enc_dec_model = EncoderDecoderModel.from_pretrained(**kwargs)
|
||||
outputs_encoder_decoder = enc_dec_model(
|
||||
input_ids=input_ids,
|
||||
decoder_input_ids=decoder_input_ids,
|
||||
attention_mask=attention_mask,
|
||||
decoder_attention_mask=decoder_attention_mask,
|
||||
)
|
||||
|
||||
self.assertEqual(outputs_encoder_decoder[0].shape, (decoder_input_ids.shape + (decoder_config.vocab_size,)))
|
||||
self.assertEqual(outputs_encoder_decoder[1].shape, (input_ids.shape + (config.hidden_size,)))
|
||||
|
||||
def create_and_check_save_and_load_encoder_decoder_model(self, config, decoder_config, **kwargs):
|
||||
encoder_model = BertModel(config)
|
||||
decoder_model = BertForMaskedLM(decoder_config)
|
||||
kwargs = {"encoder_model": encoder_model, "decoder_model": decoder_model}
|
||||
enc_dec_model = EncoderDecoderModel.from_pretrained(**kwargs)
|
||||
|
||||
with tempfile.TemporaryDirectory() as temp_dir_name:
|
||||
enc_dec_model.save_pretrained(temp_dir_name)
|
||||
enc_dec_model.from_pretrained(
|
||||
pretrained_model_name_or_path=os.path.join(temp_dir_name, "encoder"),
|
||||
decoder_pretrained_model_name_or_path=os.path.join(temp_dir_name, "decoder"),
|
||||
)
|
||||
|
||||
def check_loss_output(self, loss):
|
||||
self.assertEqual(loss.size(), ())
|
||||
|
||||
def create_and_check_bert_encoder_decoder_model_mlm_labels(
|
||||
self,
|
||||
config,
|
||||
input_ids,
|
||||
attention_mask,
|
||||
encoder_hidden_states,
|
||||
decoder_config,
|
||||
decoder_input_ids,
|
||||
decoder_attention_mask,
|
||||
masked_lm_labels,
|
||||
**kwargs
|
||||
):
|
||||
encoder_model = BertModel(config)
|
||||
decoder_model = BertForMaskedLM(decoder_config)
|
||||
enc_dec_model = EncoderDecoderModel(encoder_model, decoder_model)
|
||||
outputs_encoder_decoder = enc_dec_model(
|
||||
input_ids=input_ids,
|
||||
decoder_input_ids=decoder_input_ids,
|
||||
attention_mask=attention_mask,
|
||||
decoder_attention_mask=decoder_attention_mask,
|
||||
masked_lm_labels=masked_lm_labels,
|
||||
)
|
||||
|
||||
mlm_loss = outputs_encoder_decoder[0]
|
||||
self.check_loss_output(mlm_loss)
|
||||
# check that backprop works
|
||||
mlm_loss.backward()
|
||||
|
||||
self.assertEqual(outputs_encoder_decoder[1].shape, (decoder_input_ids.shape + (decoder_config.vocab_size,)))
|
||||
self.assertEqual(outputs_encoder_decoder[2].shape, (input_ids.shape + (config.hidden_size,)))
|
||||
|
||||
def create_and_check_bert_encoder_decoder_model_lm_labels(
|
||||
self,
|
||||
config,
|
||||
input_ids,
|
||||
attention_mask,
|
||||
encoder_hidden_states,
|
||||
decoder_config,
|
||||
decoder_input_ids,
|
||||
decoder_attention_mask,
|
||||
lm_labels,
|
||||
**kwargs
|
||||
):
|
||||
encoder_model = BertModel(config)
|
||||
decoder_model = BertForMaskedLM(decoder_config)
|
||||
enc_dec_model = EncoderDecoderModel(encoder_model, decoder_model)
|
||||
outputs_encoder_decoder = enc_dec_model(
|
||||
input_ids=input_ids,
|
||||
decoder_input_ids=decoder_input_ids,
|
||||
attention_mask=attention_mask,
|
||||
decoder_attention_mask=decoder_attention_mask,
|
||||
lm_labels=lm_labels,
|
||||
)
|
||||
|
||||
lm_loss = outputs_encoder_decoder[0]
|
||||
self.check_loss_output(lm_loss)
|
||||
# check that backprop works
|
||||
lm_loss.backward()
|
||||
|
||||
self.assertEqual(outputs_encoder_decoder[1].shape, (decoder_input_ids.shape + (decoder_config.vocab_size,)))
|
||||
self.assertEqual(outputs_encoder_decoder[2].shape, (input_ids.shape + (config.hidden_size,)))
|
||||
|
||||
def create_and_check_bert_encoder_decoder_model_generate(self, input_ids, config, decoder_config, **kwargs):
|
||||
encoder_model = BertModel(config)
|
||||
decoder_model = BertForMaskedLM(decoder_config)
|
||||
enc_dec_model = EncoderDecoderModel(encoder_model, decoder_model)
|
||||
|
||||
# Bert does not have a bos token id, so use pad_token_id instead
|
||||
generated_output = enc_dec_model.generate(input_ids, decoder_start_token_id=enc_dec_model.config.pad_token_id)
|
||||
self.assertEqual(generated_output.shape, (input_ids.shape[0],) + (decoder_config.max_length,))
|
||||
|
||||
def test_bert_encoder_decoder_model(self):
|
||||
input_ids_dict = self.prepare_config_and_inputs_bert()
|
||||
self.create_and_check_bert_encoder_decoder_model(**input_ids_dict)
|
||||
|
||||
def test_bert_encoder_decoder_model_from_pretrained(self):
|
||||
input_ids_dict = self.prepare_config_and_inputs_bert()
|
||||
self.create_and_check_bert_encoder_decoder_model_from_pretrained(**input_ids_dict)
|
||||
|
||||
def test_save_and_load_from_prertained(self):
|
||||
input_ids_dict = self.prepare_config_and_inputs_bert()
|
||||
self.create_and_check_save_and_load_encoder_decoder_model(**input_ids_dict)
|
||||
|
||||
def test_bert_encoder_decoder_model_mlm_labels(self):
|
||||
input_ids_dict = self.prepare_config_and_inputs_bert()
|
||||
self.create_and_check_bert_encoder_decoder_model_mlm_labels(**input_ids_dict)
|
||||
|
||||
def test_bert_encoder_decoder_model_lm_labels(self):
|
||||
input_ids_dict = self.prepare_config_and_inputs_bert()
|
||||
self.create_and_check_bert_encoder_decoder_model_lm_labels(**input_ids_dict)
|
||||
|
||||
def test_bert_encoder_decoder_model_generate(self):
|
||||
input_ids_dict = self.prepare_config_and_inputs_bert()
|
||||
self.create_and_check_bert_encoder_decoder_model_generate(**input_ids_dict)
|
||||
|
||||
@slow
|
||||
def test_real_bert_model_from_pretrained(self):
|
||||
model = EncoderDecoderModel.from_pretrained("bert-base-uncased", "bert-base-uncased")
|
||||
self.assertIsNotNone(model)
|
||||
Reference in New Issue
Block a user