Evaluation

This commit is contained in:
Lysandre
2020-06-01 20:23:30 -04:00
parent a6f5989bd9
commit ea3e2cebd8
5 changed files with 138 additions and 87 deletions
@@ -29,12 +29,13 @@ from itertools import chain
from pathlib import Path
from dataclasses import dataclass, field
from typing import Optional, Union
from typing import Optional, Union, Tuple, Dict, Callable
import torch
import torch.nn as nn
from torch.utils.data import IterableDataset
from tqdm import tqdm
import numpy as np
from transformers import (
AutoConfig,
@@ -46,7 +47,7 @@ from transformers import (
PreTrainedTokenizer,
TextDataset,
Trainer,
set_seed,
set_seed, EvalPrediction,
)
from transformers.modeling_utils import PreTrainedModel
@@ -151,7 +152,7 @@ class DataTrainingArguments:
max_predictions_per_sequence: int = field(
# Original implementation has a default of :(mask_probability + 0.005) * max_sequence_length
default=int((0.15 + 0.005) * 512),
default=int((0.15 + 0.005) * 128),
metadata={"help": "Maximum tokens that will be masked in a sequence."}
)
@@ -381,31 +382,13 @@ class CombinedModel(nn.Module):
sample_prob /= torch.sum(sample_prob)
# Should be passed through a log function here
# Weight of each position: 1 the position will be masked, 0 the position won't be masked
masked_lm_weights = torch.tensor([0] * self.max_predictions_per_sequence, dtype=torch.bool, device=input_ids.device)
masked_lm_weights[:number_of_tokens_to_be_masked] = True
# Sample from the probabilities
masked_lm_positions = sample_prob.multinomial(self.max_predictions_per_sequence)
# Apply the weights to the positions
masked_lm_positions *= masked_lm_weights.long()
masked_lm_positions = sample_prob.multinomial(number_of_tokens_to_be_masked)
# Gather the IDs from the positions
masked_lm_ids = input_ids.gather(-1, masked_lm_positions)
# Apply weights to the IDs
masked_lm_ids *= masked_lm_weights.long()
replace_with_mask_positions = masked_lm_positions * (torch.rand(masked_lm_positions.shape, device=masked_lm_positions.device) < 0.85)
# Replace the input IDs with masks on given positions
masked_input_ids = input_ids.scatter(-1, replace_with_mask_positions, mask_token_id)
# Updates to index 0 should be ignored
masked_input_ids[..., 0] = input_ids[..., 0]
return masked_input_ids, masked_lm_positions
return masked_lm_ids, masked_lm_positions
@staticmethod
def gather_positions(sequence, positions):
@@ -416,26 +399,6 @@ class CombinedModel(nn.Module):
gathered = flat_sequence.index_select(0, flat_positions)
return torch.reshape(gathered, [batch_size, -1, dimension])
@staticmethod
def compute_metrics(
input_ids: torch.Tensor,
masked_lm_ids: torch.Tensor,
masked_lm_preds: torch.Tensor,
input_mask: torch.Tensor,
discriminator_labels: torch.Tensor,
discriminator_predictions: torch.Tensor,
sampled_tokids: torch.Tensor
):
input_ids
masked_lm_accuracy = masked_lm_ids.eq(masked_lm_preds)
sampled_masked_lm_accuracy = masked_lm_ids.eq(sampled_tokids)
discriminator_accuracy = discriminator_labels.eq(discriminator_predictions)
input_mask
def forward(
self,
input_ids=None,
@@ -457,10 +420,15 @@ class CombinedModel(nn.Module):
masked_lm_labels = torch.full_like(input_ids, -100)
masked_lm_labels.scatter_(-1, masked_lm_positions, masked_input_ids)
# mask the inputs with the mask token
# Create a tensor filled with masks
masked_tokens = torch.full_like(masked_input_ids, self.tokenizer.mask_token_id)
masked_lm_inputs = input_ids.clone()
masked_lm_inputs.scatter_(-1, masked_lm_positions, masked_tokens)
# 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) < 0.85)
# Scatter the masks at the masked positions
masked_lm_inputs.scatter_(-1, replace_with_mask_positions, masked_tokens)
masked_lm_inputs[..., 0] = 101
generator_loss, generator_output = self.generator(
@@ -480,17 +448,20 @@ class CombinedModel(nn.Module):
# create a tensor containing the predicted tokens
fake_tokens = input_ids.scatter(-1, masked_lm_positions, fake_argmaxes)
fake_tokens[:, 0] = input_ids[:, 0]
discriminator_labels = torch.tensor(labels != fake_tokens, dtype=torch.uint8, device=input_ids.device)
discriminator_labels = (labels != fake_tokens).int()
discriminator_loss, discriminator_output = self.discriminator(
fake_tokens, attention_mask, token_type_ids, position_ids, head_mask, position_ids, labels=discriminator_labels
)[:2]
discriminator_predictions = torch.round((torch.sign(discriminator_output) + 1) / 2).int().tolist()
discriminator_predictions = torch.round((torch.sign(discriminator_output) + 1) / 2)
total_loss = (self.discriminator_weight * discriminator_loss) + (self.generator_weight * generator_loss)
return (total_loss, (generator_output, discriminator_output), (fake_tokens, discriminator_predictions))
# generator_accuracy = masked_input_ids.eq(fake_argmaxes).sum().float() / fake_argmaxes.numel()
# discriminator_accuracy = discriminator_labels.eq(discriminator_predictions).sum().float() / attention_mask.sum()
return (total_loss, (generator_output, discriminator_output), (masked_input_ids, fake_argmaxes), (discriminator_labels, discriminator_predictions))
def save_pretrained(self, directory):
if self.config.xla_device:
@@ -650,6 +621,48 @@ def main():
model = CombinedModel(discriminator, generator, tokenizer, training_args, data_args)
def compute_accuracy(evaluation_predictions: EvalPrediction) -> Dict[str, int]:
predictions: Dict[str, np.ndarray] = evaluation_predictions.predictions
labels: Dict[str, np.ndarray] = evaluation_predictions.label_ids
generator_labels, generator_predictions = labels["generator"], predictions["generator"]
discriminator_labels, discriminator_predictions = labels["discriminator"], predictions["discriminator"]
generator_accuracy = np.equal(generator_labels, generator_predictions).sum().astype(float) / generator_predictions.size
discriminator_accuracy = np.equal(discriminator_labels, discriminator_predictions).sum().astype(float) / discriminator_labels.size
return {"generator_accuracy": generator_accuracy, "discriminator_accuracy": discriminator_accuracy}
def manage_evaluation_predictions(model_outputs: Tuple[torch.Tensor]) -> Tuple[
Dict[str, torch.Tensor],
Callable[[Dict[str, np.ndarray]], EvalPrediction]
]:
total_loss, models_output, generator_evaluation_values, discriminator_evaluation_values = model_outputs
generator_labels, generator_predictions = generator_evaluation_values
discriminator_labels, discriminator_predictions = discriminator_evaluation_values
label_dictionary = {
"generator_labels": generator_labels.detach(),
"generator_predictions": generator_predictions.detach(),
"discriminator_labels": discriminator_labels.detach(),
"discriminator_predictions": discriminator_predictions.detach()
}
def mapping(dictionary: Dict[str, np.ndarray]) -> EvalPrediction:
labels = {
"generator": dictionary["generator_labels"],
"discriminator": dictionary["discriminator_labels"]
}
predictions = {
"generator": dictionary["generator_predictions"],
"discriminator": dictionary["discriminator_predictions"]
}
return EvalPrediction(labels, predictions)
return label_dictionary, mapping
# Initialize our Trainer
trainer = Trainer(
model=model,
@@ -657,6 +670,8 @@ def main():
data_collator=data_collator,
train_dataset=train_dataset,
eval_dataset=eval_dataset,
compute_metrics=compute_accuracy,
manage_evaluation_predictions=manage_evaluation_predictions,
prediction_loss_only=True,
)
@@ -676,9 +691,7 @@ def main():
if training_args.do_eval and training_args.local_rank in [-1, 0]:
logger.info("*** Evaluate ***")
eval_output = trainer.evaluate()
result = {"loss": eval_output["eval_loss"]}
result = trainer.evaluate()
output_eval_file = os.path.join(training_args.output_dir, "eval_results_lm.txt")
with open(output_eval_file, "w") as writer:
+1 -1
View File
@@ -132,7 +132,7 @@ class ElectraDiscriminatorPredictions(nn.Module):
def forward(self, discriminator_hidden_states, attention_mask):
hidden_states = self.dense(discriminator_hidden_states)
hidden_states = get_activation(self.config.hidden_act)(hidden_states)
logits = self.dense_prediction(hidden_states).squeeze()
logits = self.dense_prediction(hidden_states).squeeze(-1)
return logits
+67 -32
View File
@@ -7,7 +7,7 @@ import re
import shutil
from contextlib import contextmanager
from pathlib import Path
from typing import Callable, Dict, List, Optional, Tuple
from typing import Callable, Dict, List, Optional, Tuple, Union
import numpy as np
import torch
@@ -162,6 +162,10 @@ class Trainer:
train_dataset: Optional[Dataset]
eval_dataset: Optional[Dataset]
compute_metrics: Optional[Callable[[EvalPrediction], Dict]] = None
manage_evaluation_predictions: Optional[Callable[
[Tuple[torch.Tensor]],
Tuple[Dict[str, torch.Tensor], Callable[[Dict[str, np.ndarray]], EvalPrediction]]]
] = None
prediction_loss_only: bool
tb_writer: Optional["SummaryWriter"] = None
optimizers: Tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR] = None
@@ -176,6 +180,7 @@ class Trainer:
train_dataset: Optional[Dataset] = None,
eval_dataset: Optional[Dataset] = None,
compute_metrics: Optional[Callable[[EvalPrediction], Dict]] = None,
manage_evaluation_predictions: Optional[Callable[[Tuple[torch.Tensor]], EvalPrediction]] = None,
prediction_loss_only=False,
tb_writer: Optional["SummaryWriter"] = None,
optimizers: Tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR] = None,
@@ -197,6 +202,7 @@ class Trainer:
self.train_dataset = train_dataset
self.eval_dataset = eval_dataset
self.compute_metrics = compute_metrics
self.manage_evaluation_predictions = manage_evaluation_predictions
self.prediction_loss_only = prediction_loss_only
self.optimizers = optimizers
if tb_writer is not None:
@@ -745,8 +751,7 @@ class Trainer:
logger.info(" Num examples = %d", self.num_examples(dataloader))
logger.info(" Batch size = %d", batch_size)
eval_losses: List[float] = []
preds: torch.Tensor = None
label_ids: torch.Tensor = None
evaluation_values: Union[Dict[str, Union[torch.Tensor, None]], None] = None
model.eval()
if is_tpu_available():
@@ -763,22 +768,40 @@ class Trainer:
with torch.no_grad():
outputs = model(**inputs)
if has_labels:
step_eval_loss, logits = outputs[:2]
eval_losses += [step_eval_loss.mean().item()]
else:
logits = outputs[0]
if not prediction_loss_only:
if preds is None:
preds = logits.detach()
else:
preds = torch.cat((preds, logits.detach()), dim=0)
if inputs.get("labels") is not None:
if label_ids is None:
label_ids = inputs["labels"].detach()
if self.manage_evaluation_predictions is None:
evaluation_values = {"preds": None, "label_ids": None}
if has_labels:
step_eval_loss, logits = outputs[:2]
eval_losses += [step_eval_loss.mean().item()]
else:
label_ids = torch.cat((label_ids, inputs["labels"].detach()), dim=0)
logits = outputs[0]
if not prediction_loss_only:
if evaluation_values["preds"] is None:
evaluation_values["preds"] = logits.detach()
else:
evaluation_values["preds"] = torch.cat((evaluation_values["preds"], logits.detach()), dim=0)
if inputs.get("labels") is not None:
if evaluation_values["label_ids"] is None:
evaluation_values["label_ids"] = inputs["labels"].detach()
else:
evaluation_values["label_ids"] = torch.cat(
(evaluation_values["label_ids"], inputs["labels"].detach()),
dim=0
)
else:
if evaluation_values is None:
evaluation_items = self.manage_evaluation_predictions(outputs)
if len(evaluation_items) == 2:
evaluation_values, mapping_method = evaluation_items
else:
evaluation_values = evaluation_items[0]
mapping_method = None
else:
for key, value in self.manage_evaluation_predictions(outputs)[0].items():
evaluation_values[key] = torch.cat((evaluation_values[key], value), dim=0)
evaluation_total_steps += 1
@@ -790,25 +813,33 @@ class Trainer:
if self.args.local_rank != -1:
# In distributed mode, concatenate all results from all nodes:
total_examples = self.args.max_eval_steps if self.args.max_eval_steps > 0 else self.num_examples(dataloader)
if preds is not None:
preds = self.distributed_concat(preds, num_total_examples=total_examples)
if label_ids is not None:
label_ids = self.distributed_concat(label_ids, num_total_examples=total_examples)
for key, value in evaluation_values.items():
if value is not None:
evaluation_values[key] = self.distributed_concat(value, num_total_examples=total_examples)
elif is_tpu_available():
# tpu-comment: Get all predictions and labels from all worker shards of eval dataset
if preds is not None:
preds = xm.mesh_reduce("eval_preds", preds, torch.cat)
if label_ids is not None:
label_ids = xm.mesh_reduce("eval_label_ids", label_ids, torch.cat)
for key, value in evaluation_values.items():
if value is not None:
evaluation_values[key] = xm.mesh_reduce(key, value, torch.cat)
# Finally, turn the aggregated tensors into numpy arrays.
if preds is not None:
preds = preds.cpu().numpy()
if label_ids is not None:
label_ids = label_ids.cpu().numpy()
for key, value in evaluation_values.items():
if value is not None:
evaluation_values[key] = value.cpu().numpy()
if self.compute_metrics is not None and preds is not None and label_ids is not None:
metrics = self.compute_metrics(EvalPrediction(predictions=preds, label_ids=label_ids))
eval_predictions = EvalPrediction(None, None)
if self.compute_metrics is not None:
if self.manage_evaluation_predictions is not None and mapping_method is not None:
eval_predictions = mapping_method(evaluation_values)
metrics = self.compute_metrics(eval_predictions)
elif evaluation_values["preds"] is not None and evaluation_values["label_ids"] is not None:
eval_predictions = EvalPrediction(
predictions=evaluation_values["preds"],
label_ids=evaluation_values["label_ids"]
)
metrics = self.compute_metrics(eval_predictions)
else:
metrics = {}
else:
metrics = {}
if len(eval_losses) > 0:
@@ -819,7 +850,11 @@ class Trainer:
if not key.startswith("eval_"):
metrics[f"eval_{key}"] = metrics.pop(key)
return PredictionOutput(predictions=preds, label_ids=label_ids, metrics=metrics)
return PredictionOutput(
predictions=eval_predictions.predictions,
label_ids=eval_predictions.label_ids,
metrics=metrics
)
def distributed_concat(self, tensor: torch.Tensor, num_total_examples: int) -> torch.Tensor:
assert self.args.local_rank != -1
+3 -3
View File
@@ -1,4 +1,4 @@
from typing import Dict, NamedTuple, Optional
from typing import Dict, NamedTuple, Optional, Union
import numpy as np
@@ -9,8 +9,8 @@ class EvalPrediction(NamedTuple):
to compute metrics.
"""
predictions: np.ndarray
label_ids: np.ndarray
predictions: Union[np.ndarray, Dict[str, np.ndarray]]
label_ids: Union[np.ndarray, Dict[str, np.ndarray]]
class PredictionOutput(NamedTuple):
+3
View File
@@ -57,6 +57,9 @@ class TrainingArguments:
evaluate_during_training: bool = field(
default=False, metadata={"help": "Run evaluation during training at each logging step."},
)
evaluate_every: int = field(
default=-1, metadata={"help": "Will evaluate every n training steps."}
)
per_device_train_batch_size: int = field(
default=8, metadata={"help": "Batch size per GPU/TPU core/CPU for training."}