5. Metrics and scoring¶
metrics/my_metric.py → cfg.{train,val,tester}_metrics
cfg.val_metrics = {
"acc": {"cls_name": "metrics.accuracy",
"params": {"src_name": "logits", "tgt_name": ("targets", "labels")}},
}
Keys are namespaced {engine_type}/{key}, so this one reads valid/acc.
cfg.every_val = 1 # how often the validator runs, in epochs
cfg.every_test = 0 # 0, the default, builds no tester at all
An engine is built only if it will actually run, so every_test = 0 builds no tester during
training. gatle-ignite eval builds it anyway, since evaluating is the point of that run.
A test split has its own set of the three fields on this page: tester_metrics, tester_score_name
and tester_score_factor.
There are no built-in metrics, because a metric describes your task. For accuracy, copy the synthetic example's:
"""ignite's Accuracy reduces across ranks; a metric summing python floats is rank-local."""
from ignite.metrics import Accuracy
from gatle_ignite import get_value
def get_metric(engine_type, info, src_name="logits", tgt_name=("targets", "labels"), **kwargs):
def output_transform(output):
return get_value(output["y_pred"], src_name), get_value(output["target"], tgt_name)
return Accuracy(output_transform=output_transform, **kwargs)
Choosing the best checkpoint¶
import torch
from ignite.metrics import Metric
from ignite.metrics.metric import reinit__is_reduced, sync_all_reduce
from gatle_ignite import get_value
class TopKAccuracy(Metric):
def __init__(self, k=2, output_transform=lambda x: x, device="cpu"):
self.k = k
super().__init__(output_transform=output_transform, device=device)
@reinit__is_reduced
def reset(self):
# Tensors, not python floats: sync_all_reduce can only move tensors between ranks.
self._num_correct = torch.tensor(0, device=self._device)
self._num_examples = torch.tensor(0, device=self._device)
super().reset()
@reinit__is_reduced
def update(self, output):
y_pred, y = output
k = min(self.k, y_pred.shape[1])
topk = y_pred.topk(k, dim=1).indices
hits = (topk == y.unsqueeze(1)).any(dim=1).sum()
self._num_correct += hits.to(self._device)
self._num_examples += y.shape[0]
@sync_all_reduce("_num_correct", "_num_examples")
def compute(self):
if self._num_examples == 0:
raise ValueError("TopKAccuracy needs at least one example")
return self._num_correct.item() / self._num_examples.item()
def get_metric(engine_type, info, src_name="logits", tgt_name=("targets", "labels"), k=2, **kwargs):
n_classes = (info or {}).get("num_classes")
if n_classes is not None:
k = min(k, n_classes)
def output_transform(output):
return get_value(output["y_pred"], src_name), get_value(output["target"], tgt_name)
return TopKAccuracy(k=k, output_transform=output_transform, **kwargs)
info carries at least length, and output_transform selects what update() sees.
Rank-local metrics lie under DDP
A metric accumulating into a plain Python float computes a per-rank result, and each rank
silently reports only its own shard: right on one GPU, wrong on many. Accumulate into tensors
and decorate compute() with @sync_all_reduce(...) and reset()/update() with
@reinit__is_reduced, as examples/translation/metrics/edit_distance.py does. A metric that
is not a sum (an AUC, a distributional distance) gathers instead: idist.all_gather in compute() and no
sync_all_reduce, which would double-count on top of it.
score_factor = -1 for anything where lower is better
Checkpointing keeps the maximum, so a loss, a CER or a WER needs -1. With +1 the run
happily saves its worst model and reports success.
Score edit distance on decoded output, not teacher-forced logits: a model can teacher-force well and decode into garbage, and the logits throw away the only signal that would have noticed.
An engine you add yourself brings its own pair of fields, named from its key: EngineSpec.for_split("raw", ...)
returned from eval_specs() reads cfg.raw_metrics and cfg.every_raw, and reports under raw/.
A declared engine must set its cadence field or the run raises rather than silently skipping it.
If the config form cannot express what you need, override dict_metric_from_list and build metrics
imperatively. Call super(), or every config-declared metric is silently discarded; if
score_name named one of them, checkpointing then raises rather than scoring.
Templates: metric