torchmetrics represent uncertainty

Viewed 189

I am using torchmetrics to calculate metrics such as F1 score, Recall, Precision and Accuracy in multilabel classification setting. With random initiliazed weights the softmax output (i.e. prediction) might look like this with a batch size of 8:

import torch
y_pred = torch.tensor([[0.1944, 0.1931, 0.2184, 0.1968, 0.1973],
                       [0.2182, 0.1932, 0.1945, 0.1973, 0.1968],
                       [0.2182, 0.1932, 0.1944, 0.1973, 0.1969],
                       [0.2182, 0.1931, 0.1945, 0.1973, 0.1968],
                       [0.2184, 0.1931, 0.1944, 0.1973, 0.1968],
                       [0.2181, 0.1932, 0.1941, 0.1970, 0.1976],
                       [0.2183, 0.1932, 0.1944, 0.1974, 0.1967],
                       [0.2182, 0.1931, 0.1945, 0.1973, 0.1968]])

With the correct labels (one-hot encoded):

y_true = torch.tensor([[0, 0, 1, 0, 1],
                       [0, 1, 0, 0, 1],
                       [0, 1, 0, 0, 1],
                       [0, 0, 1, 1, 0],
                       [0, 0, 1, 1, 0],
                       [0, 1, 0, 1, 0],
                       [0, 1, 0, 1, 0],
                       [0, 0, 1, 0, 1]])

And I can calculate the metrics by taking argmax:

import torchmetrics
torchmetrics.functional.f1_score(y_pred.argmax(-1), y_true.argmax(-1))

output:

tensor(0.1250)

The first prediction happens to be correct while the rest are wrong. However, none of the predictive probabilities are above 0.3, which means that the model is generally uncertain about the predictions. I would like to encode this and say that the f1 score should be 0.0 because none of the predictive probabilities are above a 0.3 threshold.


Is this possible with torchmetrics or sklearn library?

Is this common practice?

1 Answers

You need to threshold you predictions before passing them to your torchmetrics

        t0, t1, mask_gt = batch
        mask_pred = self.forward(t0, t1)
        
        loss = self.criterion(mask_pred.squeeze().float(), mask_gt.squeeze().float())

        mask_pred = torch.sigmoid(mask_pred).squeeze()
        mask_pred = torch.where(mask_pred > 0.5, 1, 0)

        # integers to comply with metrics input type
        mask_pred = mask_pred.long()
        mask_gt = mask_gt.long()

        f1_score = self.f1(mask_pred, mask_gt)
        precision = self.precision_(mask_pred, mask_gt)
        recall = self.recall(mask_pred, mask_gt)
        jaccard = self.jaccard(mask_pred, mask_gt)

The defined torchmetrics

        self.f1 = F1Score(num_classes=2, average='macro', mdmc_average='samplewise')
        self.recall = Recall(num_classes=2, average='macro', mdmc_average='samplewise')
        self.precision_ = Precision(num_classes=2, average='macro', mdmc_average='samplewise')  # self.precision exists in torch.nn.Module. Hence '_' symbol
        self.jaccard = JaccardIndex(num_classes=2)
Related