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
20 changes: 18 additions & 2 deletions monai/metrics/confusion_matrix.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,7 +43,8 @@ class ConfusionMatrixMetric(CumulativeIterationMetric):
``"miss rate"``, ``"fall out"``, ``"false discovery rate"``, ``"false omission rate"``,
``"prevalence threshold"``, ``"threat score"``, ``"accuracy"``, ``"balanced accuracy"``,
``"f1 score"``, ``"matthews correlation coefficient"``, ``"fowlkes mallows index"``,
``"informedness"``, ``"markedness"``]
``"informedness"``, ``"markedness"``, ``"positive likelihood ratio"``,
``"negative likelihood ratio"``]
Some of the metrics have multiple aliases (as shown in the wikipedia page aforementioned),
and you can also input those names instead.
Except for input only one metric, multiple metrics are also supported via input a sequence of metric names, such as
Expand Down Expand Up @@ -185,7 +186,8 @@ def compute_confusion_matrix_metric(metric_name: str, confusion_matrix: torch.Te
``"miss rate"``, ``"fall out"``, ``"false discovery rate"``, ``"false omission rate"``,
``"prevalence threshold"``, ``"threat score"``, ``"accuracy"``, ``"balanced accuracy"``,
``"f1 score"``, ``"matthews correlation coefficient"``, ``"fowlkes mallows index"``,
``"informedness"``, ``"markedness"``]
``"informedness"``, ``"markedness"``, ``"positive likelihood ratio"``,
``"negative likelihood ratio"``]
Some of the metrics have multiple aliases (as shown in the wikipedia page aforementioned),
and you can also input those names instead.
confusion_matrix: Please see the doc string of the function ``get_confusion_matrix`` for more details.
Expand Down Expand Up @@ -263,6 +265,16 @@ def compute_confusion_matrix_metric(metric_name: str, confusion_matrix: torch.Te
npv = torch.where((tn + fn) > 0, tn / (tn + fn), nan_tensor)
numerator = ppv + npv - 1.0
denominator = 1.0
elif metric == "plr":
# LR+ = sensitivity / (1 - specificity) = tpr / fpr; fpr == 0 yields NaN
tpr = torch.where(p > 0, tp / p, nan_tensor)
fpr = torch.where(n > 0, fp / n, nan_tensor)
numerator, denominator = tpr, fpr
elif metric == "nlr":
# LR- = (1 - sensitivity) / specificity = fnr / tnr; tnr == 0 yields NaN
fnr = torch.where(p > 0, fn / p, nan_tensor)
tnr = torch.where(n > 0, tn / n, nan_tensor)
numerator, denominator = fnr, tnr
else:
raise NotImplementedError("the metric is not implemented.")

Expand Down Expand Up @@ -319,4 +331,8 @@ def check_confusion_matrix_metric_name(metric_name: str) -> str:
return "bm"
if metric_name in ["markedness", "deltap", "mk"]:
return "mk"
if metric_name in ["positive_likelihood_ratio", "plr", "lr+"]:
return "plr"
if metric_name in ["negative_likelihood_ratio", "nlr", "lr-"]:
return "nlr"
raise NotImplementedError("the metric is not implemented.")
62 changes: 62 additions & 0 deletions tests/metrics/test_compute_confusion_matrix.py
Original file line number Diff line number Diff line change
Expand Up @@ -218,6 +218,26 @@
torch.tensor([[[0.0, 0.0, 46137344.0, 0.0]]]),
]

# 5. likelihood ratios: hand-computed LR+ / LR- values, including undefined cases
# sample 0: tp=80, fp=20, tn=70, fn=20 -> tpr=0.8, fpr=2/9 -> LR+=3.6; fnr=0.2, tnr=7/9 -> LR-=0.25714...
# sample 1: tp=10, fp=0, tn=90, fn=0 -> fpr=0 -> LR+ is NaN; fnr=0 -> LR-=0.0
# sample 2: tp=5, fp=5, tn=0, fn=0 -> tpr=1, fpr=1 -> LR+=1; tnr=0 -> LR- is NaN (zero specificity denominator)
# (row order is [tp, fp, tn, fn], matching compute_confusion_matrix output)
TEST_CASE_LR = [torch.tensor([[80.0, 20.0, 70.0, 20.0], [10.0, 0.0, 90.0, 0.0], [5.0, 5.0, 0.0, 0.0]])]

