diff --git a/ignite/metrics/accuracy.py b/ignite/metrics/accuracy.py index 2d44a37d4afc..d8103ff0b993 100644 --- a/ignite/metrics/accuracy.py +++ b/ignite/metrics/accuracy.py @@ -167,6 +167,8 @@ class Accuracy(_BaseClassification): :class:`~ignite.engine.engine.Engine`'s ``process_function``'s output into the form expected by the metric. This can be useful if, for example, you have a multi-output model and you want to compute the metric with respect to one of the outputs. + average: if ``False`` and ``is_multilabel=True``, returns per-label accuracy as a tensor. + By default, ``None`` preserves the existing scalar behavior. is_multilabel: flag to use in multilabel case. By default, False. device: specifies which device updates are accumulated on. Setting the metric's device to be the same as your ``update`` arguments ensures the ``update`` method is non-blocking. By @@ -246,6 +248,29 @@ class Accuracy(_BaseClassification): 0.2 + Multilabel case with per-label accuracy (``average=False``) + + .. testcode:: 5 + + metric = Accuracy(is_multilabel=True, average=False) + metric.attach(default_evaluator, "accuracy") + y_true = torch.tensor([ + [1, 0, 1], + [0, 1, 0], + [1, 1, 0], + ]) + y_pred = torch.tensor([ + [1, 0, 0], + [0, 1, 0], + [1, 0, 0], + ]) + state = default_evaluator.run([[y_pred, y_true]]) + print(state.metrics["accuracy"]) + + .. testoutput:: 5 + + tensor([1.0000, 0.6667, 0.3333], dtype=torch.float64) + In binary and multilabel cases, the elements of `y` and `y_pred` should have 0 or 1 values. Thresholding of predictions can be done as below: @@ -279,7 +304,14 @@ def __init__( is_multilabel: bool = False, device: str | torch.device = torch.device("cpu"), skip_unrolling: bool = False, + average: bool | str | None = None, ): + if average is True: + self._average = "macro" + elif average is False or average is None: + self._average = average + else: + raise ValueError("Argument average should be None or a boolean.") super().__init__( output_transform=output_transform, is_multilabel=is_multilabel, device=device, skip_unrolling=skip_unrolling ) @@ -307,15 +339,28 @@ def update(self, output: Sequence[torch.Tensor]) -> None: last_dim = y_pred.ndimension() y_pred = torch.transpose(y_pred, 1, last_dim - 1).reshape(-1, num_classes) y = torch.transpose(y, 1, last_dim - 1).reshape(-1, num_classes) - correct = torch.all(y == y_pred.type_as(y), dim=-1) + if self._average is False: + correct = (y == y_pred.type_as(y)).to(dtype=torch.float64, device=self._device) + else: + correct = torch.all(y == y_pred.type_as(y), dim=-1) else: raise ValueError(f"Unexpected type: {self._type}") - self._num_correct += torch.sum(correct).to(self._device) + if self._type == "multilabel" and self._average is False: + self._num_correct = self._num_correct.to(self._device) + if self._num_correct.numel() == 1: + self._num_correct = torch.zeros(correct.size(1), device=self._device, dtype=torch.float64) + self._num_correct += correct.sum(dim=0).to(self._device) + else: + self._num_correct += torch.sum(correct).to(self._device) self._num_examples += correct.shape[0] @sync_all_reduce("_num_examples", "_num_correct") - def compute(self) -> float: + def compute(self) -> float | torch.Tensor: if self._num_examples == 0: raise NotComputableError("Accuracy must have at least one example before it can be computed.") + + if self._type == "multilabel" and self._average is False: + return self._num_correct / self._num_examples + return self._num_correct.item() / self._num_examples diff --git a/tests/ignite/metrics/test_accuracy.py b/tests/ignite/metrics/test_accuracy.py index cc87b72be36f..cd521a06b53d 100644 --- a/tests/ignite/metrics/test_accuracy.py +++ b/tests/ignite/metrics/test_accuracy.py @@ -176,6 +176,21 @@ def test_multilabel_input(n_times, available_device, test_data_multilabel): assert accuracy_score(np_y, np_y_pred) == pytest.approx(acc.compute()) +def test_multilabel_average_false(): + y_pred = torch.tensor([[1, 0, 0], [0, 1, 0], [1, 0, 0]], dtype=torch.long) + y = torch.tensor([[1, 0, 1], [0, 1, 0], [1, 1, 0]], dtype=torch.long) + + acc = Accuracy(is_multilabel=True, average=False) + acc.update((y_pred, y)) + + expected = torch.tensor([1.0, 2.0 / 3.0, 1.0 / 3.0], dtype=torch.float64) + assert torch.allclose(acc.compute(), expected) + + acc = Accuracy(is_multilabel=True) + acc.update((y_pred, y)) + assert isinstance(acc.compute(), float) + + def test_incorrect_type(): acc = Accuracy()