Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d1406e9a63 | ||
|
|
76f5cd8b7c | ||
|
|
ae98046330 |
@@ -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
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user