This commit is contained in:
Lysandre
2020-06-01 20:23:30 -04:00
parent ea3e2cebd8
commit f842f7c764
2 changed files with 47 additions and 47 deletions
@@ -621,47 +621,52 @@ def main():
model = CombinedModel(discriminator, generator, tokenizer, training_args, data_args)
def compute_accuracy(evaluation_predictions: EvalPrediction) -> Dict[str, int]:
def compute_metrics(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"]
true_positives = np.logical_and(np.equal(discriminator_predictions, discriminator_labels), np.equal(discriminator_labels, 1)).sum().astype(float)
false_negatives = np.logical_and(np.not_equal(discriminator_predictions, discriminator_labels), np.equal(discriminator_labels, 1)).sum().astype(float)
false_positives = np.logical_and(np.not_equal(discriminator_predictions, discriminator_labels), np.equal(discriminator_labels, 0)).sum().astype(float)
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
discriminator_precision = true_positives / (true_positives + false_positives)
discriminator_recall = true_positives / (true_positives + false_negatives)
return {"generator_accuracy": generator_accuracy, "discriminator_accuracy": discriminator_accuracy}
return {
"generator_accuracy": generator_accuracy,
"discriminator_accuracy": discriminator_accuracy,
"discriminator_precision": discriminator_precision,
"discriminator_recall": discriminator_recall
}
def manage_evaluation_predictions(model_outputs: Tuple[torch.Tensor]) -> Tuple[
Dict[str, torch.Tensor],
Callable[[Dict[str, np.ndarray]], EvalPrediction]
]:
def manage_evaluation_predictions(model_outputs: Tuple[torch.Tensor]) -> Dict[str, torch.Tensor]:
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 = {
return {
"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
def eval_prediction_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)
# Initialize our Trainer
trainer = Trainer(
@@ -670,8 +675,9 @@ def main():
data_collator=data_collator,
train_dataset=train_dataset,
eval_dataset=eval_dataset,
compute_metrics=compute_accuracy,
compute_metrics=compute_metrics,
manage_evaluation_predictions=manage_evaluation_predictions,
eval_prediction_mapping=eval_prediction_mapping,
prediction_loss_only=True,
)
+20 -26
View File
@@ -162,10 +162,8 @@ 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
manage_evaluation_predictions: Optional[Callable[[Tuple[torch.Tensor]], Dict[str, torch.Tensor]]] = None
eval_prediction_mapping: Optional[Callable[[Dict[str, np.ndarray]], EvalPrediction]]
prediction_loss_only: bool
tb_writer: Optional["SummaryWriter"] = None
optimizers: Tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR] = None
@@ -180,7 +178,8 @@ 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,
manage_evaluation_predictions: Optional[Callable[[Tuple[torch.Tensor]], Dict[str, torch.Tensor]]] = None,
eval_prediction_mapping: Optional[Callable[[Dict[str, np.ndarray]], EvalPrediction]] = None,
prediction_loss_only=False,
tb_writer: Optional["SummaryWriter"] = None,
optimizers: Tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR] = None,
@@ -203,6 +202,7 @@ class Trainer:
self.eval_dataset = eval_dataset
self.compute_metrics = compute_metrics
self.manage_evaluation_predictions = manage_evaluation_predictions
self.eval_prediction_mapping = eval_prediction_mapping
self.prediction_loss_only = prediction_loss_only
self.optimizers = optimizers
if tb_writer is not None:
@@ -751,7 +751,7 @@ class Trainer:
logger.info(" Num examples = %d", self.num_examples(dataloader))
logger.info(" Batch size = %d", batch_size)
eval_losses: List[float] = []
evaluation_values: Union[Dict[str, Union[torch.Tensor, None]], None] = None
evaluation_values: Dict[str, torch.Tensor] = {}
model.eval()
if is_tpu_available():
@@ -770,7 +770,6 @@ class Trainer:
outputs = model(**inputs)
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()]
@@ -778,12 +777,12 @@ class Trainer:
logits = outputs[0]
if not prediction_loss_only:
if evaluation_values["preds"] is None:
if "preds" not in evaluation_values["preds"]:
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:
if "label_ids" not in evaluation_values["label_ids"]:
evaluation_values["label_ids"] = inputs["labels"].detach()
else:
evaluation_values["label_ids"] = torch.cat(
@@ -791,25 +790,18 @@ class Trainer:
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
if not len(evaluation_values):
evaluation_values = self.manage_evaluation_predictions(outputs)
else:
for key, value in self.manage_evaluation_predictions(outputs)[0].items():
for key, value in self.manage_evaluation_predictions(outputs).items():
evaluation_values[key] = torch.cat((evaluation_values[key], value), dim=0)
evaluation_total_steps += 1
if evaluation_total_steps > self.args.max_eval_steps:
evaluation_iterator.close()
break
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)
@@ -823,25 +815,27 @@ class Trainer:
evaluation_values[key] = xm.mesh_reduce(key, value, torch.cat)
# Finally, turn the aggregated tensors into numpy arrays.
evaluation_values_numpy: Dict[str, np.ndarray] = {}
for key, value in evaluation_values.items():
if value is not None:
evaluation_values[key] = value.cpu().numpy()
evaluation_values_numpy[key] = value.cpu().numpy()
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)
if self.manage_evaluation_predictions is not None and self.eval_prediction_mapping is not None:
eval_predictions = self.eval_prediction_mapping(evaluation_values_numpy)
metrics = self.compute_metrics(eval_predictions)
elif evaluation_values["preds"] is not None and evaluation_values["label_ids"] is not None:
elif evaluation_values_numpy["preds"] is not None and evaluation_values_numpy["label_ids"] is not None:
eval_predictions = EvalPrediction(
predictions=evaluation_values["preds"],
label_ids=evaluation_values["label_ids"]
predictions=evaluation_values_numpy["preds"],
label_ids=evaluation_values_numpy["label_ids"]
)
metrics = self.compute_metrics(eval_predictions)
else:
metrics = {}
eval_predictions = EvalPrediction(np.array(0), np.array(0))
else:
metrics = {}
eval_predictions = EvalPrediction(np.array(0), np.array(0))
if len(eval_losses) > 0:
metrics["eval_loss"] = np.mean(eval_losses)