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
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:]
positions = positions[:, -1:] # happens after we embed them
assert input_ids.ne(self.padding_idx).any()
@@ -476,7 +476,7 @@ class BartDecoder(nn.Module):
causal_mask=decoder_causal_mask,
)
if self.output_past:
if generation_mode:
next_decoder_cache.append(layer_past.copy())
if self.output_hidden_states:
all_hidden_states += (x,)
@@ -488,7 +488,7 @@ class BartDecoder(nn.Module):
x = x.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)
else:
next_cache = None
@@ -907,23 +907,18 @@ 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
if not past[1]:
encoder_outputs, decoder_cached_states = past, None
else:
encoder_outputs, decoder_cached_states = past
# first step, decoder_cached_states are empty (None)
(encoder_outputs, encoder_attention_mask), 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
@@ -931,15 +926,19 @@ class BartForConditionalGeneration(PretrainedBartModel):
@staticmethod
def _reorder_cache(past, beam_idx):
((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)
past = ((new_enc_out, new_enc_mask), reordered_past)
+131 -109
View File
@@ -658,24 +658,26 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
def generate(
self,
input_ids=None,
max_length=None,
attention_mask=None,
encoder_input_ids=None,
decoder_input_ids=None,
min_length=None,
do_sample=None,
max_length=None,
early_stopping=None,
num_return_sequences=None,
num_beams=None,
do_sample=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,
num_return_sequences=None,
attention_mask=None,
decoder_start_token_id=None,
repetition_penalty=None,
use_cache=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.
@@ -688,24 +690,42 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
Parameters:
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
it as an empty `torch.LongTensor` of shape `(1,)`.
Short-hand for either encoder_input_ids in seq2seq models or decoder_input_ids for language models.
max_length: (`optional`) int
The max length of the sequence to be generated. Between `min_length` and infinity. Default to 20.
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.
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`.
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.
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
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
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.
@@ -732,23 +749,16 @@ 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)`.
num_return_sequences: (`optional`) int
The number of independently computed returned sequences for each element in the batch. Default to 1.
use_cache: (`optional`) bool
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:
@@ -800,16 +810,18 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
"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
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
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
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
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
@@ -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
)
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
@@ -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 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."
@@ -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 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, (
"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."
)
input_ids = torch.full(
decoder_ids = torch.full(
(batch_size, 1), bos_token_id, dtype=torch.long, device=next(self.parameters()).device,
)
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
if do_sample is False:
@@ -876,7 +898,6 @@ 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 (
@@ -885,10 +906,23 @@ 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 (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)
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
# set pad_token_id to eos_token_id if not set. Important that this is done after
# attention_mask is created
@@ -909,45 +943,19 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
effective_batch_size = batch_size
effective_batch_mult = 1
if self.config.is_encoder_decoder:
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
# Expand decoder_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
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:
# 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
encoder = self.get_encoder()
encoder_outputs = encoder(encoder_ids, attention_mask=enc_attention_mask)
assert (
batch_size == encoder_outputs[0].shape[0]
@@ -959,19 +967,30 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
.view(-1, 1)
.repeat(1, num_beams * effective_batch_mult)
.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:])
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
else:
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:
output = self._generate_beam_search(
input_ids,
input_ids=decoder_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,
@@ -984,20 +1003,21 @@ 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,
encoder_outputs=encoder_outputs,
attention_mask=attention_mask,
use_cache=use_cache,
)
else:
output = self._generate_no_beam_search(
input_ids,
input_ids=decoder_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,
@@ -1009,11 +1029,9 @@ 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,
encoder_outputs=encoder_outputs,
attention_mask=attention_mask,
use_cache=use_cache,
)
return output
@@ -1022,6 +1040,9 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
self,
input_ids,
cur_len,
encoder_outputs,
encoder_attention_mask,
decoder_attention_mask,
max_length,
min_length,
do_sample,
@@ -1034,10 +1055,8 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
bos_token_id,
pad_token_id,
eos_token_id,
decoder_start_token_id,
batch_size,
encoder_outputs,
attention_mask,
use_cache,
):
""" Generate sequences for each example without beam search (num_beams == 1).
All returned sequence are generated independantly.
@@ -1046,17 +1065,18 @@ 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 = 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:
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)
next_token_logits = outputs[0][:, -1, :]
# if model has past, then set the past variable to speed up decoding
if self._do_output_past(outputs):
past = outputs[1]
if use_cache:
_, decoder_cache = outputs[1]
past = ((encoder_outputs, encoder_attention_mask), decoder_cache)
# repetition penalty from CTRL paper (https://arxiv.org/abs/1909.05858)
if repetition_penalty != 1.0:
@@ -1139,6 +1159,9 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
self,
input_ids,
cur_len,
encoder_outputs,
encoder_attention_mask,
decoder_attention_mask,
max_length,
min_length,
do_sample,
@@ -1152,14 +1175,12 @@ 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,
encoder_outputs,
attention_mask,
use_cache,
):
""" 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,)
# 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 = [False for _ in range(batch_size)]
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)
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 self._do_output_past(outputs):
past = outputs[1]
if use_cache:
_, decoder_cache = outputs[1]
past = ((encoder_outputs, encoder_attention_mask), decoder_cache)
# repetition penalty (from CTRL paper https://arxiv.org/abs/1909.05858)
if repetition_penalty != 1.0: