Skip to content
Merged
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
7 changes: 7 additions & 0 deletions sdm/processing/numerical/_stats.py
Original file line number Diff line number Diff line change
@@ -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,
Expand Down
17 changes: 8 additions & 9 deletions sdm/processing/numerical/clip_soft.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@

from sdm import Stype, TableTensor
from sdm.processing import Processor
from sdm.processing.numerical._stats import _isfinite


class ClipSoft(Processor):
Expand Down Expand Up @@ -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:
Expand Down
41 changes: 22 additions & 19 deletions sdm/processing/numerical/power.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand All @@ -139,17 +143,17 @@ 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
log_jacobian = magnitude_log.copysign(inp).nansum(dim=-2, keepdim=True)
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
Expand Down Expand Up @@ -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,
Expand All @@ -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
Expand All @@ -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,
Expand Down
7 changes: 4 additions & 3 deletions sdm/processing/numerical/robust_scale.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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(
Expand All @@ -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:
Expand Down
30 changes: 19 additions & 11 deletions sdm/processing/numerical/sigma_clip.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@

from sdm import Stype, TableTensor
from sdm.processing import Processor
from sdm.processing.numerical._stats import _isfinite


class ClipSigma(Processor):
Expand Down Expand Up @@ -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)
Comment thread
ValterH marked this conversation as resolved.

# 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
Expand Down
23 changes: 12 additions & 11 deletions sdm/processing/numerical/standardize.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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:
Expand Down
20 changes: 20 additions & 0 deletions test/processing/numerical/test_standardize.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
Loading