Compare commits

..
Author SHA1 Message Date
Thomas Wolf 3bfeebe722 Less general but avoid hook issues 2020-04-02 10:00:35 +02:00
Thomas Wolf ffed6a8a5d Update modeling_t5.py
Style and quality
2020-04-01 23:28:36 +02:00
Thomas Wolf 426b5106e7 Adding spread_on_devices 2020-04-01 23:21:10 +02:00
6 changed files with 213 additions and 183 deletions
-1
View File
@@ -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()
+1 -1
View File
@@ -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),
+24 -29
View File
@@ -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)
+68
View File
@@ -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
View File
@@ -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: