|
|
|
@@ -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,66 +943,54 @@ 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
|
|
|
|
|
batch_idx = self.encoder_outputs_batch_dim_idx
|
|
|
|
|
encoder = self.get_encoder()
|
|
|
|
|
encoder_outputs = encoder(encoder_ids, attention_mask=enc_attention_mask)
|
|
|
|
|
|
|
|
|
|
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 = (
|
|
|
|
|
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]} "
|
|
|
|
|
|
|
|
|
|
# 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 = (
|
|
|
|
|
torch.arange(batch_size)
|
|
|
|
|
.view(-1, 1)
|
|
|
|
|
.repeat(1, num_beams * effective_batch_mult)
|
|
|
|
|
.view(-1)
|
|
|
|
|
.to(input_ids.device)
|
|
|
|
|
.to(decoder_ids.device)
|
|
|
|
|
)
|
|
|
|
|
encoder_outputs = (encoder_outputs[0].index_select(batch_idx, expanded_idx), *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,
|
|
|
|
@@ -981,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,
|
|
|
|
@@ -1006,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
|
|
|
|
@@ -1019,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,
|
|
|
|
@@ -1031,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.
|
|
|
|
@@ -1043,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:
|
|
|
|
@@ -1136,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,
|
|
|
|
@@ -1149,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.
|
|
|
|
|
"""
|
|
|
|
@@ -1176,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:
|
|
|
|
|