diff --git a/sdm/processing/numerical/impute.py b/sdm/processing/numerical/impute.py index 239461144..b8dbde2a2 100644 --- a/sdm/processing/numerical/impute.py +++ b/sdm/processing/numerical/impute.py @@ -5,14 +5,18 @@ from sdm import Stype, TableTensor from sdm.processing import Processor +from sdm.processing.numerical._stats import _isfinite class ImputeMean(Processor): """Replace NaN feature values with fitted per-column means. + Infinite values are ignored when fitting the means and preserved during + the transform. + Args: fill_value: Finite value used for columns whose fitted mean is - undefined (e.g. all-NaN columns). + undefined (e.g. columns without finite values). """ handles_stypes = frozenset({Stype.numerical}) @@ -35,7 +39,7 @@ def _fit( ) -> None: numerical = table.numerical mean = torch.nanmean( - numerical, + numerical.masked_fill(~_isfinite(numerical), torch.nan), dim=-2, keepdim=True, ) diff --git a/test/processing/numerical/test_impute.py b/test/processing/numerical/test_impute.py index be8872c26..137c797aa 100644 --- a/test/processing/numerical/test_impute.py +++ b/test/processing/numerical/test_impute.py @@ -54,3 +54,33 @@ def test_impute_mean(device: torch.device, dtype: torch.dtype | None) -> None: device=device, ), ) + + +@withCUDA +def test_impute_mean_ignores_infinities(device: torch.device) -> None: + inf = torch.inf + inp = torch.tensor( + [ + [1.0, inf, inf], + [3.0, 2.0, -inf], + [torch.nan, torch.nan, torch.nan], + [-inf, 4.0, torch.nan], + ], + device=device, + ) + + processor = ImputeMean(fill_value=-5.0).fit(TableTensor.from_tensor(inp)) + transformed = processor.transform(TableTensor.from_tensor(inp)).numerical + + assert torch.equal( + transformed, + torch.tensor( + [ + [1.0, inf, inf], + [3.0, 2.0, -inf], + [2.0, 3.0, -5.0], + [-inf, 4.0, -5.0], + ], + device=device, + ), + )