Compare commits

...
Author SHA1 Message Date
yjernite d1406e9a63 resolved_decode_with_prefix 2020-04-02 14:19:12 -04:00
yjernite 76f5cd8b7c resolved_decode_with_prefix 2020-04-02 14:14:11 -04:00
yjernite ae98046330 decode_with_prefix 2020-04-02 14:09:31 -04:00
2 changed files with 149 additions and 128 deletions
+18 -19
View File
@@ -440,7 +440,7 @@ class BartDecoder(nn.Module):
# embed positions # embed positions
positions = self.embed_positions(input_ids, generation_mode=generation_mode) positions = self.embed_positions(input_ids, generation_mode=generation_mode)
if generation_mode: if generation_mode and decoder_cached_states is not None:
input_ids = input_ids[:, -1:] input_ids = input_ids[:, -1:]
positions = positions[:, -1:] # happens after we embed them positions = positions[:, -1:] # happens after we embed them
assert input_ids.ne(self.padding_idx).any() assert input_ids.ne(self.padding_idx).any()
@@ -476,7 +476,7 @@ class BartDecoder(nn.Module):
causal_mask=decoder_causal_mask, causal_mask=decoder_causal_mask,
) )
if self.output_past: if generation_mode:
next_decoder_cache.append(layer_past.copy()) next_decoder_cache.append(layer_past.copy())
if self.output_hidden_states: if self.output_hidden_states:
all_hidden_states += (x,) all_hidden_states += (x,)
@@ -488,7 +488,7 @@ class BartDecoder(nn.Module):
x = x.transpose(0, 1) x = x.transpose(0, 1)
encoder_hidden_states = encoder_hidden_states.transpose(0, 1) encoder_hidden_states = encoder_hidden_states.transpose(0, 1)
if self.output_past: if generation_mode:
next_cache = ((encoder_hidden_states, encoder_padding_mask), next_decoder_cache) next_cache = ((encoder_hidden_states, encoder_padding_mask), next_decoder_cache)
else: else:
next_cache = None next_cache = None
@@ -907,23 +907,18 @@ class BartForConditionalGeneration(PretrainedBartModel):
def prepare_inputs_for_generation(self, decoder_input_ids, past, attention_mask, **kwargs): 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" assert past is not None, "past has to be defined for encoder_outputs"
# first step, decoder_cached_states are empty # first step, decoder_cached_states are empty (None)
if not past[1]: (encoder_outputs, encoder_attention_mask), decoder_cached_states = past
encoder_outputs, decoder_cached_states = past, None
else:
encoder_outputs, decoder_cached_states = past
return { return {
"input_ids": None, # encoder_outputs is defined. input_ids not needed "input_ids": None, # encoder_outputs is defined. input_ids not needed
"encoder_outputs": encoder_outputs, "encoder_outputs": encoder_outputs,
"attention_mask": encoder_attention_mask,
"decoder_cached_states": decoder_cached_states, "decoder_cached_states": decoder_cached_states,
"decoder_input_ids": decoder_input_ids, "decoder_input_ids": decoder_input_ids,
"attention_mask": attention_mask,
"generation_mode": True, "generation_mode": True,
} }
def prepare_scores_for_generation(self, scores, cur_len, max_length): 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: 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) self._force_token_ids_generation(scores, self.config.eos_token_id)
return scores return scores
@@ -931,15 +926,19 @@ class BartForConditionalGeneration(PretrainedBartModel):
@staticmethod @staticmethod
def _reorder_cache(past, beam_idx): def _reorder_cache(past, beam_idx):
((enc_out, enc_mask), decoder_cached_states) = past ((enc_out, enc_mask), decoder_cached_states) = past
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)
new_enc_out = enc_out if enc_out is None else enc_out.index_select(0, beam_idx) 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:])
new_enc_mask = enc_mask if enc_mask is None else enc_mask.index_select(0, 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) past = ((new_enc_out, new_enc_mask), reordered_past)
+131 -109
View File
@@ -658,24 +658,26 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
def generate( def generate(
self, self,
input_ids=None, input_ids=None,
max_length=None, attention_mask=None,
encoder_input_ids=None,
decoder_input_ids=None,
min_length=None, min_length=None,
do_sample=None, max_length=None,
early_stopping=None, early_stopping=None,
num_return_sequences=None,
num_beams=None, num_beams=None,
do_sample=None,
temperature=None, temperature=None,
top_k=None, top_k=None,
top_p=None, top_p=None,
repetition_penalty=None,
bad_words_ids=None, bad_words_ids=None,
bos_token_id=None, bos_token_id=None,
pad_token_id=None, pad_token_id=None,
eos_token_id=None, eos_token_id=None,
length_penalty=None, length_penalty=None,
no_repeat_ngram_size=None, no_repeat_ngram_size=None,
num_return_sequences=None, repetition_penalty=None,
attention_mask=None, use_cache=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. 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.
@@ -688,24 +690,42 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
Parameters: Parameters:
input_ids: (`optional`) `torch.LongTensor` of shape `(batch_size, sequence_length)` input_ids: (`optional`) `torch.LongTensor` of shape `(batch_size, sequence_length)`
The sequence used as a prompt for the generation. If `None` the method initializes Short-hand for either encoder_input_ids in seq2seq models or decoder_input_ids for language models.
it as an empty `torch.LongTensor` of shape `(1,)`.
max_length: (`optional`) int attention_mask (`optional`) obj: `torch.LongTensor` of same shape as `input_ids`
The max length of the sequence to be generated. Between `min_length` and infinity. Default to 20. 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 min_length: (`optional`) int
The min length of the sequence to be generated. Between 0 and infinity. Default to 0. The min length of the sequence to be generated. Between 0 and infinity. Default to 0.
do_sample: (`optional`) bool max_length: (`optional`) int
If set to `False` greedy decoding is used. Otherwise sampling is used. Defaults to `False` as defined in `configuration_utils.PretrainedConfig`. The max length of the sequence to be generated. Between `min_length` and infinity. Default to 20.
early_stopping: (`optional`) bool 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`. 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 num_beams: (`optional`) int
Number of beams for beam search. Must be between 1 and infinity. 1 means no beam search. Default to 1. Number of beams for beam search. Must be between 1 and infinity. 1 means no beam search. Default to 1.
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`.
temperature: (`optional`) float temperature: (`optional`) float
The value used to module the next token probabilities. Must be strictly positive. Default to 1.0. The value used to module the next token probabilities. Must be strictly positive. Default to 1.0.
@@ -715,9 +735,6 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
top_p: (`optional`) float 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. 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 pad_token_id: (`optional`) int
Padding token. Default to specicic model pad_token_id or None if it does not exist. Padding token. Default to specicic model pad_token_id or None if it does not exist.
@@ -732,23 +749,16 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
no_repeat_ngram_size: (`optional`) int no_repeat_ngram_size: (`optional`) int
If set to int > 0, all ngrams of size `no_repeat_ngram_size` can only occur once. 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: (`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)`. `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)`.
num_return_sequences: (`optional`) int use_cache: (`optional`) bool
The number of independently computed returned sequences for each element in the batch. Default to 1. If set to `True` the model re-uses pre-computed decoder hidden states from one time step to the next.
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: Return:
@@ -800,16 +810,18 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
"Please use another model class (e.g. `OpenAIGPTLMHeadModel`, `XLNetLMHeadModel`, `GPT2LMHeadModel`, `CTRLLMHeadModel`, `T5WithLMHeadModel`, `TransfoXLLMHeadModel`, `XLMWithLMHeadModel`, `BartForConditionalGeneration` )" "Please use another model class (e.g. `OpenAIGPTLMHeadModel`, `XLNetLMHeadModel`, `GPT2LMHeadModel`, `CTRLLMHeadModel`, `T5WithLMHeadModel`, `TransfoXLLMHeadModel`, `XLMWithLMHeadModel`, `BartForConditionalGeneration` )"
) )
max_length = max_length if max_length is not None else self.config.max_length
min_length = min_length if min_length is not None else self.config.min_length 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 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 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 num_beams = num_beams if num_beams is not None else self.config.num_beams
do_sample = do_sample if do_sample is not None else self.config.do_sample
temperature = temperature if temperature is not None else self.config.temperature 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_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 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 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 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 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 length_penalty = length_penalty if length_penalty is not None else self.config.length_penalty
@@ -817,15 +829,12 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
no_repeat_ngram_size if no_repeat_ngram_size is not None else self.config.no_repeat_ngram_size 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 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: if input_ids is not None:
batch_size = input_ids.shape[0] # overriden by the input batch_size 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: else:
batch_size = 1 batch_size = 1
@@ -838,9 +847,6 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
assert isinstance(top_k, int) and top_k >= 0, "`top_k` should be a positive integer." 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 0 <= top_p <= 1, "`top_p` should be between 0 and 1."
assert repetition_penalty >= 1.0, "`repetition_penalty` should be >= 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 ( assert pad_token_id is None or (
isinstance(pad_token_id, int) and (pad_token_id >= 0) isinstance(pad_token_id, int) and (pad_token_id >= 0)
), "`pad_token_id` should be a positive integer." ), "`pad_token_id` should be a positive integer."
@@ -858,16 +864,32 @@ 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 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" ), "`bad_words_ids` is either `None` or a list of lists of tokens that should not be generated"
if input_ids is None: # 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:
assert isinstance(bos_token_id, int) and bos_token_id >= 0, ( assert isinstance(bos_token_id, int) and bos_token_id >= 0, (
"you should either supply a context to complete as `input_ids` input " "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." "or a `bos_token_id` (integer >= 0) as a first token to start the generation."
) )
input_ids = torch.full( decoder_ids = torch.full(
(batch_size, 1), bos_token_id, dtype=torch.long, device=next(self.parameters()).device, (batch_size, 1), bos_token_id, dtype=torch.long, device=next(self.parameters()).device,
) )
else: else:
assert input_ids.dim() == 2, "Input prompt should be of shape (batch_size, sequence length)." assert decoder_ids.dim() == 2, "Input prompt should be of shape (batch_size, sequence length)."
# not allow to duplicate outputs when greedy decoding # not allow to duplicate outputs when greedy decoding
if do_sample is False: if do_sample is False:
@@ -876,7 +898,6 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
assert ( assert (
num_return_sequences == 1 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" ), "Greedy decoding will always produce the same output for num_beams == 1 and num_return_sequences > 1. Please set num_return_sequences = 1"
else: else:
# beam_search greedy generation conditions # beam_search greedy generation conditions
assert ( assert (
@@ -885,10 +906,23 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
# create attention mask if necessary # 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 # TODO (PVP): this should later be handled by the forward fn() in each model in the future see PR 3140
if (attention_mask is None) and (pad_token_id is not None) and (pad_token_id in input_ids): if self.config.is_encoder_decoder:
attention_mask = input_ids.ne(pad_token_id).long() if (attention_mask is None) and (pad_token_id is not None) and (pad_token_id in encoder_ids):
elif attention_mask is None: enc_attention_mask = encoder_ids.ne(pad_token_id).long()
attention_mask = input_ids.new_ones(input_ids.shape) 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
# set pad_token_id to eos_token_id if not set. Important that this is done after # set pad_token_id to eos_token_id if not set. Important that this is done after
# attention_mask is created # attention_mask is created
@@ -909,45 +943,19 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
effective_batch_size = batch_size effective_batch_size = batch_size
effective_batch_mult = 1 effective_batch_mult = 1
if self.config.is_encoder_decoder: # Expand decoder_ids if num_beams > 1 or num_return_sequences > 1
if decoder_start_token_id is None:
decoder_start_token_id = bos_token_id
assert (
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)
# 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: if num_return_sequences > 1 or num_beams > 1:
input_ids_len = input_ids.shape[-1] decoder_ids_len = decoder_ids.shape[-1]
input_ids = input_ids.unsqueeze(1).expand(batch_size, effective_batch_mult * num_beams, input_ids_len) decoder_ids = decoder_ids.unsqueeze(1).expand(batch_size, effective_batch_mult * num_beams, decoder_ids_len)
attention_mask = attention_mask.unsqueeze(1).expand( decoder_ids = decoder_ids.contiguous().view(
batch_size, effective_batch_mult * num_beams, input_ids_len effective_batch_size * num_beams, decoder_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) ) # 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: if self.config.is_encoder_decoder:
# create empty decoder_input_ids encoder = self.get_encoder()
input_ids = torch.full( encoder_outputs = encoder(encoder_ids, attention_mask=enc_attention_mask)
(effective_batch_size * num_beams, 1),
decoder_start_token_id,
dtype=torch.long,
device=next(self.parameters()).device,
)
cur_len = 1
assert ( assert (
batch_size == encoder_outputs[0].shape[0] batch_size == encoder_outputs[0].shape[0]
@@ -959,19 +967,30 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
.view(-1, 1) .view(-1, 1)
.repeat(1, num_beams * effective_batch_mult) .repeat(1, num_beams * effective_batch_mult)
.view(-1) .view(-1)
.to(input_ids.device) .to(decoder_ids.device)
) )
# expand encoder_outputs encoder_outputs = (encoder_outputs[0].index_select(0, expanded_batch_idxs), *encoder_outputs[1:]) # (x, encoder_states, all_attentions)
encoder_outputs = (encoder_outputs[0].index_select(0, expanded_batch_idxs), *encoder_outputs[1:]) enc_attention_mask = enc_attention_mask.index_select(0, expanded_batch_idxs)
# TODO (Yacine): attention_mask was prepared for the encoder, not the decoder
else: else:
encoder_outputs = None encoder_outputs = None
cur_len = input_ids.shape[-1] 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]
if num_beams > 1: if num_beams > 1:
output = self._generate_beam_search( output = self._generate_beam_search(
input_ids, input_ids=decoder_ids,
cur_len=cur_len, cur_len=cur_len,
encoder_outputs=encoder_outputs,
encoder_attention_mask=enc_attention_mask,
decoder_attention_mask=dec_attention_mask,
max_length=max_length, max_length=max_length,
min_length=min_length, min_length=min_length,
do_sample=do_sample, do_sample=do_sample,
@@ -984,20 +1003,21 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
bad_words_ids=bad_words_ids, bad_words_ids=bad_words_ids,
bos_token_id=bos_token_id, bos_token_id=bos_token_id,
pad_token_id=pad_token_id, pad_token_id=pad_token_id,
decoder_start_token_id=decoder_start_token_id,
eos_token_id=eos_token_id, eos_token_id=eos_token_id,
batch_size=effective_batch_size, batch_size=effective_batch_size,
num_return_sequences=num_return_sequences, num_return_sequences=num_return_sequences,
length_penalty=length_penalty, length_penalty=length_penalty,
num_beams=num_beams, num_beams=num_beams,
vocab_size=vocab_size, vocab_size=vocab_size,
encoder_outputs=encoder_outputs, use_cache=use_cache,
attention_mask=attention_mask,
) )
else: else:
output = self._generate_no_beam_search( output = self._generate_no_beam_search(
input_ids, input_ids=decoder_ids,
cur_len=cur_len, cur_len=cur_len,
encoder_outputs=encoder_outputs,
encoder_attention_mask=enc_attention_mask,
decoder_attention_mask=dec_attention_mask,
max_length=max_length, max_length=max_length,
min_length=min_length, min_length=min_length,
do_sample=do_sample, do_sample=do_sample,
@@ -1009,11 +1029,9 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
bad_words_ids=bad_words_ids, bad_words_ids=bad_words_ids,
bos_token_id=bos_token_id, bos_token_id=bos_token_id,
pad_token_id=pad_token_id, pad_token_id=pad_token_id,
decoder_start_token_id=decoder_start_token_id,
eos_token_id=eos_token_id, eos_token_id=eos_token_id,
batch_size=effective_batch_size, batch_size=effective_batch_size,
encoder_outputs=encoder_outputs, use_cache=use_cache,
attention_mask=attention_mask,
) )
return output return output
@@ -1022,6 +1040,9 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
self, self,
input_ids, input_ids,
cur_len, cur_len,
encoder_outputs,
encoder_attention_mask,
decoder_attention_mask,
max_length, max_length,
min_length, min_length,
do_sample, do_sample,
@@ -1034,10 +1055,8 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
bos_token_id, bos_token_id,
pad_token_id, pad_token_id,
eos_token_id, eos_token_id,
decoder_start_token_id,
batch_size, batch_size,
encoder_outputs, use_cache,
attention_mask,
): ):
""" Generate sequences for each example without beam search (num_beams == 1). """ Generate sequences for each example without beam search (num_beams == 1).
All returned sequence are generated independantly. All returned sequence are generated independantly.
@@ -1046,17 +1065,18 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
unfinished_sents = input_ids.new(batch_size).fill_(1) unfinished_sents = input_ids.new(batch_size).fill_(1)
sent_lengths = input_ids.new(batch_size).fill_(max_length) sent_lengths = input_ids.new(batch_size).fill_(max_length)
past = encoder_outputs # defined for encoder-decoder models, None for decoder-only models past = None if encoder_outputs is None else ((encoder_outputs, encoder_attention_mask), None) # defined for encoder-decoder models, None for decoder-only models
while cur_len < max_length: while cur_len < max_length:
model_inputs = self.prepare_inputs_for_generation(input_ids, past=past, attention_mask=attention_mask) model_inputs = self.prepare_inputs_for_generation(input_ids, past=past, attention_mask=decoder_attention_mask)
outputs = self(**model_inputs) outputs = self(**model_inputs)
next_token_logits = outputs[0][:, -1, :] next_token_logits = outputs[0][:, -1, :]
# if model has past, then set the past variable to speed up decoding # if model has past, then set the past variable to speed up decoding
if self._do_output_past(outputs): if use_cache:
past = outputs[1] _, decoder_cache = outputs[1]
past = ((encoder_outputs, encoder_attention_mask), decoder_cache)
# repetition penalty from CTRL paper (https://arxiv.org/abs/1909.05858) # repetition penalty from CTRL paper (https://arxiv.org/abs/1909.05858)
if repetition_penalty != 1.0: if repetition_penalty != 1.0:
@@ -1139,6 +1159,9 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
self, self,
input_ids, input_ids,
cur_len, cur_len,
encoder_outputs,
encoder_attention_mask,
decoder_attention_mask,
max_length, max_length,
min_length, min_length,
do_sample, do_sample,
@@ -1152,14 +1175,12 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
bos_token_id, bos_token_id,
pad_token_id, pad_token_id,
eos_token_id, eos_token_id,
decoder_start_token_id,
batch_size, batch_size,
num_return_sequences, num_return_sequences,
length_penalty, length_penalty,
num_beams, num_beams,
vocab_size, vocab_size,
encoder_outputs, use_cache,
attention_mask,
): ):
""" Generate sequences for each example with beam search. """ Generate sequences for each example with beam search.
""" """
@@ -1179,19 +1200,20 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
beam_scores = beam_scores.view(-1) # shape (batch_size * num_beams,) beam_scores = beam_scores.view(-1) # shape (batch_size * num_beams,)
# cache compute states # cache compute states
past = encoder_outputs # defined for encoder-decoder models, None for decoder-only models past = None if encoder_outputs is None else ((encoder_outputs, encoder_attention_mask), None) # defined for encoder-decoder models, None for decoder-only models
# done sentences # done sentences
done = [False for _ in range(batch_size)] done = [False for _ in range(batch_size)]
while cur_len < max_length: while cur_len < max_length:
model_inputs = self.prepare_inputs_for_generation(input_ids, past=past, attention_mask=attention_mask) model_inputs = self.prepare_inputs_for_generation(input_ids, past=past, attention_mask=decoder_attention_mask)
outputs = self(**model_inputs) # (batch_size * num_beams, cur_len, vocab_size) outputs = self(**model_inputs) # (batch_size * num_beams, cur_len, vocab_size)
next_token_logits = outputs[0][:, -1, :] # (batch_size * num_beams, 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 model has past, then set the past variable to speed up decoding
if self._do_output_past(outputs): if use_cache:
past = outputs[1] _, decoder_cache = outputs[1]
past = ((encoder_outputs, encoder_attention_mask), decoder_cache)
# repetition penalty (from CTRL paper https://arxiv.org/abs/1909.05858) # repetition penalty (from CTRL paper https://arxiv.org/abs/1909.05858)
if repetition_penalty != 1.0: if repetition_penalty != 1.0: