diff --git a/sdm/processing/numerical/_stats.py b/sdm/processing/numerical/_stats.py index 002d435c7..30d21f0e1 100644 --- a/sdm/processing/numerical/_stats.py +++ b/sdm/processing/numerical/_stats.py @@ -1,10 +1,17 @@ # SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 +import math + import torch from torch import Tensor +def _isfinite(x: Tensor) -> Tensor: + # Equal to 'x.isfinite()', which allocates 'x.abs()' on the way. + return x.gt(-math.inf).logical_and_(x.lt(math.inf)) + + def _constant_feature_mask( var: Tensor, mean: Tensor, diff --git a/sdm/processing/numerical/clip_soft.py b/sdm/processing/numerical/clip_soft.py index 716008655..ffd6cd84b 100644 --- a/sdm/processing/numerical/clip_soft.py +++ b/sdm/processing/numerical/clip_soft.py @@ -7,6 +7,7 @@ from sdm import Stype, TableTensor from sdm.processing import Processor +from sdm.processing.numerical._stats import _isfinite class ClipSoft(Processor): @@ -38,16 +39,14 @@ def __init__(self, max_absolute_value: float) -> None: def _transform(self, table: TableTensor) -> TableTensor: numerical = table.numerical + assert not numerical.requires_grad bound = self.max_absolute_value - ratio = (numerical / bound).abs() - squared = 1 + ratio.square() - unit = torch.where( - squared.isfinite(), - ratio / squared.sqrt(), - 1.0, - ) - clipped = numerical.sign() * bound * unit - clipped = torch.where(numerical.isnan(), numerical, clipped) + unit = numerical.div(bound).abs_() + root = unit.square().add_(1).sqrt_() + unit.div_(root).masked_fill_(~_isfinite(root), 1.0) + del root + clipped = unit.mul_(bound).mul_(numerical.sign()) + torch.where(numerical.isnan(), numerical, clipped, out=clipped) return table.replace_blocks(numerical=clipped) def __repr__(self, *, indent: int = 0) -> str: diff --git a/sdm/processing/numerical/power.py b/sdm/processing/numerical/power.py index c62424226..73f34a4b2 100644 --- a/sdm/processing/numerical/power.py +++ b/sdm/processing/numerical/power.py @@ -8,7 +8,10 @@ from sdm import Stype, TableTensor from sdm.processing import InvertibleMixin, Processor -from sdm.processing.numerical._stats import _constant_feature_mask +from sdm.processing.numerical._stats import ( + _constant_feature_mask, + _isfinite, +) # Keep GPU execution batched; adaptive per-column stopping would resynchronize. # For float32 overflow-safe bounds, 44 golden steps reaches ~1.48e-8. @@ -68,7 +71,7 @@ def _yeojohnson_inverse_transform(inp: Tensor, lambdas: Tensor) -> Tensor: def _yeojohnson_bounds(inp: Tensor) -> tuple[Tensor, Tensor]: missing = inp.isnan() - max_abs = inp.abs().nan_to_num(nan=0.0).amax(dim=-2, keepdim=True) + max_abs = inp.abs().nan_to_num_(nan=0.0).amax(dim=-2, keepdim=True) log1p_max_x = (20 * max_abs).log1p() log1p_max_x = torch.where( max_abs == 0, @@ -124,8 +127,9 @@ def _yeojohnson_log_likelihood( exponents=exponents, out=transformed, ) - mean = transformed.nanmean(dim=-2, keepdim=True) - variance = transformed.sub_(mean).square_().nanmean(dim=-2, keepdim=True) + mean = transformed.nansum(dim=-2, keepdim=True).div_(count) + variance = transformed.sub_(mean).square_().nansum(dim=-2, keepdim=True) + variance /= count tiny = torch.finfo(inp.dtype).tiny valid = variance.isfinite() & (variance >= tiny) loglike = variance.log_().mul_(-count / 2) @@ -139,6 +143,10 @@ def _optimize_lambdas( *, count: Tensor, ) -> Tensor: + left, right = _yeojohnson_bounds(inp) + left = left.masked_fill(constant_features, 1.0) + right = right.masked_fill(constant_features, 1.0) + # Reuse full-table workspaces throughout the golden-section search. magnitude_log = inp.abs().log1p_() positive = inp >= 0 @@ -146,10 +154,6 @@ def _optimize_lambdas( exponents = torch.empty_like(inp) transformed = torch.empty_like(inp) - left, right = _yeojohnson_bounds(inp) - left = left.masked_fill(constant_features, 1.0) - right = right.masked_fill(constant_features, 1.0) - invphi = (math.sqrt(5) - 1) / 2 span = (right - left).mul_(invphi) c = right - span @@ -252,14 +256,13 @@ def _fit( *, generator: torch.Generator | None = None, ) -> None: - finite = table.numerical.isfinite() + finite = _isfinite(table.numerical) finite_or_nan = table.numerical.masked_fill(~finite, torch.nan) - count = finite.sum(dim=-2, keepdim=True) + count = finite.sum(dim=-2, keepdim=True).clamp_(min=1) - mean = finite_or_nan.nanmean(-2, keepdim=True) - mean.masked_fill_(mean.isnan(), 0.0) - var = (finite_or_nan - mean).square().nanmean(-2, keepdim=True) - var.masked_fill_(var.isnan(), 0.0) + mean = finite_or_nan.nansum(-2, keepdim=True).div_(count) + var = finite_or_nan.sub(mean).square_().nansum(-2, keepdim=True) + var /= count constant_features = _constant_feature_mask( var, mean, @@ -283,10 +286,10 @@ def _fit( if self.standardize: transformed = _yeojohnson_transform(finite_or_nan, self.lambdas) - mean = transformed.nanmean(dim=-2, keepdim=True) - mean.masked_fill_(mean.isnan(), 0.0) - var = (transformed - mean).square().nanmean(dim=-2, keepdim=True) - var.masked_fill_(var.isnan(), 0.0) + del finite_or_nan + mean = transformed.nansum(dim=-2, keepdim=True).div_(count) + var = transformed.sub_(mean).square_().nansum(-2, keepdim=True) + var /= count scale = var.sqrt() scale[_constant_feature_mask(var, mean, count)] = 1.0 self.mean = mean @@ -312,7 +315,7 @@ def _inverse_transform(self, table: TableTensor) -> TableTensor: # Above the fitted upper bound the inverse diverges, either to # infinity or, past the asymptote, to NaN. - diverged = ~inverse.isfinite() & ~unscaled.isnan() + diverged = ~_isfinite(inverse) & ~unscaled.isnan() return table.replace_blocks( numerical=torch.where( diverged, diff --git a/sdm/processing/numerical/robust_scale.py b/sdm/processing/numerical/robust_scale.py index f9dd58b4d..f2ba25506 100644 --- a/sdm/processing/numerical/robust_scale.py +++ b/sdm/processing/numerical/robust_scale.py @@ -5,6 +5,7 @@ from sdm import Stype, TableTensor from sdm.processing import InvertibleMixin, Processor +from sdm.processing.numerical._stats import _isfinite class RobustScale(Processor, InvertibleMixin): @@ -44,7 +45,7 @@ def _fit( generator: torch.Generator | None = None, ) -> None: numerical = table.numerical - finite_or_nan = numerical.masked_fill(~numerical.isfinite(), torch.nan) + finite_or_nan = numerical.masked_fill(~_isfinite(numerical), torch.nan) q_low, q_high = (value / 100.0 for value in self.quantile_range) # 'nanquantile' requires single or double precision input. quantile_input = finite_or_nan.to( @@ -60,11 +61,11 @@ def _fit( self.scale = scale.to(dtype=numerical.dtype) def _transform(self, table: TableTensor) -> TableTensor: - numerical = (table.numerical - self.median) / self.scale + numerical = (table.numerical - self.median).div_(self.scale) return table.replace_blocks(numerical=numerical) def _inverse_transform(self, table: TableTensor) -> TableTensor: - numerical = table.numerical * self.scale + self.median + numerical = (table.numerical * self.scale).add_(self.median) return table.replace_blocks(numerical=numerical) def __repr__(self, *, indent: int = 0) -> str: diff --git a/sdm/processing/numerical/sigma_clip.py b/sdm/processing/numerical/sigma_clip.py index 81444cffa..511b10df5 100644 --- a/sdm/processing/numerical/sigma_clip.py +++ b/sdm/processing/numerical/sigma_clip.py @@ -5,6 +5,7 @@ from sdm import Stype, TableTensor from sdm.processing import Processor +from sdm.processing.numerical._stats import _isfinite class ClipSigma(Processor): @@ -43,30 +44,37 @@ def _fit( generator: torch.Generator | None = None, ) -> None: - finite = table.numerical.isfinite() - finite_or_nan = table.numerical.masked_fill(~finite, torch.nan) + numerical = table.numerical + assert not numerical.requires_grad + finite = _isfinite(numerical) + count_finite = finite.sum(dim=-2, keepdim=True) + finite_or_nan = numerical.masked_fill(~finite, torch.nan) # Compute finite mean and standard deviation: - mean = finite_or_nan.nanmean(-2, keepdim=True) + mean = finite_or_nan.nansum(-2, keepdim=True).div_(count_finite) mean.masked_fill_(mean.isnan(), 0.0) - var = (finite_or_nan - mean).square().nansum(-2, keepdim=True) - var /= (finite.sum(-2, keepdim=True) - 1).clamp_(min=1) + var = finite_or_nan.sub_(mean).square_().nansum(-2, keepdim=True) + del finite_or_nan + var /= (count_finite - 1).clamp_(min=1) std = var.sqrt().clamp(min=1e-6) - # Find values within range: + # Find values within range (non-finite values are never kept): lower = mean - self.threshold * std upper = mean + self.threshold * std - keep = finite & (finite_or_nan >= lower) & (finite_or_nan <= upper) - count = keep.sum(-2, keepdim=True) + keep = (numerical >= lower).logical_and_(numerical <= upper) + keep &= finite + del finite + count = keep.sum(dim=-2, keepdim=True) # Compute mean and standard deviation of kept values: - kept_mean = torch.where(keep, finite_or_nan, 0.0).sum(-2, keepdim=True) + kept_mean = torch.where(keep, numerical, 0.0).sum(-2, keepdim=True) kept_mean /= count.clamp(min=1) - centered = torch.where(keep, finite_or_nan - kept_mean, 0.0) + centered = numerical.sub(kept_mean).masked_fill_(~keep, 0.0) denominator = (count - 1).clamp(min=1) - kept_var = centered.square().sum(-2, keepdim=True) / denominator + kept_var = centered.square_().sum(-2, keepdim=True) / denominator + del centered kept_std = kept_var.sqrt().clamp(min=1e-6) has_kept = count > 0 diff --git a/sdm/processing/numerical/standardize.py b/sdm/processing/numerical/standardize.py index 8326d7db6..376a91555 100644 --- a/sdm/processing/numerical/standardize.py +++ b/sdm/processing/numerical/standardize.py @@ -5,7 +5,10 @@ from sdm import Stype, TableTensor from sdm.processing import InvertibleMixin, Processor -from sdm.processing.numerical._stats import _constant_feature_mask +from sdm.processing.numerical._stats import ( + _constant_feature_mask, + _isfinite, +) class Standardize(Processor, InvertibleMixin): @@ -37,35 +40,33 @@ def _fit( generator: torch.Generator | None = None, ) -> None: - finite = table.numerical.isfinite() + finite = _isfinite(table.numerical) + count = finite.sum(dim=-2, keepdim=True) finite_or_nan = table.numerical.masked_fill(~finite, torch.nan) finite_or_nan = finite_or_nan.double() # Ensure high precision. - self.mean = finite_or_nan.nanmean(-2, keepdim=True) + self.mean = finite_or_nan.nansum(-2, keepdim=True).div_(count) self.mean.masked_fill_(self.mean.isnan(), 0.0) - var = (finite_or_nan - self.mean).square().nanmean(-2, keepdim=True) + var = finite_or_nan.sub_(self.mean).square_().nansum(-2, keepdim=True) + var /= count var.masked_fill_(var.isnan(), 0.0) self.scale = var.sqrt() if self.eps == 0: - mask = _constant_feature_mask( - var, - self.mean, - num_samples=finite.sum(-2, keepdim=True), - ) + mask = _constant_feature_mask(var, self.mean, num_samples=count) self.scale[mask] = 1.0 else: self.scale += self.eps def _transform(self, table: TableTensor) -> TableTensor: dtype = table.numerical.dtype - numerical = (table.numerical - self.mean) / self.scale + numerical = (table.numerical - self.mean).div_(self.scale) return table.replace_blocks(numerical=numerical.to(dtype=dtype)) def _inverse_transform(self, table: TableTensor) -> TableTensor: dtype = table.numerical.dtype - numerical = table.numerical * self.scale + self.mean + numerical = (table.numerical * self.scale).add_(self.mean) return table.replace_blocks(numerical=numerical.to(dtype=dtype)) def __repr__(self, *, indent: int = 0) -> str: diff --git a/test/processing/numerical/test_standardize.py b/test/processing/numerical/test_standardize.py index 375c736d4..8c709942d 100644 --- a/test/processing/numerical/test_standardize.py +++ b/test/processing/numerical/test_standardize.py @@ -76,6 +76,26 @@ def test_standardize(device: torch.device) -> None: ) +@withCUDA +def test_standardize_ignores_non_finite_values_in_many_rows( + device: torch.device, +) -> None: + inp = torch.randn(2, 600, 3, dtype=torch.float64, device=device) + inp[inp > 1.0] = float("nan") + inp[inp < -1.5] = float("inf") + + out = Standardize().fit_transform(TableTensor.from_tensor(inp)) + + finite = inp.masked_fill(~inp.isfinite(), float("nan")) + mean = finite.nanmean(-2, keepdim=True) + std = (finite - mean).square().nanmean(-2, keepdim=True).sqrt() + torch.testing.assert_close( + out.numerical, + (inp - mean) / std, + equal_nan=True, + ) + + @withCUDA def test_standardize_single_sample_uses_unit_scale( device: torch.device,