Lightning-AI / Lightning-AI/torchmetrics

`BinaryAccuracy()` sometimes gives incorrect answers due to non-deterministic sigmoiding

Open
#1,604 4 comments 2 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

bug / fix help wanted v1.1.x
Dominant language
Python
Stars
2.5k
Forks
526
Avg merge
6d 11h
Merged PRs (30d)
5

Description

## 🐛 Bug

`torchmetrics.classification.BinaryAccuracy` will apply a sigmoid to some inputs but not others leading to incorrect behavior.

## Details
The current behavior of BinaryAccuracy() is to apply a sigmoid transformation if the inputs are outside of [0, 1] before binarizing

> If preds is a floating point tensor with values outside [0,1] range we consider the input to be logits and will auto apply sigmoid per element.

i.e.
y_hat = 1(sigmoid(z) >= threshold) if z outside [0, 1]
y_hat = 1(z >= threshold) if z inside [0, 1]

I assume z inside [0, 1] is checked for then entire batch (i.e. if one element of the batch is outside [0, 1] then we apply the sigmoid to everyone).

This will cause silent errors. In particular, if the user inputs logits then they expect the logits to always be sigmoided. However, it is totally possible for all of the logits to lie in [0, 1] for some batches in which case the input will **not** be sigmoided which will cause incorrect thresholding.

### To Reproduce

Here is a simple example. Support our network outputs logits.

```python
from torchmetrics.classification import BinaryAccuracy
from scipy.special import expit # expit = sigmoid
import numpy as np
import torch
```

This example should lead to a correct prediction
``` python
probability_thresh = 0.5
logits = np.array([0.49]) # network output
target = np.array([1])

# logits of 0.49 give a probability of 0.62 indicating class 1, the correct prediction
expit(logits)
array([0.62010643])
int(expit(logits) >= probability_thresh) == target
True
```

`BinaryAccuracy()` however thinks it's an incorrect prediction~
``` python
# torchmetrics, however, thinks we have the inccorect prediction because it does NOT sigmoid the logits
ba = BinaryAccuracy(threshold=probability_thresh)
ba.forward(preds=torch.tensor(logits), target=torch.tensor(target))
tensor(0.)
```

### Suggested Fix

I suggest adding an argument indicating whether or not the input predictions are sigmoided so the inputs are either always sigmoided or never sigmoided

Contributor guide

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Start at the BinaryAccuracy entry point and reproduce the provided logits example with threshold 0.5. Check how input ranges determine sigmoid application, then define and test consistent behavior for logits and probabilities; done means the example returns the expected accuracy without batch-dependent results.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
machine-learning
Issue type
Bug
Difficulty
4/5
Estimated time
3-5 days
Activity status
Quiet
Clarity
Mostly clear
Newbie friendliness
45/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.