From 5dea9b349aee7d92be79dd3ead66efc84a87aa07 Mon Sep 17 00:00:00 2001 From: Dante Rigo Date: Fri, 28 Aug 2026 07:42:23 -0700 Subject: [PATCH] Fix FID metric for scipy >= 1.18: sqrtm no longer accepts disp scipy 1.18 removed the `disp` parameter from `scipy.linalg.sqrtm`, and with it the 2-tuple return that `disp=False` produced. `_sqrtm` still called `sqrtm(..., disp=False)` and unpacked two values, so every FID computation raised `TypeError` on scipy >= 1.18. Call `sqrtm` without `disp` and use its return value directly. This works across MONAI's whole supported range (scipy >= 1.12), since older versions also return the matrix alone when `disp` is left at its default. Verified on scipy 1.18.1 / py3.12 and scipy 1.12.0 / py3.11. Signed-off-by: Dante Rigo --- monai/metrics/fid.py | 2 +- tests/metrics/test_compute_fid_metric.py | 7 +++++++ 2 files changed, 8 insertions(+), 1 deletion(-) diff --git a/monai/metrics/fid.py b/monai/metrics/fid.py index 596f9aef7cb..1df9ef00fa2 100644 --- a/monai/metrics/fid.py +++ b/monai/metrics/fid.py @@ -82,7 +82,7 @@ def _cov(input_data: torch.Tensor, rowvar: bool = True) -> torch.Tensor: def _sqrtm(input_data: torch.Tensor) -> torch.Tensor: """Compute the square root of a matrix.""" - scipy_res, _ = scipy.linalg.sqrtm(input_data.detach().cpu().numpy().astype(np.float64), disp=False) + scipy_res = scipy.linalg.sqrtm(input_data.detach().cpu().numpy().astype(np.float64)) return torch.from_numpy(scipy_res) diff --git a/tests/metrics/test_compute_fid_metric.py b/tests/metrics/test_compute_fid_metric.py index bd867f5296e..cd1c1cfe94e 100644 --- a/tests/metrics/test_compute_fid_metric.py +++ b/tests/metrics/test_compute_fid_metric.py @@ -17,6 +17,7 @@ import torch from monai.metrics import FIDMetric +from monai.metrics.fid import _sqrtm from monai.utils import optional_import _, has_scipy = optional_import("scipy") @@ -25,6 +26,12 @@ @unittest.skipUnless(has_scipy, "Requires scipy") class TestFIDMetric(unittest.TestCase): + def test_sqrtm_returns_tensor(self): + """``scipy.linalg.sqrtm`` dropped its ``disp`` argument, and with it the 2-tuple return.""" + result = _sqrtm(torch.tensor([[4.0, 0.0], [0.0, 9.0]], dtype=torch.float64)) + self.assertIsInstance(result, torch.Tensor) + np.testing.assert_allclose(result.cpu().numpy(), np.array([[2.0, 0.0], [0.0, 3.0]]), atol=1e-6) + def test_results(self): x = torch.Tensor([[1, 2], [1, 2], [1, 2]]) y = torch.Tensor([[2, 2], [1, 2], [1, 2]])