Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3bfeebe722 | ||
|
|
ffed6a8a5d | ||
|
|
426b5106e7 |
@@ -624,7 +624,6 @@ def main():
|
||||
and os.listdir(args.output_dir)
|
||||
and args.do_train
|
||||
and not args.overwrite_output_dir
|
||||
and not args.should_continue
|
||||
):
|
||||
raise ValueError(
|
||||
"Output directory ({}) already exists and is not empty. Use --overwrite_output_dir to overcome.".format(
|
||||
|
||||
@@ -29,14 +29,7 @@ class TestT5Examples(unittest.TestCase):
|
||||
output_file_name = Path(tempfile.gettempdir()) / "utest_output_t5_sum.hypo"
|
||||
score_file_name = Path(tempfile.gettempdir()) / "utest_score_t5_sum.hypo"
|
||||
|
||||
testargs = [
|
||||
"evaluate_cnn.py",
|
||||
"patrickvonplaten/t5-tiny-random",
|
||||
str(tmp),
|
||||
str(output_file_name),
|
||||
str(tmp),
|
||||
str(score_file_name),
|
||||
]
|
||||
testargs = ["evaluate_cnn.py", "t5-small", str(tmp), str(output_file_name), str(tmp), str(score_file_name)]
|
||||
|
||||
with patch.object(sys, "argv", testargs):
|
||||
run_generate()
|
||||
|
||||
@@ -37,7 +37,7 @@ class TestT5Examples(unittest.TestCase):
|
||||
|
||||
testargs = [
|
||||
"evaluate_wmt.py",
|
||||
"patrickvonplaten/t5-tiny-random",
|
||||
"t5-small",
|
||||
str(tmp_source),
|
||||
str(output_file_name),
|
||||
str(tmp_target),
|
||||
|
||||
@@ -116,6 +116,7 @@ class PretrainedBartModel(PreTrainedModel):
|
||||
config_class = BartConfig
|
||||
base_model_prefix = "model"
|
||||
pretrained_model_archive_map = BART_PRETRAINED_MODEL_ARCHIVE_MAP
|
||||
encoder_outputs_batch_dim_idx = 1 # outputs shaped (seq_len, bs, ...)
|
||||
|
||||
def _init_weights(self, module):
|
||||
std = self.config.init_std
|
||||
@@ -293,10 +294,7 @@ class BartEncoder(nn.Module):
|
||||
if self.output_hidden_states:
|
||||
encoder_states.append(x)
|
||||
|
||||
# T x B x C -> B x T x C
|
||||
encoder_states = [hidden_state.transpose(0, 1) for hidden_state in encoder_states]
|
||||
x = x.transpose(0, 1)
|
||||
|
||||
return x, encoder_states, all_attentions
|
||||
|
||||
|
||||
@@ -440,7 +438,7 @@ class BartDecoder(nn.Module):
|
||||
# embed positions
|
||||
positions = self.embed_positions(input_ids, generation_mode=generation_mode)
|
||||
|
||||
if generation_mode and decoder_cached_states is not None:
|
||||
if generation_mode:
|
||||
input_ids = input_ids[:, -1:]
|
||||
positions = positions[:, -1:] # happens after we embed them
|
||||
assert input_ids.ne(self.padding_idx).any()
|
||||
@@ -450,11 +448,7 @@ class BartDecoder(nn.Module):
|
||||
|
||||
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)
|
||||
x = x.transpose(0, 1)
|
||||
encoder_hidden_states = encoder_hidden_states.transpose(0, 1)
|
||||
|
||||
x = x.transpose(0, 1) # (seq_len, BS, model_dim)
|
||||
# decoder layers
|
||||
all_hidden_states = ()
|
||||
all_self_attns = ()
|
||||
@@ -476,19 +470,18 @@ class BartDecoder(nn.Module):
|
||||
causal_mask=decoder_causal_mask,
|
||||
)
|
||||
|
||||
if generation_mode:
|
||||
if self.output_past:
|
||||
next_decoder_cache.append(layer_past.copy())
|
||||
if self.output_hidden_states:
|
||||
all_hidden_states += (x,)
|
||||
if self.output_attentions:
|
||||
all_self_attns += (layer_self_attn,)
|
||||
|
||||
# Convert to standart output format: (seq_len, BS, model_dim) -> (BS, seq_len, model_dim)
|
||||
# Convert shapes from (seq_len, BS, model_dim) to (BS, seq_len, model_dim)
|
||||
all_hidden_states = [hidden_state.transpose(0, 1) for hidden_state in all_hidden_states]
|
||||
x = x.transpose(0, 1)
|
||||
encoder_hidden_states = encoder_hidden_states.transpose(0, 1)
|
||||
|
||||
if generation_mode:
|
||||
if self.output_past:
|
||||
next_cache = ((encoder_hidden_states, encoder_padding_mask), next_decoder_cache)
|
||||
else:
|
||||
next_cache = None
|
||||
@@ -907,18 +900,23 @@ class BartForConditionalGeneration(PretrainedBartModel):
|
||||
def prepare_inputs_for_generation(self, decoder_input_ids, past, attention_mask, **kwargs):
|
||||
assert past is not None, "past has to be defined for encoder_outputs"
|
||||
|
||||
# first step, decoder_cached_states are empty (None)
|
||||
(encoder_outputs, encoder_attention_mask), decoder_cached_states = past
|
||||
# first step, decoder_cached_states are empty
|
||||
if not past[1]:
|
||||
encoder_outputs, decoder_cached_states = past, None
|
||||
else:
|
||||
encoder_outputs, decoder_cached_states = past
|
||||
return {
|
||||
"input_ids": None, # encoder_outputs is defined. input_ids not needed
|
||||
"encoder_outputs": encoder_outputs,
|
||||
"attention_mask": encoder_attention_mask,
|
||||
"decoder_cached_states": decoder_cached_states,
|
||||
"decoder_input_ids": decoder_input_ids,
|
||||
"attention_mask": attention_mask,
|
||||
"generation_mode": True,
|
||||
}
|
||||
|
||||
def prepare_scores_for_generation(self, scores, cur_len, max_length):
|
||||
if cur_len == 1:
|
||||
self._force_token_ids_generation(scores, self.config.bos_token_id)
|
||||
if cur_len == max_length - 1 and self.config.eos_token_id is not None:
|
||||
self._force_token_ids_generation(scores, self.config.eos_token_id)
|
||||
return scores
|
||||
@@ -926,19 +924,16 @@ class BartForConditionalGeneration(PretrainedBartModel):
|
||||
@staticmethod
|
||||
def _reorder_cache(past, beam_idx):
|
||||
((enc_out, enc_mask), decoder_cached_states) = past
|
||||
|
||||
if decoder_cached_states is not None:
|
||||
reordered_past = []
|
||||
for layer_past in decoder_cached_states:
|
||||
# get the correct batch idx from decoder layer's batch dim for cross and self-attn
|
||||
layer_past_new = {
|
||||
attn_key: _reorder_buffer(attn_cache, beam_idx) for attn_key, attn_cache in layer_past.items()
|
||||
}
|
||||
reordered_past.append(layer_past_new)
|
||||
else:
|
||||
reordered_past = None
|
||||
|
||||
new_enc_out = enc_out if enc_out is None else (enc_out[0].index_select(1, beam_idx), *enc_out[1:])
|
||||
reordered_past = []
|
||||
for layer_past in decoder_cached_states:
|
||||
# get the correct batch idx from decoder layer's batch dim for cross and self-attn
|
||||
layer_past_new = {
|
||||
attn_key: _reorder_buffer(attn_cache, beam_idx) for attn_key, attn_cache in layer_past.items()
|
||||
}
|
||||
# reordered_layer_past = [layer_past[:, i].unsqueeze(1).clone().detach() for i in beam_idx]
|
||||
# reordered_layer_past = torch.cat(reordered_layer_past, dim=1)
|
||||
reordered_past.append(layer_past_new)
|
||||
new_enc_out = enc_out if enc_out is None else enc_out.index_select(1, beam_idx)
|
||||
new_enc_mask = enc_mask if enc_mask is None else enc_mask.index_select(0, beam_idx)
|
||||
|
||||
past = ((new_enc_out, new_enc_mask), reordered_past)
|
||||
|
||||
@@ -20,6 +20,7 @@ import itertools
|
||||
import logging
|
||||
import math
|
||||
import os
|
||||
from typing import List, Optional
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
@@ -404,6 +405,8 @@ class T5Block(nn.Module):
|
||||
def __init__(self, config, has_relative_attention_bias=False):
|
||||
super().__init__()
|
||||
self.is_decoder = config.is_decoder
|
||||
self.device = None
|
||||
self.next_device = None
|
||||
self.layer = nn.ModuleList()
|
||||
self.layer.append(T5LayerSelfAttention(config, has_relative_attention_bias=has_relative_attention_bias))
|
||||
if self.is_decoder:
|
||||
@@ -422,6 +425,21 @@ class T5Block(nn.Module):
|
||||
encoder_decoder_position_bias=None,
|
||||
head_mask=None,
|
||||
):
|
||||
if self.device is not None:
|
||||
(hidden_states,
|
||||
attention_mask,
|
||||
position_bias,
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
encoder_decoder_position_bias,
|
||||
head_mask) = tuple(t.to(self.device) for t in (hidden_states,
|
||||
attention_mask,
|
||||
position_bias,
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
encoder_decoder_position_bias,
|
||||
head_mask)) # Model parallelism
|
||||
|
||||
self_attention_outputs = self.layer[0](
|
||||
hidden_states, attention_mask=attention_mask, position_bias=position_bias, head_mask=head_mask
|
||||
)
|
||||
@@ -445,6 +463,10 @@ class T5Block(nn.Module):
|
||||
hidden_states = self.layer[2](hidden_states)
|
||||
|
||||
outputs = (hidden_states,) + outputs # add attentions if we output them
|
||||
|
||||
if self.next_device is not None:
|
||||
outputs = tuple(t.to(self.device) for t in outputs) # Model parallelism
|
||||
|
||||
return outputs # hidden-states, (self-attention weights), (self-attention position bias), (cross-attention weights), (cross-attention position bias)
|
||||
|
||||
|
||||
@@ -457,6 +479,7 @@ class T5PreTrainedModel(PreTrainedModel):
|
||||
pretrained_model_archive_map = T5_PRETRAINED_MODEL_ARCHIVE_MAP
|
||||
load_tf_weights = load_tf_weights_in_t5
|
||||
base_model_prefix = "transformer"
|
||||
encoder_outputs_batch_dim_idx = 0 # outputs shaped (bs, ...)
|
||||
|
||||
@property
|
||||
def dummy_inputs(self):
|
||||
@@ -540,6 +563,9 @@ class T5Stack(T5PreTrainedModel):
|
||||
|
||||
self.init_weights()
|
||||
|
||||
def get_block_list(self):
|
||||
return list(self.block)
|
||||
|
||||
def get_input_embeddings(self):
|
||||
return self.embed_tokens
|
||||
|
||||
@@ -772,6 +798,48 @@ class T5Model(T5PreTrainedModel):
|
||||
|
||||
self.init_weights()
|
||||
|
||||
def spread_on_devices(self, devices: Optional[List] = None):
|
||||
""" Spread a transformers model on several devices by moving block on several devices (simple model parallelism)
|
||||
|
||||
The blocks of the transformers are spread among the given device list
|
||||
or on all visible CUDA devices if no device list is given.
|
||||
|
||||
The first device will host in addition the embeddings and the input/output tensors.
|
||||
|
||||
"""
|
||||
if devices is None and torch.cuda.is_available():
|
||||
devices = list(range(torch.cuda.device_count()))
|
||||
if len(devices) < 2:
|
||||
self.to(devices[0] if devices else None)
|
||||
return
|
||||
|
||||
modules_to_move = set(self.modules)
|
||||
|
||||
# Evenly spread the blocks on devices
|
||||
block_list = self.get_block_list()
|
||||
group_size = len(block_list) // len(devices)
|
||||
for i, block in enumerate(block_list):
|
||||
device = devices[i // group_size]
|
||||
# Note that we cannot easily use `forward_pre_hook` to move tensors around since this type of hooks currently
|
||||
# only act on the positional arguments send to the forward pass (PyTorch 1.4.0).
|
||||
# So you should call your model's forward pass with tensors as positional arguments
|
||||
# see: https://github.com/pytorch/pytorch/blob/master/torch/nn/modules/module.py#L548-L554
|
||||
# block.register_forward_pre_hook(lambda module, input: tuple(t.to(device) for t in input))
|
||||
block.to(device)
|
||||
block.device = device
|
||||
modules_to_move.remove(block)
|
||||
|
||||
# Take care of brining back the tensors to the first device at the end of the last block's forward
|
||||
block.next_device = device[0]
|
||||
# block.register_forward_hook(lambda module, input, output: tuple(t.to(device[0]) for t in output))
|
||||
|
||||
# Move the remaining modules (embeddings) on the first device
|
||||
for module in list(modules_to_move):
|
||||
module.to(devices[0])
|
||||
|
||||
def get_block_list(self):
|
||||
return list(self.encoder.get_block_list()) + list(self.decoder.get_block_list())
|
||||
|
||||
def get_input_embeddings(self):
|
||||
return self.shared
|
||||
|
||||
|
||||
+119
-144
@@ -658,26 +658,24 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
def generate(
|
||||
self,
|
||||
input_ids=None,
|
||||
attention_mask=None,
|
||||
encoder_input_ids=None,
|
||||
decoder_input_ids=None,
|
||||
min_length=None,
|
||||
max_length=None,
|
||||
early_stopping=None,
|
||||
num_return_sequences=None,
|
||||
num_beams=None,
|
||||
min_length=None,
|
||||
do_sample=None,
|
||||
early_stopping=None,
|
||||
num_beams=None,
|
||||
temperature=None,
|
||||
top_k=None,
|
||||
top_p=None,
|
||||
repetition_penalty=None,
|
||||
bad_words_ids=None,
|
||||
bos_token_id=None,
|
||||
pad_token_id=None,
|
||||
eos_token_id=None,
|
||||
length_penalty=None,
|
||||
no_repeat_ngram_size=None,
|
||||
repetition_penalty=None,
|
||||
use_cache=None,
|
||||
num_return_sequences=None,
|
||||
attention_mask=None,
|
||||
decoder_start_token_id=None,
|
||||
):
|
||||
r""" Generates sequences for models with a LM head. The method currently supports greedy decoding, beam-search decoding, sampling with temperature, sampling with top-k or nucleus sampling.
|
||||
|
||||
@@ -690,42 +688,24 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
Parameters:
|
||||
|
||||
input_ids: (`optional`) `torch.LongTensor` of shape `(batch_size, sequence_length)`
|
||||
Short-hand for either encoder_input_ids in seq2seq models or decoder_input_ids for language models.
|
||||
|
||||
attention_mask (`optional`) obj: `torch.LongTensor` of same shape as `input_ids`
|
||||
Mask to avoid performing attention on padding token indices.
|
||||
Mask values selected in ``[0, 1]``:
|
||||
``1`` for tokens that are NOT MASKED, ``0`` for MASKED tokens.
|
||||
Defaults to `None`.
|
||||
|
||||
`What are attention masks? <../glossary.html#attention-mask>`__
|
||||
|
||||
encoder_input_ids: (`optional`) `torch.LongTensor` of shape `(batch_size, sequence_length)`
|
||||
encoder inputs in the encoder-decoder setting (uses input_ids if None)
|
||||
|
||||
decoder_input_ids: (`optional`) `torch.LongTensor` of shape `(batch_size, sequence_length)`
|
||||
The sequence used as a prompt for the generation in the encoder-decoder setting.
|
||||
If `None` in the seq2seq setting, the method initializes it as a `torch.LongTensor` of shape `(batch_size, 1)` filled with bos_token_id.
|
||||
If `None` in language modeling setting, uses input_ids if available, or initializes with bos_token_id otherwise.
|
||||
|
||||
min_length: (`optional`) int
|
||||
The min length of the sequence to be generated. Between 0 and infinity. Default to 0.
|
||||
The sequence used as a prompt for the generation. If `None` the method initializes
|
||||
it as an empty `torch.LongTensor` of shape `(1,)`.
|
||||
|
||||
max_length: (`optional`) int
|
||||
The max length of the sequence to be generated. Between `min_length` and infinity. Default to 20.
|
||||
|
||||
early_stopping: (`optional`) bool
|
||||
if set to `True` beam search is stopped when at least `num_beams` sentences finished per batch. Defaults to `False` as defined in `configuration_utils.PretrainedConfig`.
|
||||
|
||||
num_return_sequences: (`optional`) int
|
||||
The number of independently computed returned sequences for each element in the batch. Default to 1.
|
||||
|
||||
num_beams: (`optional`) int
|
||||
Number of beams for beam search. Must be between 1 and infinity. 1 means no beam search. Default to 1.
|
||||
min_length: (`optional`) int
|
||||
The min length of the sequence to be generated. Between 0 and infinity. Default to 0.
|
||||
|
||||
do_sample: (`optional`) bool
|
||||
If set to `False` greedy decoding is used. Otherwise sampling is used. Defaults to `False` as defined in `configuration_utils.PretrainedConfig`.
|
||||
|
||||
early_stopping: (`optional`) bool
|
||||
if set to `True` beam search is stopped when at least `num_beams` sentences finished per batch. Defaults to `False` as defined in `configuration_utils.PretrainedConfig`.
|
||||
|
||||
num_beams: (`optional`) int
|
||||
Number of beams for beam search. Must be between 1 and infinity. 1 means no beam search. Default to 1.
|
||||
|
||||
temperature: (`optional`) float
|
||||
The value used to module the next token probabilities. Must be strictly positive. Default to 1.0.
|
||||
|
||||
@@ -735,6 +715,9 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
top_p: (`optional`) float
|
||||
The cumulative probability of parameter highest probability vocabulary tokens to keep for nucleus sampling. Must be between 0 and 1. Default to 1.
|
||||
|
||||
repetition_penalty: (`optional`) float
|
||||
The parameter for repetition penalty. Between 1.0 and infinity. 1.0 means no penalty. Default to 1.0.
|
||||
|
||||
pad_token_id: (`optional`) int
|
||||
Padding token. Default to specicic model pad_token_id or None if it does not exist.
|
||||
|
||||
@@ -749,16 +732,23 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
|
||||
no_repeat_ngram_size: (`optional`) int
|
||||
If set to int > 0, all ngrams of size `no_repeat_ngram_size` can only occur once.
|
||||
|
||||
repetition_penalty: (`optional`) float
|
||||
The parameter for repetition penalty. Between 1.0 and infinity. 1.0 means no penalty. Default to 1.0.
|
||||
|
||||
bad_words_ids: (`optional`) list of lists of int
|
||||
`bad_words_ids` contains tokens that are not allowed to be generated. In order to get the tokens of the words that should not appear in the generated text, use `tokenizer.encode(bad_word, add_prefix_space=True)`.
|
||||
|
||||
use_cache: (`optional`) bool
|
||||
If set to `True` the model re-uses pre-computed decoder hidden states from one time step to the next.
|
||||
num_return_sequences: (`optional`) int
|
||||
The number of independently computed returned sequences for each element in the batch. Default to 1.
|
||||
|
||||
attention_mask (`optional`) obj: `torch.LongTensor` of same shape as `input_ids`
|
||||
Mask to avoid performing attention on padding token indices.
|
||||
Mask values selected in ``[0, 1]``:
|
||||
``1`` for tokens that are NOT MASKED, ``0`` for MASKED tokens.
|
||||
Defaults to `None`.
|
||||
|
||||
`What are attention masks? <../glossary.html#attention-mask>`__
|
||||
|
||||
decoder_start_token_id=None: (`optional`) int
|
||||
If an encoder-decoder model starts decoding with a different token than BOS.
|
||||
Defaults to `None` and is changed to `BOS` later.
|
||||
|
||||
Return:
|
||||
|
||||
@@ -810,18 +800,16 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
"Please use another model class (e.g. `OpenAIGPTLMHeadModel`, `XLNetLMHeadModel`, `GPT2LMHeadModel`, `CTRLLMHeadModel`, `T5WithLMHeadModel`, `TransfoXLLMHeadModel`, `XLMWithLMHeadModel`, `BartForConditionalGeneration` )"
|
||||
)
|
||||
|
||||
min_length = min_length if min_length is not None else self.config.min_length
|
||||
max_length = max_length if max_length is not None else self.config.max_length
|
||||
early_stopping = early_stopping if early_stopping is not None else self.config.early_stopping
|
||||
num_return_sequences = (
|
||||
num_return_sequences if num_return_sequences is not None else self.config.num_return_sequences
|
||||
)
|
||||
num_beams = num_beams if num_beams is not None else self.config.num_beams
|
||||
min_length = min_length if min_length is not None else self.config.min_length
|
||||
do_sample = do_sample if do_sample is not None else self.config.do_sample
|
||||
early_stopping = early_stopping if early_stopping is not None else self.config.early_stopping
|
||||
num_beams = num_beams if num_beams is not None else self.config.num_beams
|
||||
temperature = temperature if temperature is not None else self.config.temperature
|
||||
top_k = top_k if top_k is not None else self.config.top_k
|
||||
top_p = top_p if top_p is not None else self.config.top_p
|
||||
repetition_penalty = repetition_penalty if repetition_penalty is not None else self.config.repetition_penalty
|
||||
bos_token_id = bos_token_id if bos_token_id is not None else self.config.bos_token_id
|
||||
pad_token_id = pad_token_id if pad_token_id is not None else self.config.pad_token_id
|
||||
eos_token_id = eos_token_id if eos_token_id is not None else self.config.eos_token_id
|
||||
length_penalty = length_penalty if length_penalty is not None else self.config.length_penalty
|
||||
@@ -829,12 +817,15 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
no_repeat_ngram_size if no_repeat_ngram_size is not None else self.config.no_repeat_ngram_size
|
||||
)
|
||||
bad_words_ids = bad_words_ids if bad_words_ids is not None else self.config.bad_words_ids
|
||||
|
||||
num_return_sequences = (
|
||||
num_return_sequences if num_return_sequences is not None else self.config.num_return_sequences
|
||||
)
|
||||
decoder_start_token_id = (
|
||||
decoder_start_token_id if decoder_start_token_id is not None else self.config.decoder_start_token_id
|
||||
)
|
||||
|
||||
if input_ids is not None:
|
||||
batch_size = input_ids.shape[0] # overriden by the input batch_size
|
||||
elif encoder_input_ids is not None:
|
||||
batch_size = encoder_input_ids.shape[0]
|
||||
else:
|
||||
batch_size = 1
|
||||
|
||||
@@ -847,6 +838,9 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
assert isinstance(top_k, int) and top_k >= 0, "`top_k` should be a positive integer."
|
||||
assert 0 <= top_p <= 1, "`top_p` should be between 0 and 1."
|
||||
assert repetition_penalty >= 1.0, "`repetition_penalty` should be >= 1."
|
||||
assert input_ids is not None or (
|
||||
isinstance(bos_token_id, int) and bos_token_id >= 0
|
||||
), "If input_ids is not defined, `bos_token_id` should be a positive integer."
|
||||
assert pad_token_id is None or (
|
||||
isinstance(pad_token_id, int) and (pad_token_id >= 0)
|
||||
), "`pad_token_id` should be a positive integer."
|
||||
@@ -864,32 +858,16 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
bad_words_ids is None or isinstance(bad_words_ids, list) and isinstance(bad_words_ids[0], list)
|
||||
), "`bad_words_ids` is either `None` or a list of lists of tokens that should not be generated"
|
||||
|
||||
# different requirements fod decoder-only and encoder-decoder
|
||||
if self.config.is_encoder_decoder:
|
||||
assert hasattr(self, "get_encoder"), "{} should have a 'get_encoder' function defined".format(self)
|
||||
assert callable(self.get_encoder), "{} should be a method".format(self.get_encoder)
|
||||
assert input_ids is not None or encoder_input_ids is not None, "the encoder inputs should to be provided"
|
||||
bos_token_id = bos_token_id if bos_token_id is not None else self.config.decoder_start_token_id
|
||||
encoder_ids = encoder_input_ids if encoder_input_ids is not None else input_ids # for back-compatibility, esp summarization examples
|
||||
decoder_ids = decoder_input_ids
|
||||
else:
|
||||
decoder_ids = decoder_input_ids if decoder_input_ids is not None else input_ids
|
||||
|
||||
bos_token_id = bos_token_id if bos_token_id is not None else self.config.bos_token_id
|
||||
|
||||
assert decoder_ids is not None or (
|
||||
isinstance(bos_token_id, int) and bos_token_id >= 0
|
||||
), "If input_ids is not defined, `bos_token_id` should be a positive integer."
|
||||
if decoder_ids is None:
|
||||
if input_ids is None:
|
||||
assert isinstance(bos_token_id, int) and bos_token_id >= 0, (
|
||||
"you should either supply a context to complete as `input_ids` input "
|
||||
"or a `bos_token_id` (integer >= 0) as a first token to start the generation."
|
||||
)
|
||||
decoder_ids = torch.full(
|
||||
input_ids = torch.full(
|
||||
(batch_size, 1), bos_token_id, dtype=torch.long, device=next(self.parameters()).device,
|
||||
)
|
||||
else:
|
||||
assert decoder_ids.dim() == 2, "Input prompt should be of shape (batch_size, sequence length)."
|
||||
assert input_ids.dim() == 2, "Input prompt should be of shape (batch_size, sequence length)."
|
||||
|
||||
# not allow to duplicate outputs when greedy decoding
|
||||
if do_sample is False:
|
||||
@@ -898,6 +876,7 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
assert (
|
||||
num_return_sequences == 1
|
||||
), "Greedy decoding will always produce the same output for num_beams == 1 and num_return_sequences > 1. Please set num_return_sequences = 1"
|
||||
|
||||
else:
|
||||
# beam_search greedy generation conditions
|
||||
assert (
|
||||
@@ -906,23 +885,10 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
|
||||
# create attention mask if necessary
|
||||
# TODO (PVP): this should later be handled by the forward fn() in each model in the future see PR 3140
|
||||
if self.config.is_encoder_decoder:
|
||||
if (attention_mask is None) and (pad_token_id is not None) and (pad_token_id in encoder_ids):
|
||||
enc_attention_mask = encoder_ids.ne(pad_token_id).long()
|
||||
elif attention_mask is None:
|
||||
enc_attention_mask = encoder_ids.new_ones(decoder_ids.shape)
|
||||
else:
|
||||
enc_attention_mask = attention_mask
|
||||
# TODO (Yacine): implement behavior for decoder side in encoder-decoder
|
||||
dec_attention_mask = None
|
||||
else:
|
||||
enc_attention_mask = None
|
||||
if (attention_mask is None) and (pad_token_id is not None) and (pad_token_id in decoder_ids):
|
||||
dec_attention_mask = decoder_ids.ne(pad_token_id).long()
|
||||
elif attention_mask is None:
|
||||
dec_attention_mask = decoder_ids.new_ones(decoder_ids.shape)
|
||||
else:
|
||||
dec_attention_mask = attention_mask
|
||||
if (attention_mask is None) and (pad_token_id is not None) and (pad_token_id in input_ids):
|
||||
attention_mask = input_ids.ne(pad_token_id).long()
|
||||
elif attention_mask is None:
|
||||
attention_mask = input_ids.new_ones(input_ids.shape)
|
||||
|
||||
# set pad_token_id to eos_token_id if not set. Important that this is done after
|
||||
# attention_mask is created
|
||||
@@ -943,54 +909,66 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
effective_batch_size = batch_size
|
||||
effective_batch_mult = 1
|
||||
|
||||
# Expand decoder_ids if num_beams > 1 or num_return_sequences > 1
|
||||
if num_return_sequences > 1 or num_beams > 1:
|
||||
decoder_ids_len = decoder_ids.shape[-1]
|
||||
decoder_ids = decoder_ids.unsqueeze(1).expand(batch_size, effective_batch_mult * num_beams, decoder_ids_len)
|
||||
decoder_ids = decoder_ids.contiguous().view(
|
||||
effective_batch_size * num_beams, decoder_ids_len
|
||||
) # shape: (batch_size * num_return_sequences * num_beams, cur_len)
|
||||
|
||||
|
||||
# compute encoder outputs once if necessary, and expand to match number of generated sequences
|
||||
if self.config.is_encoder_decoder:
|
||||
encoder = self.get_encoder()
|
||||
encoder_outputs = encoder(encoder_ids, attention_mask=enc_attention_mask)
|
||||
if decoder_start_token_id is None:
|
||||
decoder_start_token_id = bos_token_id
|
||||
|
||||
assert (
|
||||
batch_size == encoder_outputs[0].shape[0]
|
||||
), f"expected encoder_outputs[0] to have 1st dimension bs={batch_size}, got {encoder_outputs[0].shape[0]} "
|
||||
decoder_start_token_id is not None
|
||||
), "decoder_start_token_id or bos_token_id has to be defined for encoder-decoder generation"
|
||||
assert hasattr(self, "get_encoder"), "{} should have a 'get_encoder' function defined".format(self)
|
||||
assert callable(self.get_encoder), "{} should be a method".format(self.get_encoder)
|
||||
|
||||
# expand batch_idx to assign correct encoder output for expanded input_ids (due to num_beams > 1 and num_return_sequences > 1)
|
||||
expanded_batch_idxs = (
|
||||
# get encoder and store encoder outputs
|
||||
encoder = self.get_encoder()
|
||||
|
||||
encoder_outputs = encoder(input_ids, attention_mask=attention_mask)
|
||||
|
||||
# Expand input ids if num_beams > 1 or num_return_sequences > 1
|
||||
if num_return_sequences > 1 or num_beams > 1:
|
||||
input_ids_len = input_ids.shape[-1]
|
||||
input_ids = input_ids.unsqueeze(1).expand(batch_size, effective_batch_mult * num_beams, input_ids_len)
|
||||
attention_mask = attention_mask.unsqueeze(1).expand(
|
||||
batch_size, effective_batch_mult * num_beams, input_ids_len
|
||||
)
|
||||
|
||||
input_ids = input_ids.contiguous().view(
|
||||
effective_batch_size * num_beams, input_ids_len
|
||||
) # shape: (batch_size * num_return_sequences * num_beams, cur_len)
|
||||
attention_mask = attention_mask.contiguous().view(
|
||||
effective_batch_size * num_beams, input_ids_len
|
||||
) # shape: (batch_size * num_return_sequences * num_beams, cur_len)
|
||||
|
||||
if self.config.is_encoder_decoder:
|
||||
# create empty decoder_input_ids
|
||||
input_ids = torch.full(
|
||||
(effective_batch_size * num_beams, 1),
|
||||
decoder_start_token_id,
|
||||
dtype=torch.long,
|
||||
device=next(self.parameters()).device,
|
||||
)
|
||||
cur_len = 1
|
||||
batch_idx = self.encoder_outputs_batch_dim_idx
|
||||
assert (
|
||||
batch_size == encoder_outputs[0].shape[batch_idx]
|
||||
), f"expected encoder_outputs[0] to have 1st dimension bs={batch_size}, got {encoder_outputs[0].shape[1]} "
|
||||
expanded_idx = (
|
||||
torch.arange(batch_size)
|
||||
.view(-1, 1)
|
||||
.repeat(1, num_beams * effective_batch_mult)
|
||||
.view(-1)
|
||||
.to(decoder_ids.device)
|
||||
.to(input_ids.device)
|
||||
)
|
||||
encoder_outputs = (encoder_outputs[0].index_select(0, expanded_batch_idxs), *encoder_outputs[1:]) # (x, encoder_states, all_attentions)
|
||||
enc_attention_mask = enc_attention_mask.index_select(0, expanded_batch_idxs)
|
||||
# TODO (Yacine): attention_mask was prepared for the encoder, not the decoder
|
||||
encoder_outputs = (encoder_outputs[0].index_select(batch_idx, expanded_idx), *encoder_outputs[1:])
|
||||
|
||||
else:
|
||||
encoder_outputs = None
|
||||
dec_attention_mask = dec_attention_mask.unsqueeze(1).expand(
|
||||
batch_size, effective_batch_mult * num_beams, decoder_ids_len
|
||||
)
|
||||
dec_attention_mask = dec_attention_mask.contiguous().view(
|
||||
effective_batch_size * num_beams, decoder_ids_len
|
||||
) # shape: (batch_size * num_return_sequences * num_beams, cur_len)
|
||||
|
||||
cur_len = decoder_ids.shape[-1]
|
||||
cur_len = input_ids.shape[-1]
|
||||
|
||||
if num_beams > 1:
|
||||
output = self._generate_beam_search(
|
||||
input_ids=decoder_ids,
|
||||
input_ids,
|
||||
cur_len=cur_len,
|
||||
encoder_outputs=encoder_outputs,
|
||||
encoder_attention_mask=enc_attention_mask,
|
||||
decoder_attention_mask=dec_attention_mask,
|
||||
max_length=max_length,
|
||||
min_length=min_length,
|
||||
do_sample=do_sample,
|
||||
@@ -1003,21 +981,20 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
bad_words_ids=bad_words_ids,
|
||||
bos_token_id=bos_token_id,
|
||||
pad_token_id=pad_token_id,
|
||||
decoder_start_token_id=decoder_start_token_id,
|
||||
eos_token_id=eos_token_id,
|
||||
batch_size=effective_batch_size,
|
||||
num_return_sequences=num_return_sequences,
|
||||
length_penalty=length_penalty,
|
||||
num_beams=num_beams,
|
||||
vocab_size=vocab_size,
|
||||
use_cache=use_cache,
|
||||
encoder_outputs=encoder_outputs,
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
else:
|
||||
output = self._generate_no_beam_search(
|
||||
input_ids=decoder_ids,
|
||||
input_ids,
|
||||
cur_len=cur_len,
|
||||
encoder_outputs=encoder_outputs,
|
||||
encoder_attention_mask=enc_attention_mask,
|
||||
decoder_attention_mask=dec_attention_mask,
|
||||
max_length=max_length,
|
||||
min_length=min_length,
|
||||
do_sample=do_sample,
|
||||
@@ -1029,9 +1006,11 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
bad_words_ids=bad_words_ids,
|
||||
bos_token_id=bos_token_id,
|
||||
pad_token_id=pad_token_id,
|
||||
decoder_start_token_id=decoder_start_token_id,
|
||||
eos_token_id=eos_token_id,
|
||||
batch_size=effective_batch_size,
|
||||
use_cache=use_cache,
|
||||
encoder_outputs=encoder_outputs,
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
|
||||
return output
|
||||
@@ -1040,9 +1019,6 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
self,
|
||||
input_ids,
|
||||
cur_len,
|
||||
encoder_outputs,
|
||||
encoder_attention_mask,
|
||||
decoder_attention_mask,
|
||||
max_length,
|
||||
min_length,
|
||||
do_sample,
|
||||
@@ -1055,8 +1031,10 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
bos_token_id,
|
||||
pad_token_id,
|
||||
eos_token_id,
|
||||
decoder_start_token_id,
|
||||
batch_size,
|
||||
use_cache,
|
||||
encoder_outputs,
|
||||
attention_mask,
|
||||
):
|
||||
""" Generate sequences for each example without beam search (num_beams == 1).
|
||||
All returned sequence are generated independantly.
|
||||
@@ -1065,18 +1043,17 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
unfinished_sents = input_ids.new(batch_size).fill_(1)
|
||||
sent_lengths = input_ids.new(batch_size).fill_(max_length)
|
||||
|
||||
past = None if encoder_outputs is None else ((encoder_outputs, encoder_attention_mask), None) # defined for encoder-decoder models, None for decoder-only models
|
||||
past = encoder_outputs # defined for encoder-decoder models, None for decoder-only models
|
||||
|
||||
while cur_len < max_length:
|
||||
model_inputs = self.prepare_inputs_for_generation(input_ids, past=past, attention_mask=decoder_attention_mask)
|
||||
model_inputs = self.prepare_inputs_for_generation(input_ids, past=past, attention_mask=attention_mask)
|
||||
|
||||
outputs = self(**model_inputs)
|
||||
next_token_logits = outputs[0][:, -1, :]
|
||||
|
||||
# if model has past, then set the past variable to speed up decoding
|
||||
if use_cache:
|
||||
_, decoder_cache = outputs[1]
|
||||
past = ((encoder_outputs, encoder_attention_mask), decoder_cache)
|
||||
if self._do_output_past(outputs):
|
||||
past = outputs[1]
|
||||
|
||||
# repetition penalty from CTRL paper (https://arxiv.org/abs/1909.05858)
|
||||
if repetition_penalty != 1.0:
|
||||
@@ -1159,9 +1136,6 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
self,
|
||||
input_ids,
|
||||
cur_len,
|
||||
encoder_outputs,
|
||||
encoder_attention_mask,
|
||||
decoder_attention_mask,
|
||||
max_length,
|
||||
min_length,
|
||||
do_sample,
|
||||
@@ -1175,12 +1149,14 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
bos_token_id,
|
||||
pad_token_id,
|
||||
eos_token_id,
|
||||
decoder_start_token_id,
|
||||
batch_size,
|
||||
num_return_sequences,
|
||||
length_penalty,
|
||||
num_beams,
|
||||
vocab_size,
|
||||
use_cache,
|
||||
encoder_outputs,
|
||||
attention_mask,
|
||||
):
|
||||
""" Generate sequences for each example with beam search.
|
||||
"""
|
||||
@@ -1200,20 +1176,19 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
beam_scores = beam_scores.view(-1) # shape (batch_size * num_beams,)
|
||||
|
||||
# cache compute states
|
||||
past = None if encoder_outputs is None else ((encoder_outputs, encoder_attention_mask), None) # defined for encoder-decoder models, None for decoder-only models
|
||||
past = encoder_outputs # defined for encoder-decoder models, None for decoder-only models
|
||||
|
||||
# done sentences
|
||||
done = [False for _ in range(batch_size)]
|
||||
|
||||
while cur_len < max_length:
|
||||
model_inputs = self.prepare_inputs_for_generation(input_ids, past=past, attention_mask=decoder_attention_mask)
|
||||
model_inputs = self.prepare_inputs_for_generation(input_ids, past=past, attention_mask=attention_mask)
|
||||
outputs = self(**model_inputs) # (batch_size * num_beams, cur_len, vocab_size)
|
||||
next_token_logits = outputs[0][:, -1, :] # (batch_size * num_beams, vocab_size)
|
||||
|
||||
# if model has past, then set the past variable to speed up decoding
|
||||
if use_cache:
|
||||
_, decoder_cache = outputs[1]
|
||||
past = ((encoder_outputs, encoder_attention_mask), decoder_cache)
|
||||
if self._do_output_past(outputs):
|
||||
past = outputs[1]
|
||||
|
||||
# repetition penalty (from CTRL paper https://arxiv.org/abs/1909.05858)
|
||||
if repetition_penalty != 1.0:
|
||||
|
||||
Reference in New Issue
Block a user