Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
dac7031a62 |
@@ -70,9 +70,9 @@ class PretrainedConfig(object):
|
||||
self.top_k = kwargs.pop("top_k", 50)
|
||||
self.top_p = kwargs.pop("top_p", 1.0)
|
||||
self.repetition_penalty = kwargs.pop("repetition_penalty", 1.0)
|
||||
self.bos_token_id = kwargs.pop("bos_token_id", 0)
|
||||
self.pad_token_id = kwargs.pop("pad_token_id", 0)
|
||||
self.eos_token_ids = kwargs.pop("eos_token_ids", 0)
|
||||
self.bos_token_id = kwargs.pop("bos_token_id", None)
|
||||
self.pad_token_id = kwargs.pop("pad_token_id", None)
|
||||
self.eos_token_ids = kwargs.pop("eos_token_ids", None)
|
||||
self.length_penalty = kwargs.pop("length_penalty", 1.0)
|
||||
self.num_return_sequences = kwargs.pop("num_return_sequences", 1)
|
||||
|
||||
|
||||
@@ -709,17 +709,20 @@ 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 isinstance(bos_token_id, int) and bos_token_id >= 0, "`bos_token_id` should be a positive integer."
|
||||
assert isinstance(pad_token_id, int) and pad_token_id >= 0, "`pad_token_id` should be a positive integer."
|
||||
assert isinstance(eos_token_ids, (list, tuple)) and (
|
||||
assert input_ids is not None or (isinstance(bos_token_id, int) and bos_token_id >= 0), "`bos_token_id` should be a positive integer."
|
||||
assert (eos_token_ids is None) or (isinstance(pad_token_id, int) and pad_token_id >= 0), "`pad_token_id` should be a positive integer."
|
||||
assert (eos_token_ids is None) or (isinstance(eos_token_ids, (list, tuple)) and (
|
||||
e >= 0 for e in eos_token_ids
|
||||
), "`eos_token_ids` should be a positive integer or a list/tuple of positive integers."
|
||||
)), "`eos_token_ids` should be a positive integer or a list/tuple of positive integers."
|
||||
assert length_penalty > 0, "`length_penalty` should be strictely positive."
|
||||
assert (
|
||||
isinstance(num_return_sequences, int) and num_return_sequences > 0
|
||||
), "`num_return_sequences` should be a strictely positive integer."
|
||||
|
||||
if input_ids is None:
|
||||
assert (isinstance(bos_token_id, int) and bos_token_id >= 0,
|
||||
"you should either supply a context to complete as `input_ids` input "
|
||||
"or a `bos_token_id` (integer >= 0) as a first token to start the generation.")
|
||||
input_ids = torch.full(
|
||||
(batch_size, 1), bos_token_id, dtype=torch.long, device=next(self.parameters()).device
|
||||
)
|
||||
@@ -830,18 +833,22 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
next_token = torch.argmax(next_token_logits, dim=-1)
|
||||
|
||||
# update generations and finished sentences
|
||||
tokens_to_add = next_token * unfinished_sents + pad_token_id * (1 - unfinished_sents)
|
||||
if eos_token_ids is not None:
|
||||
tokens_to_add = next_token * unfinished_sents + pad_token_id * (1 - unfinished_sents)
|
||||
else:
|
||||
tokens_to_add = next_token
|
||||
input_ids = torch.cat([input_ids, tokens_to_add.unsqueeze(-1)], dim=-1)
|
||||
for eos_token_id in eos_token_ids:
|
||||
unfinished_sents.mul_(tokens_to_add.ne(eos_token_id).long())
|
||||
if eos_token_ids is not None:
|
||||
for eos_token_id in eos_token_ids:
|
||||
unfinished_sents.mul_(tokens_to_add.ne(eos_token_id).long())
|
||||
cur_len = cur_len + 1
|
||||
|
||||
# stop when there is a </s> in each sentence, or if we exceed the maximul length
|
||||
if unfinished_sents.max() == 0:
|
||||
break
|
||||
|
||||
# add eos_token_ids to unfinished sentences
|
||||
if cur_len == max_length:
|
||||
# add the first eos_token_ids to unfinished sentences <= TODO should we do that?
|
||||
if cur_len == max_length and eos_token_ids is not None:
|
||||
input_ids[:, -1].masked_fill_(unfinished_sents.to(dtype=torch.bool), eos_token_ids[0])
|
||||
|
||||
return input_ids
|
||||
@@ -942,7 +949,7 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
|
||||
# if we are done with this sentence
|
||||
done[batch_ex] = done[batch_ex] or generated_hyps[batch_ex].is_done(next_scores[batch_ex].max().item())
|
||||
if done[batch_ex]:
|
||||
if done[batch_ex] and pad_token_id is not None:
|
||||
next_batch_beam.extend([(0, pad_token_id, 0)] * num_beams) # pad the batch
|
||||
continue
|
||||
|
||||
@@ -957,7 +964,7 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
word_id = idx % vocab_size
|
||||
|
||||
# end of sentence, or next word
|
||||
if word_id.item() in eos_token_ids or cur_len + 1 == max_length:
|
||||
if (eos_token_ids is not None and word_id.item() in eos_token_ids) or cur_len + 1 == max_length:
|
||||
generated_hyps[batch_ex].add(
|
||||
input_ids[batch_ex * num_beams + beam_id, :cur_len].clone(), score.item()
|
||||
)
|
||||
@@ -970,7 +977,7 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
|
||||
# update next beam content
|
||||
assert len(next_sent_beam) == 0 if cur_len + 1 == max_length else num_beams
|
||||
if len(next_sent_beam) == 0:
|
||||
if len(next_sent_beam) == 0 and pad_token_id is not None:
|
||||
next_sent_beam = [(0, pad_token_id, 0)] * num_beams # pad the batch
|
||||
next_batch_beam.extend(next_sent_beam)
|
||||
assert len(next_batch_beam) == num_beams * (batch_ex + 1)
|
||||
@@ -1005,15 +1012,6 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
if all(done):
|
||||
break
|
||||
|
||||
# visualize hypotheses
|
||||
# print([len(x) for x in generated_hyps], cur_len)
|
||||
# globals().update( locals() );
|
||||
# !import code; code.interact(local=vars())
|
||||
# for ii in range(batch_size):
|
||||
# for ss, ww in sorted(generated_hyps[ii].hyp, key=lambda x: x[0], reverse=True):
|
||||
# print("%.3f " % ss + " ".join(self.dico[x] for x in ww.tolist()))
|
||||
# print("")
|
||||
|
||||
# select the best hypotheses
|
||||
tgt_len = input_ids.new(batch_size)
|
||||
best = []
|
||||
@@ -1024,10 +1022,11 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
|
||||
best.append(best_hyp)
|
||||
|
||||
# generate target batch
|
||||
decoded = input_ids.new(batch_size, tgt_len.max().item()).fill_(pad_token_id)
|
||||
decoded = input_ids.new(batch_size, tgt_len.max().item()).fill_(pad_token_id if pad_token_id is not None else -1)
|
||||
for i, hypo in enumerate(best):
|
||||
decoded[i, : tgt_len[i] - 1] = hypo
|
||||
decoded[i, tgt_len[i] - 1] = eos_token_ids[0]
|
||||
if eos_token_ids is not None:
|
||||
decoded[i, tgt_len[i] - 1] = eos_token_ids[0]
|
||||
|
||||
return decoded
|
||||
|
||||
|
||||
Reference in New Issue
Block a user