From 2a09198d4106edb5819e770174c7ef015aa91d3a Mon Sep 17 00:00:00 2001 From: Jingang Qu Date: Mon, 28 Sep 2026 22:40:27 +0200 Subject: [PATCH 1/4] Reduce numerical preprocessing temporaries - Find finite values without an `abs()` copy and count them without an int64 copy of the mask, and compute Standardize, ClipSigma and PowerTransform statistics with fewer full-size temporaries. - Find the Yeo-Johnson bounds before allocating the PowerTransform workspaces. - Transform in Standardize, ClipSoft and RobustScale with fewer temporaries, keeping out-of-place ops where gradients are required. Signed-off-by: Jingang Qu Co-authored-by: Cedric Lorenz --- sdm/processing/numerical/_stats.py | 18 ++++++++ sdm/processing/numerical/clip_soft.py | 25 +++++++---- sdm/processing/numerical/power.py | 42 ++++++++++++------- sdm/processing/numerical/robust_scale.py | 7 ++-- sdm/processing/numerical/sigma_clip.py | 37 ++++++++++------ sdm/processing/numerical/standardize.py | 26 +++++++----- test/processing/numerical/test_sigma_clip.py | 25 +++++++++++ test/processing/numerical/test_standardize.py | 20 +++++++++ 8 files changed, 150 insertions(+), 50 deletions(-) diff --git a/sdm/processing/numerical/_stats.py b/sdm/processing/numerical/_stats.py index 002d435c7..df3d84adf 100644 --- a/sdm/processing/numerical/_stats.py +++ b/sdm/processing/numerical/_stats.py @@ -1,10 +1,28 @@ # 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 _count(mask: Tensor) -> Tensor: + # [..., N, C] -> [..., 1, C] int64 number of true values per column. + # Summing bool first casts all of 'mask' to int64. Sum blocks of 255 rows + # as uint8 instead, which cannot overflow. + num_blocks = mask.size(-2) // 255 + blocks = mask[..., : num_blocks * 255, :].unflatten(-2, (num_blocks, 255)) + count = blocks.view(torch.uint8).sum(-2, dtype=torch.uint8) + remainder = mask[..., num_blocks * 255 :, :] + return count.sum(-2, keepdim=True) + remainder.sum(-2, keepdim=True) + + 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..1d6ec874c 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): @@ -39,15 +40,21 @@ def __init__(self, max_absolute_value: float) -> None: def _transform(self, table: TableTensor) -> TableTensor: numerical = table.numerical 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) + if torch.is_grad_enabled() and numerical.requires_grad: + 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) + return table.replace_blocks(numerical=clipped) + + unit = numerical.div(bound).abs_() + root = unit.square().add_(1).sqrt_() + # 'root' is finite exactly where '1 + unit ** 2' is. + unit.div_(root).masked_fill_(~_isfinite(root), 1.0) + del root + clipped = numerical.sign().mul_(bound).mul_(unit) + 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..2fdffe9b1 100644 --- a/sdm/processing/numerical/power.py +++ b/sdm/processing/numerical/power.py @@ -8,7 +8,11 @@ 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, + _count, + _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 +72,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 +128,11 @@ 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) + # 'transformed' is NaN exactly where 'inp' is, so this equals 'nanmean', + # which would copy its input to count values. + mean = transformed.nansum(dim=-2, keepdim=True) / 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 +146,11 @@ def _optimize_lambdas( *, count: Tensor, ) -> Tensor: + # Find the bounds before allocating the workspaces. + 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 +158,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,13 +260,15 @@ 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 = _count(finite) - mean = finite_or_nan.nanmean(-2, keepdim=True) + # Equal to 'nanmean', which would copy its input to count values: + mean = finite_or_nan.nansum(-2, keepdim=True) / count mean.masked_fill_(mean.isnan(), 0.0) - var = (finite_or_nan - mean).square().nanmean(-2, keepdim=True) + var = finite_or_nan.sub(mean).square_().nansum(-2, keepdim=True) + var /= count var.masked_fill_(var.isnan(), 0.0) constant_features = _constant_feature_mask( var, @@ -283,9 +293,11 @@ def _fit( if self.standardize: transformed = _yeojohnson_transform(finite_or_nan, self.lambdas) - mean = transformed.nanmean(dim=-2, keepdim=True) + del finite_or_nan + mean = transformed.nansum(dim=-2, keepdim=True) / count mean.masked_fill_(mean.isnan(), 0.0) - var = (transformed - mean).square().nanmean(dim=-2, keepdim=True) + var = transformed.sub_(mean).square_().nansum(-2, keepdim=True) + var /= count var.masked_fill_(var.isnan(), 0.0) scale = var.sqrt() scale[_constant_feature_mask(var, mean, count)] = 1.0 @@ -312,7 +324,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..e8e23ef1e 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.sub(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.mul(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..904b49ec0 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 _count, _isfinite class ClipSigma(Processor): @@ -43,30 +44,42 @@ 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 + finite = _isfinite(numerical) + count_finite = _count(finite) + finite_or_nan = numerical.masked_fill(~finite, torch.nan) - # Compute finite mean and standard deviation: - mean = finite_or_nan.nanmean(-2, keepdim=True) + # Compute finite mean and standard deviation (equal to 'nanmean', + # which would copy its input to count values): + mean = finite_or_nan.nansum(-2, keepdim=True) / 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) + centered = ( + finite_or_nan.sub(mean) + if torch.is_grad_enabled() + else finite_or_nan.sub_(mean) + ) + var = centered.square_().nansum(-2, keepdim=True) + del centered, 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.logical_and_(finite) + del finite + count = _count(keep) # 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..8beb10fb0 100644 --- a/sdm/processing/numerical/standardize.py +++ b/sdm/processing/numerical/standardize.py @@ -5,7 +5,11 @@ 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, + _count, + _isfinite, +) class Standardize(Processor, InvertibleMixin): @@ -37,35 +41,35 @@ def _fit( generator: torch.Generator | None = None, ) -> None: - finite = table.numerical.isfinite() + finite = _isfinite(table.numerical) + count = _count(finite) 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) + # Equal to 'nanmean', which would copy its input to count values. + self.mean = finite_or_nan.nansum(-2, keepdim=True) / count self.mean.masked_fill_(self.mean.isnan(), 0.0) - var = (finite_or_nan - self.mean).square().nanmean(-2, keepdim=True) + # 'finite_or_nan' is a private copy, so it can be modified in place. + 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.sub(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.mul(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_sigma_clip.py b/test/processing/numerical/test_sigma_clip.py index d29cd94f5..155750666 100644 --- a/test/processing/numerical/test_sigma_clip.py +++ b/test/processing/numerical/test_sigma_clip.py @@ -153,3 +153,28 @@ def test_clip_sigma_preserves_nonfinite( ), equal_nan=True, ) + + +@withCUDA +@pytest.mark.parametrize("fit_transform", [False, True]) +def test_clip_sigma_gradients( + device: torch.device, + fit_transform: bool, +) -> None: + context = torch.tensor( + [[1.0], [2.0], [4.0], [100.0]], + dtype=torch.float64, + device=device, + requires_grad=True, + ) + query = TableTensor.from_tensor(context.detach().clone()) + + def transform(values: torch.Tensor) -> torch.Tensor: + processor = ClipSigma(threshold=1.0) + table = TableTensor.from_tensor(values) + if fit_transform: + return processor.fit_transform(table).numerical + processor.fit(table) + return processor.transform(query).numerical + + assert torch.autograd.gradcheck(transform, (context,)) 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, From cb8ba437da1b945b26fb83accba90e877e783655 Mon Sep 17 00:00:00 2001 From: Valter Hudovernik Date: Tue, 29 Sep 2026 11:09:39 +0200 Subject: [PATCH 2/4] Update sdm/processing/numerical/power.py Co-authored-by: Matthias Fey --- sdm/processing/numerical/power.py | 1 - 1 file changed, 1 deletion(-) diff --git a/sdm/processing/numerical/power.py b/sdm/processing/numerical/power.py index 2fdffe9b1..c4198c1a8 100644 --- a/sdm/processing/numerical/power.py +++ b/sdm/processing/numerical/power.py @@ -146,7 +146,6 @@ def _optimize_lambdas( *, count: Tensor, ) -> Tensor: - # Find the bounds before allocating the workspaces. left, right = _yeojohnson_bounds(inp) left = left.masked_fill(constant_features, 1.0) right = right.masked_fill(constant_features, 1.0) From 8c1e7c6abef29696acda950b6ebaf0bbcbc026ce Mon Sep 17 00:00:00 2001 From: Valter Hudovernik Date: Tue, 29 Sep 2026 11:38:23 +0200 Subject: [PATCH 3/4] address review feedback --- sdm/processing/numerical/_stats.py | 11 --------- sdm/processing/numerical/clip_soft.py | 12 ++-------- sdm/processing/numerical/power.py | 16 ++++--------- sdm/processing/numerical/robust_scale.py | 4 ++-- sdm/processing/numerical/sigma_clip.py | 23 +++++++----------- sdm/processing/numerical/standardize.py | 11 ++++----- test/processing/numerical/test_sigma_clip.py | 25 -------------------- 7 files changed, 21 insertions(+), 81 deletions(-) diff --git a/sdm/processing/numerical/_stats.py b/sdm/processing/numerical/_stats.py index df3d84adf..30d21f0e1 100644 --- a/sdm/processing/numerical/_stats.py +++ b/sdm/processing/numerical/_stats.py @@ -12,17 +12,6 @@ def _isfinite(x: Tensor) -> Tensor: return x.gt(-math.inf).logical_and_(x.lt(math.inf)) -def _count(mask: Tensor) -> Tensor: - # [..., N, C] -> [..., 1, C] int64 number of true values per column. - # Summing bool first casts all of 'mask' to int64. Sum blocks of 255 rows - # as uint8 instead, which cannot overflow. - num_blocks = mask.size(-2) // 255 - blocks = mask[..., : num_blocks * 255, :].unflatten(-2, (num_blocks, 255)) - count = blocks.view(torch.uint8).sum(-2, dtype=torch.uint8) - remainder = mask[..., num_blocks * 255 :, :] - return count.sum(-2, keepdim=True) + remainder.sum(-2, keepdim=True) - - 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 1d6ec874c..b192966ad 100644 --- a/sdm/processing/numerical/clip_soft.py +++ b/sdm/processing/numerical/clip_soft.py @@ -39,21 +39,13 @@ 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 - if torch.is_grad_enabled() and numerical.requires_grad: - 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) - return table.replace_blocks(numerical=clipped) - unit = numerical.div(bound).abs_() root = unit.square().add_(1).sqrt_() - # 'root' is finite exactly where '1 + unit ** 2' is. unit.div_(root).masked_fill_(~_isfinite(root), 1.0) del root - clipped = numerical.sign().mul_(bound).mul_(unit) + clipped = (numerical.sign() * bound).mul_(unit) torch.where(numerical.isnan(), numerical, clipped, out=clipped) return table.replace_blocks(numerical=clipped) diff --git a/sdm/processing/numerical/power.py b/sdm/processing/numerical/power.py index c4198c1a8..73f34a4b2 100644 --- a/sdm/processing/numerical/power.py +++ b/sdm/processing/numerical/power.py @@ -10,7 +10,6 @@ from sdm.processing import InvertibleMixin, Processor from sdm.processing.numerical._stats import ( _constant_feature_mask, - _count, _isfinite, ) @@ -128,9 +127,7 @@ def _yeojohnson_log_likelihood( exponents=exponents, out=transformed, ) - # 'transformed' is NaN exactly where 'inp' is, so this equals 'nanmean', - # which would copy its input to count values. - mean = transformed.nansum(dim=-2, keepdim=True) / count + 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 @@ -261,14 +258,11 @@ def _fit( ) -> None: finite = _isfinite(table.numerical) finite_or_nan = table.numerical.masked_fill(~finite, torch.nan) - count = _count(finite) + count = finite.sum(dim=-2, keepdim=True).clamp_(min=1) - # Equal to 'nanmean', which would copy its input to count values: - mean = finite_or_nan.nansum(-2, keepdim=True) / count - mean.masked_fill_(mean.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 - var.masked_fill_(var.isnan(), 0.0) constant_features = _constant_feature_mask( var, mean, @@ -293,11 +287,9 @@ def _fit( if self.standardize: transformed = _yeojohnson_transform(finite_or_nan, self.lambdas) del finite_or_nan - mean = transformed.nansum(dim=-2, keepdim=True) / count - mean.masked_fill_(mean.isnan(), 0.0) + mean = transformed.nansum(dim=-2, keepdim=True).div_(count) var = transformed.sub_(mean).square_().nansum(-2, keepdim=True) var /= count - var.masked_fill_(var.isnan(), 0.0) scale = var.sqrt() scale[_constant_feature_mask(var, mean, count)] = 1.0 self.mean = mean diff --git a/sdm/processing/numerical/robust_scale.py b/sdm/processing/numerical/robust_scale.py index e8e23ef1e..f2ba25506 100644 --- a/sdm/processing/numerical/robust_scale.py +++ b/sdm/processing/numerical/robust_scale.py @@ -61,11 +61,11 @@ def _fit( self.scale = scale.to(dtype=numerical.dtype) def _transform(self, table: TableTensor) -> TableTensor: - numerical = table.numerical.sub(self.median).div_(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.mul(self.scale).add_(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 904b49ec0..511b10df5 100644 --- a/sdm/processing/numerical/sigma_clip.py +++ b/sdm/processing/numerical/sigma_clip.py @@ -5,7 +5,7 @@ from sdm import Stype, TableTensor from sdm.processing import Processor -from sdm.processing.numerical._stats import _count, _isfinite +from sdm.processing.numerical._stats import _isfinite class ClipSigma(Processor): @@ -45,22 +45,17 @@ def _fit( ) -> None: numerical = table.numerical + assert not numerical.requires_grad finite = _isfinite(numerical) - count_finite = _count(finite) + count_finite = finite.sum(dim=-2, keepdim=True) finite_or_nan = numerical.masked_fill(~finite, torch.nan) - # Compute finite mean and standard deviation (equal to 'nanmean', - # which would copy its input to count values): - mean = finite_or_nan.nansum(-2, keepdim=True) / count_finite + # Compute finite mean and standard deviation: + mean = finite_or_nan.nansum(-2, keepdim=True).div_(count_finite) mean.masked_fill_(mean.isnan(), 0.0) - centered = ( - finite_or_nan.sub(mean) - if torch.is_grad_enabled() - else finite_or_nan.sub_(mean) - ) - var = centered.square_().nansum(-2, keepdim=True) - del centered, finite_or_nan + 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) @@ -68,9 +63,9 @@ def _fit( lower = mean - self.threshold * std upper = mean + self.threshold * std keep = (numerical >= lower).logical_and_(numerical <= upper) - keep.logical_and_(finite) + keep &= finite del finite - count = _count(keep) + count = keep.sum(dim=-2, keepdim=True) # Compute mean and standard deviation of kept values: kept_mean = torch.where(keep, numerical, 0.0).sum(-2, keepdim=True) diff --git a/sdm/processing/numerical/standardize.py b/sdm/processing/numerical/standardize.py index 8beb10fb0..376a91555 100644 --- a/sdm/processing/numerical/standardize.py +++ b/sdm/processing/numerical/standardize.py @@ -7,7 +7,6 @@ from sdm.processing import InvertibleMixin, Processor from sdm.processing.numerical._stats import ( _constant_feature_mask, - _count, _isfinite, ) @@ -42,15 +41,13 @@ def _fit( ) -> None: finite = _isfinite(table.numerical) - count = _count(finite) + 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. - # Equal to 'nanmean', which would copy its input to count values. - self.mean = finite_or_nan.nansum(-2, keepdim=True) / count + self.mean = finite_or_nan.nansum(-2, keepdim=True).div_(count) self.mean.masked_fill_(self.mean.isnan(), 0.0) - # 'finite_or_nan' is a private copy, so it can be modified in place. var = finite_or_nan.sub_(self.mean).square_().nansum(-2, keepdim=True) var /= count var.masked_fill_(var.isnan(), 0.0) @@ -64,12 +61,12 @@ def _fit( def _transform(self, table: TableTensor) -> TableTensor: dtype = table.numerical.dtype - numerical = table.numerical.sub(self.mean).div_(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.mul(self.scale).add_(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_sigma_clip.py b/test/processing/numerical/test_sigma_clip.py index 155750666..d29cd94f5 100644 --- a/test/processing/numerical/test_sigma_clip.py +++ b/test/processing/numerical/test_sigma_clip.py @@ -153,28 +153,3 @@ def test_clip_sigma_preserves_nonfinite( ), equal_nan=True, ) - - -@withCUDA -@pytest.mark.parametrize("fit_transform", [False, True]) -def test_clip_sigma_gradients( - device: torch.device, - fit_transform: bool, -) -> None: - context = torch.tensor( - [[1.0], [2.0], [4.0], [100.0]], - dtype=torch.float64, - device=device, - requires_grad=True, - ) - query = TableTensor.from_tensor(context.detach().clone()) - - def transform(values: torch.Tensor) -> torch.Tensor: - processor = ClipSigma(threshold=1.0) - table = TableTensor.from_tensor(values) - if fit_transform: - return processor.fit_transform(table).numerical - processor.fit(table) - return processor.transform(query).numerical - - assert torch.autograd.gradcheck(transform, (context,)) From 6580f8eb6463da06726f5b7ed925cc9bd3eb009b Mon Sep 17 00:00:00 2001 From: Valter Hudovernik Date: Tue, 29 Sep 2026 11:56:39 +0200 Subject: [PATCH 4/4] update --- sdm/processing/numerical/clip_soft.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/sdm/processing/numerical/clip_soft.py b/sdm/processing/numerical/clip_soft.py index b192966ad..ffd6cd84b 100644 --- a/sdm/processing/numerical/clip_soft.py +++ b/sdm/processing/numerical/clip_soft.py @@ -45,7 +45,7 @@ def _transform(self, table: TableTensor) -> TableTensor: root = unit.square().add_(1).sqrt_() unit.div_(root).masked_fill_(~_isfinite(root), 1.0) del root - clipped = (numerical.sign() * bound).mul_(unit) + clipped = unit.mul_(bound).mul_(numerical.sign()) torch.where(numerical.isnan(), numerical, clipped, out=clipped) return table.replace_blocks(numerical=clipped)