Addressing @mfuntowicz's comments

This commit is contained in:
Lysandre
2020-06-01 20:34:12 -04:00
parent ce1e6c3d74
commit 50fd3479f3
@@ -371,13 +371,10 @@ class CombinedModel(nn.Module):
# Identify the number of tokens to be masked, which should be: 1 < num < max_predictions per seq.
# It is set to be: n_tokens * mask_probability, but is truncated if it goes beyond bounds.
number_of_tokens_to_be_masked = torch.max(
torch.tensor(1),
torch.min(
torch.tensor(self.max_predictions_per_sequence, dtype=torch.long),
torch.tensor(int(total_number_of_tokens * self.mask_probability), dtype=torch.long),
),
)
number_of_tokens_to_be_masked = torch.min(
torch.tensor(self.max_predictions_per_sequence, dtype=torch.long),
torch.tensor(int(total_number_of_tokens * self.mask_probability), dtype=torch.long),
).clamp(1)
# The probability of each token being masked
sample_prob = proposal_distribution * inputs_which_can_be_masked
@@ -395,8 +392,8 @@ class CombinedModel(nn.Module):
def gather_positions(sequence, positions):
batch_size, sequence_length, dimension = sequence.shape
position_shift = (sequence_length * torch.arange(batch_size, device=sequence.device)).unsqueeze(-1)
flat_positions = torch.reshape(positions + position_shift, [-1]).long()
flat_sequence = torch.reshape(sequence, [batch_size * sequence_length, dimension])
flat_positions = torch.view(positions + position_shift, [-1]).long()
flat_sequence = torch.view(sequence, [batch_size * sequence_length, dimension])
gathered = flat_sequence.index_select(0, flat_positions)
return torch.reshape(gathered, [batch_size, -1, dimension])
@@ -417,9 +414,10 @@ class CombinedModel(nn.Module):
masked_lm_inputs = input_ids.clone()
# Of the evaluated tokens, 15% of those will keep their original tokens
replace_with_mask_positions = masked_lm_positions * (
torch.rand(masked_lm_positions.shape, device=masked_lm_positions.device) < (1 - self.mask_probability)
)
replace_with_mask_positions = torch.rand_like(
masked_lm_positions,
device=masked_lm_positions.device
) < (1 - self.mask_probability)
# Scatter the masks at the masked positions
masked_lm_inputs.scatter_(-1, replace_with_mask_positions, masked_tokens)
@@ -454,7 +452,7 @@ class CombinedModel(nn.Module):
labels=discriminator_labels,
)[:2]
discriminator_predictions = torch.round((torch.sign(discriminator_output) + 1) / 2)
discriminator_predictions = torch.round((torch.sign(discriminator_output) + 1.0) * 0.5)
total_loss = (self.discriminator_weight * discriminator_loss) + (self.generator_weight * generator_loss)