fix beam search bug in tf as well (#4745)

This commit is contained in:
Patrick von Platen
2020-06-03 12:53:23 -04:00
committed by GitHub
parent 1b5820a565
commit ed4df85572
2 changed files with 2 additions and 2 deletions
+1 -1
View File
@@ -1218,7 +1218,7 @@ class TFPreTrainedModel(tf.keras.Model, TFModelUtilsMixin):
continue
# test that beam scores match previously calculated scores if not eos and batch_idx not done
if eos_token_id is not None and all(
(token_id % vocab_size).numpy().item() is not eos_token_id for token_id in next_tokens[batch_idx]
(token_id % vocab_size).numpy().item() != eos_token_id for token_id in next_tokens[batch_idx]
):
assert tf.reduce_all(
next_scores[batch_idx, :num_beams] == tf.reshape(beam_scores, (batch_size, num_beams))[batch_idx]
+1 -1
View File
@@ -1528,7 +1528,7 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin):
# test that beam scores match previously calculated scores if not eos and batch_idx not done
if eos_token_id is not None and all(
(token_id % vocab_size).item() is not eos_token_id for token_id in next_tokens[batch_idx]
(token_id % vocab_size).item() != eos_token_id for token_id in next_tokens[batch_idx]
):
assert torch.all(
next_scores[batch_idx, :num_beams] == beam_scores.view(batch_size, num_beams)[batch_idx]