|
| 1 | +import pytest |
| 2 | +import torch |
| 3 | + |
| 4 | +from ignite.metrics.utils import get_sequence_transform |
| 5 | + |
| 6 | +def test_get_sequence_transform(): |
| 7 | + # test (N, L, C) |
| 8 | + y_pred = torch.tensor( |
| 9 | + [ |
| 10 | + [[0.1, 0.9], [0.8, 0.2], [0.3, 0.7], [0.5, 0.5]], |
| 11 | + [[0.9, 0.1], [0.2, 0.8], [0.4, 0.6], [0.5, 0.5]], |
| 12 | + ] |
| 13 | + ) # shape: (2, 4, 2) |
| 14 | + y = torch.tensor([[1, 0, 1, -1], [0, 1, 0, -1]]) # shape: (2, 4) |
| 15 | + |
| 16 | + transform = get_sequence_transform(ignore_index=-1) |
| 17 | + y_pred_t, y_t = transform((y_pred, y)) |
| 18 | + |
| 19 | + assert y_pred_t.shape == (6, 2) |
| 20 | + assert y_t.shape == (6,) |
| 21 | + assert y_t.tolist() == [1, 0, 1, 0, 1, 0] |
| 22 | + assert y_pred_t[:, 1].tolist() == pytest.approx([0.9, 0.2, 0.7, 0.1, 0.8, 0.6]) |
| 23 | + |
| 24 | + # test (N, C, L) |
| 25 | + y_pred_ncl = y_pred.transpose(1, 2).contiguous() # (2, 2, 4) |
| 26 | + y_pred_t2, y_t2 = transform((y_pred_ncl, y)) |
| 27 | + assert y_pred_t2.shape == (6, 2) |
| 28 | + assert torch.all(y_pred_t2 == y_pred_t) |
| 29 | + assert torch.all(y_t2 == y_t) |
| 30 | + |
| 31 | + # test binary (N, L) |
| 32 | + y_pred_bin = torch.tensor([[1, 0, 1, 1], [0, 1, 0, 0]]) |
| 33 | + y_bin = torch.tensor([[1, 0, 1, 2], [0, 1, 0, 2]]) |
| 34 | + transform_bin = get_sequence_transform(ignore_index=2) |
| 35 | + y_pred_bin_t, y_bin_t = transform_bin((y_pred_bin, y_bin)) |
| 36 | + |
| 37 | + assert y_pred_bin_t.shape == (6,) |
| 38 | + assert y_bin_t.shape == (6,) |
| 39 | + assert y_bin_t.tolist() == [1, 0, 1, 0, 1, 0] |
| 40 | + assert y_pred_bin_t.tolist() == [1, 0, 1, 0, 1, 0] |
| 41 | + |
| 42 | + # test without padding |
| 43 | + transform_nopad = get_sequence_transform() |
| 44 | + y_pred_nopad, y_nopad = transform_nopad((y_pred_bin, y_bin)) |
| 45 | + assert y_pred_nopad.shape == (8,) |
| 46 | + assert y_nopad.shape == (8,) |
| 47 | + |
| 48 | + # test multiple ignore_index values |
| 49 | + y_bin = torch.tensor([[1, -1, 1, 2], [0, 1, -1, 2]]) |
| 50 | + transform_multi = get_sequence_transform(ignore_index=[-1, 2]) |
| 51 | + y_pred_multi_t, y_multi_t = transform_multi((y_pred_bin, y_bin)) |
| 52 | + assert y_pred_multi_t.shape == (4,) |
| 53 | + assert y_multi_t.shape == (4,) |
| 54 | + assert y_multi_t.tolist() == [1, 1, 0, 1] |
| 55 | + |
| 56 | + # test bad shapes |
| 57 | + y_bad = torch.tensor([1, 0, 1]) |
| 58 | + with pytest.raises(ValueError, match="must be 3D and 2D arrays, or both 2D arrays"): |
| 59 | + transform((y_pred_bin, y_bad)) |
| 60 | + |
| 61 | + y_pred_bad = torch.tensor([[[1], [2]], [[3], [4]]]) |
| 62 | + y_bad = torch.tensor([[1, 2, 3], [4, 5, 6]]) |
| 63 | + with pytest.raises(ValueError, match="incompatible sequence shapes"): |
| 64 | + transform((y_pred_bad, y_bad)) |
0 commit comments