Flattens predictions in the batch
(probas, labels, ignore=None)
| 650 | |
| 651 | |
| 652 | def _flatten_probas(probas, labels, ignore=None): |
| 653 | """Flattens predictions in the batch""" |
| 654 | if probas.dim() == 3: |
| 655 | # assumes output of a sigmoid layer |
| 656 | B, H, W = probas.size() |
| 657 | probas = probas.view(B, 1, H, W) |
| 658 | |
| 659 | C = probas.size(1) |
| 660 | probas = torch.movedim(probas, 1, -1) # [B, C, Di, Dj, ...] -> [B, Di, Dj, ..., C] |
| 661 | probas = probas.contiguous().view(-1, C) # [P, C] |
| 662 | |
| 663 | labels = labels.view(-1) |
| 664 | if ignore is None: |
| 665 | return probas, labels |
| 666 | valid = labels != ignore |
| 667 | vprobas = probas[valid] |
| 668 | vlabels = labels[valid] |
| 669 | return vprobas, vlabels |
| 670 | |
| 671 | |
| 672 | def isnan(x): |