diff --git a/sdm/models/kumo/relational/model.py b/sdm/models/kumo/relational/model.py index 54247ad98..35b234214 100644 --- a/sdm/models/kumo/relational/model.py +++ b/sdm/models/kumo/relational/model.py @@ -590,13 +590,9 @@ def _get_rel_time( rel_time.nanmean(dim=-2, keepdim=True).nan_to_num(0.0), rel_time, ) - rel_time = standardizer.fit_transform( - TableTensor.from_tensor(rel_time) - ).numerical + rel_time = standardizer._fit_transform_tensor(rel_time) else: - rel_time = standardizer.transform( - TableTensor.from_tensor(rel_time) - ).numerical + rel_time = standardizer._transform_tensor(rel_time) rel_time[na_mask] = 0.0 return rel_time diff --git a/sdm/processing/numerical/standardize.py b/sdm/processing/numerical/standardize.py index 492cff947..8333d0147 100644 --- a/sdm/processing/numerical/standardize.py +++ b/sdm/processing/numerical/standardize.py @@ -2,6 +2,7 @@ # SPDX-License-Identifier: Apache-2.0 import torch +from torch import Tensor from sdm import Stype, TableTensor from sdm.processing import InvertibleMixin, Processor @@ -40,9 +41,12 @@ def _fit( generator: torch.Generator | None = None, ) -> None: - finite = _isfinite(table.numerical) + self._fit_tensor(table.numerical) + + def _fit_tensor(self, numerical: Tensor) -> None: + finite = _isfinite(numerical) count = finite.sum(dim=-2, keepdim=True) - finite_or_nan = table.numerical.masked_fill(~finite, torch.nan) + finite_or_nan = numerical.masked_fill(~finite, torch.nan) self.mean = finite_or_nan.nansum(-2, keepdim=True).div_(count) self.mean.masked_fill_(self.mean.isnan(), 0.0) @@ -59,8 +63,18 @@ def _fit( self.scale += self.eps def _transform(self, table: TableTensor) -> TableTensor: - numerical = (table.numerical - self.mean).div_(self.scale) - return table.replace_blocks(numerical=numerical) + return table.replace_blocks( + numerical=self._transform_tensor(table.numerical) + ) + + def _transform_tensor(self, numerical: Tensor) -> Tensor: + return (numerical - self.mean).div_(self.scale) + + def _fit_transform_tensor(self, numerical: Tensor) -> Tensor: + self._fit_tensor(numerical) + out = self._transform_tensor(numerical) + self._set_fitted(numerical.device) + return out def _inverse_transform(self, table: TableTensor) -> TableTensor: dtype = table.numerical.dtype diff --git a/test/processing/numerical/test_standardize.py b/test/processing/numerical/test_standardize.py index f36e053e9..9786b5d42 100644 --- a/test/processing/numerical/test_standardize.py +++ b/test/processing/numerical/test_standardize.py @@ -8,6 +8,25 @@ from sdm.testing import withCUDA +@withCUDA +def test_standardize_tensor_fit_preserves_fitted_behavior( + device: torch.device, +) -> None: + context = torch.tensor([[1.0, 5.0], [3.0, 5.0]], device=device) + query = TableTensor.from_tensor(context + 2.0) + processor = Standardize() + reference = Standardize() + + out = processor._fit_transform_tensor(context) + expected = reference.fit_transform(TableTensor.from_tensor(context)) + torch.testing.assert_close(out, expected.numerical) + assert processor.is_fitted + torch.testing.assert_close( + processor.transform(query).numerical, + reference.transform(query).numerical, + ) + + @withCUDA def test_standardize_fit_transform_and_inverse_round_trip( device: torch.device,