# classification-style input, channel-wise hand-computed values, no undefined cases:
# ch0: tp=2, fp=1, tn=1, fn=1 -> LR+=4/3, LR-=2/3; ch1: tp=1, fp=1, tn=2, fn=1 -> LR+=3/2, LR-=3/4
TEST_CASE_LR_CLF = [
{
"y_pred": torch.tensor([[1, 0], [1, 0], [0, 1], [1, 0], [0, 1]]),
"y": torch.tensor([[1, 0], [1, 0], [0, 1], [0, 1], [1, 0]]),
"include_background": True,
"metric_name": ["positive likelihood ratio", "negative likelihood ratio"],
"reduction": "sum_batch",
},
[torch.tensor([4.0 / 3.0, 3.0 / 2.0]), torch.tensor([2.0 / 3.0, 3.0 / 4.0])],
]


class TestConfusionMatrix(unittest.TestCase):
@parameterized.expand([TEST_CASE_CONFUSION_MATRIX])
Expand Down Expand Up @@ -289,6 +309,48 @@ def test_precision(self, input_data, expected_value):
assert_allclose(result, expected_value, atol=1e-4, rtol=1e-4)
np.testing.assert_equal(result.device, input_data["y_pred"].device)

@parameterized.expand([TEST_CASE_LR])
def test_likelihood_ratios(self, confusion_matrix):
Comment thread
coderabbitai[bot] marked this conversation as resolved.
"""Check likelihood-ratio aliases and edge cases on a per-class confusion matrix.

Args:
confusion_matrix: a stacked [2, 4] confusion-matrix tensor, each row in
``[tp, fp, tn, fn]`` order as produced by ``compute_confusion_matrix``.
"""
# every advertised spelling must resolve to the same result (case/space-insensitive)
for alias in ("lr+", "plr", "Positive Likelihood Ratio", "POSITIVE_LIKELIHOOD_RATIO"):
plr = compute_confusion_matrix_metric(alias, confusion_matrix)
assert_allclose(plr[0], torch.tensor(3.6), atol=1e-4, rtol=1e-4)
self.assertTrue(torch.isnan(plr[1]))
assert_allclose(plr[2], torch.tensor(1.0), atol=1e-4, rtol=1e-4)
for alias in ("lr-", "nlr", "Negative Likelihood Ratio", "NEGATIVE_LIKELIHOOD_RATIO"):
nlr = compute_confusion_matrix_metric(alias, confusion_matrix)
assert_allclose(nlr[0], torch.tensor(0.2 / (70.0 / 90.0)), atol=1e-4, rtol=1e-4)
assert_allclose(nlr[1], torch.tensor(0.0), atol=1e-4, rtol=1e-4)
# sample 2 has tn == 0 (zero specificity), so LR- is undefined -> NaN
self.assertTrue(torch.isnan(nlr[2]))

@parameterized.expand([TEST_CASE_LR_CLF])
def test_likelihood_ratios_clf(self, input_data, expected_values):
"""Check likelihood ratios through the ``ConfusionMatrixMetric`` classification API.

Args:
input_data: keyword arguments for ``ConfusionMatrixMetric`` plus ``y_pred``/``y``
to feed the metric.
expected_values: expected per-channel LR+ / LR- values after aggregation.
"""
params = input_data.copy()
vals = {}
vals["y_pred"] = params.pop("y_pred")
vals["y"] = params.pop("y")
metric = ConfusionMatrixMetric(**params)
metric(**vals)
results = metric.aggregate()
# one aggregated channel per requested metric, in the same order as expected_values
self.assertEqual(len(results), len(expected_values))
for result, expected_value in zip(results, expected_values):
Comment thread
coderabbitai[bot] marked this conversation as resolved.
assert_allclose(result, expected_value, atol=1e-4, rtol=1e-4)


if __name__ == "__main__":
unittest.main()