Skip to content
Draft
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
8 changes: 6 additions & 2 deletions sdm/processing/numerical/impute.py
Original file line number Diff line number Diff line change
Expand Up @@ -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})
Expand All @@ -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,
)
Expand Down
30 changes: 30 additions & 0 deletions test/processing/numerical/test_impute.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
),
)