Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
36 changes: 34 additions & 2 deletions ignite/metrics/metric_group.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,11 @@
from collections.abc import Callable, Sequence
from collections.abc import Callable, Mapping, Sequence
from typing import Any

import torch

from ignite.engine import Engine
from ignite.metrics import Metric
from ignite.metrics.metric import _is_list_of_tensors_or_numbers, _to_batched_tensor


class MetricGroup(Metric):
Expand Down Expand Up @@ -57,7 +59,37 @@ def reset(self) -> None:
for m in self.metrics.values():
m.reset()

def update(self, output: Sequence[torch.Tensor]) -> None:
def iteration_completed(self, engine: Engine) -> None:
# Overridden because, unlike a "leaf" metric, a MetricGroup does not itself consume a
# ``(y_pred, y)``-shaped output: each metric in the group applies its own
# ``output_transform`` in ``update`` to pull whatever it needs out of the group's
# (transformed) output. So, unlike ``Metric.iteration_completed``, a mapping output is
# passed straight through to ``update`` rather than being validated/unpacked against
# ``required_output_keys``, which only makes sense for a single metric's ``(y_pred, y)``.
output = self._output_transform(engine.state.output)
if isinstance(output, Mapping):
self.update(output)
return

if (
(not self._skip_unrolling)
and isinstance(output, Sequence)
and all(_is_list_of_tensors_or_numbers(o) for o in output)
):
if not (len(output) == 2 and len(output[0]) == len(output[1])):
raise ValueError(
f"Output should have 2 items of the same length, "
f"got {len(output)} and {len(output[0])}, {len(output[1])}"
)
for o1, o2 in zip(output[0], output[1]):
# o1 and o2 are list of tensors or numbers
tensor_o1 = _to_batched_tensor(o1)
tensor_o2 = _to_batched_tensor(o2, device=tensor_o1.device)
self.update((tensor_o1, tensor_o2))
else:
self.update(output)

def update(self, output: Sequence[torch.Tensor] | Mapping[Any, Any]) -> None:
for m in self.metrics.values():
m.update(m._output_transform(output))

Expand Down
33 changes: 32 additions & 1 deletion tests/ignite/metrics/test_metric_group.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@

from ignite import distributed as idist
from ignite.engine import Engine
from ignite.metrics import Accuracy, MetricGroup, Precision
from ignite.metrics import Accuracy, Loss, MetricGroup, Precision

torch.manual_seed(41)

Expand Down Expand Up @@ -48,6 +48,37 @@ def drop_first(output):
assert accuracy.state_dict() == group.metrics["accuracy"].state_dict()


def test_mapping_output_with_custom_keys():
# Regression test for https://github.com/pytorch/ignite/issues/3806 :
# a MetricGroup should not require the engine's output mapping to contain
# ('y_pred', 'y'); each metric in the group applies its own output_transform
# to pull what it needs from the mapping, same as attaching it directly.
def step(engine, batch):
return {
"outputs_1": (torch.rand(4, 3), torch.rand(4, 3)),
"masks": (torch.rand(4, 3), torch.rand(4, 3)),
}

loss_fn = torch.nn.MSELoss()

direct_engine = Engine(step)
direct_loss = Loss(loss_fn, output_transform=lambda o: o["outputs_1"])
direct_loss.attach(direct_engine, "loss")

group_engine = Engine(step)
group_loss = Loss(loss_fn, output_transform=lambda o: o["outputs_1"])
group = MetricGroup({"loss": group_loss})
group.attach(group_engine, "metrics")

torch.manual_seed(0)
direct_engine.run([0])
torch.manual_seed(0)
group_engine.run([0])

assert group_engine.state.metrics["metrics"] == {"loss": direct_engine.state.metrics["loss"]}
assert direct_loss.state_dict() == group_loss.state_dict()


def test_compute():
precision = Precision()
accuracy = Accuracy()
Expand Down