diff --git a/sdm/models/kumo/tabular/recipe.py b/sdm/models/kumo/tabular/recipe.py index 38173f3b4..05ac30d5c 100644 --- a/sdm/models/kumo/tabular/recipe.py +++ b/sdm/models/kumo/tabular/recipe.py @@ -22,6 +22,7 @@ def numerical_processor() -> sp.Sequential: sp.RobustScale(), sp.ClipSoft(3.0), ], + sp.RankGaussian(max_knots=8192), method="round_robin", ), sp.ClipSigma(threshold=4.0), diff --git a/sdm/processing/__init__.py b/sdm/processing/__init__.py index 400a11d9a..ad88657d6 100644 --- a/sdm/processing/__init__.py +++ b/sdm/processing/__init__.py @@ -32,6 +32,7 @@ ImputeMean, PowerTransform, QuantileTransform, + RankGaussian, Standardize, RobustScale, FlipSign, @@ -76,6 +77,7 @@ "ImputeMean", "PowerTransform", "QuantileTransform", + "RankGaussian", "Standardize", "RobustScale", "FlipSign", diff --git a/sdm/processing/numerical/__init__.py b/sdm/processing/numerical/__init__.py index 6f670deb4..fbbba0ec8 100644 --- a/sdm/processing/numerical/__init__.py +++ b/sdm/processing/numerical/__init__.py @@ -11,6 +11,7 @@ from sdm.processing.numerical.impute import ImputeMean from sdm.processing.numerical.power import PowerTransform from sdm.processing.numerical.quantile import QuantileTransform +from sdm.processing.numerical.rank_gaussian import RankGaussian from sdm.processing.numerical.standardize import Standardize from sdm.processing.numerical.robust_scale import RobustScale from sdm.processing.numerical.flip_sign import FlipSign @@ -27,6 +28,7 @@ "ImputeMean", "PowerTransform", "QuantileTransform", + "RankGaussian", "Standardize", "RobustScale", "FlipSign", diff --git a/sdm/processing/numerical/rank_gaussian.py b/sdm/processing/numerical/rank_gaussian.py new file mode 100644 index 000000000..1907570ff --- /dev/null +++ b/sdm/processing/numerical/rank_gaussian.py @@ -0,0 +1,168 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import torch + +from sdm import Stype, TableTensor +from sdm._memory import split_size +from sdm.processing import Processor +from sdm.processing.numerical._stats import _isfinite +from sdm.processing.numerical.quantile import _batched_interp + + +class RankGaussian(Processor): + """Map interpolated empirical mid-ranks to standard normal quantiles. + + Each fitted value receives probability ``(L + R) / (2 * N)``, where + ``L`` and ``R`` count finite fitted values strictly below and at or below + it, and ``N`` is the number of finite fitted values. Ties share a rank. + Query probabilities interpolate between retained fitted value-rank pairs + and clamp to the endpoint probabilities. + + When ``max_knots`` is set, each column retains at most that many pairs, + selected at evenly spaced normal quantiles between the fitted extremes. + + NaN and infinite values are ignored during fitting. NaNs are preserved + during transformation. Constant columns map to zero; columns without + finite fitted values produce NaNs. + + Args: + max_knots: Maximum number of knots per column, at least ``2``. If + ``None``, retain all fitted values. + """ + + handles_stypes = frozenset({Stype.numerical}) + requires_fit = True + + def __init__(self, *, max_knots: int | None = None) -> None: + super().__init__() + if max_knots is not None and max_knots < 2: + raise ValueError("max_knots must be at least 2.") + self.max_knots = max_knots + self.register_buffer("_values", torch.empty(0)) + self.register_buffer("_probabilities", torch.empty(0)) + + def _fit( + self, + table: TableTensor, + *, + generator: torch.Generator | None = None, + ) -> None: + numerical = table.numerical + dtype = torch.promote_types(numerical.dtype, torch.float32) + # Columns are fitted on their own, so chunks of columns bound the + # sorting and ranking temporaries of about ten values per cell. + size = split_size( + num_items=numerical.size(-1), + item_bytes=numerical[..., :1].numel() * 10 * dtype.itemsize, + device=numerical.device, + ) + knots = [] + for chunk in numerical.split(size, dim=-1): + num_rows = chunk.size(-2) + max_knots = num_rows if self.max_knots is None else self.max_knots + columns = chunk.to(dtype).movedim(-1, -2) # [..., C, N] + finite = _isfinite(columns) + count = finite.sum(dim=-1, keepdim=True) # [..., C, 1] + values = columns.masked_fill(~finite, torch.inf) + values = values.sort(dim=-1).values.contiguous() + left = torch.searchsorted(values, values, right=False) + right = torch.searchsorted(values, values, right=True) + probabilities = (left + right).to(values.dtype) / ( + 2 * count.clamp_min(1) + ) + # Rounding to one would give an infinite normal quantile. + probabilities.clamp_(max=1 - torch.finfo(dtype).eps) + last = (count - 1).clamp_min(0) + rows = torch.arange(num_rows, device=values.device) + + if num_rows <= max_knots: + positions = rows.minimum(last) + else: + # Index among distinct column values; a value starts a new + # one at its first occurrence. + distinct = (left == rows).cumsum(dim=-1).sub_(1) + num_distinct = distinct.gather(-1, last) + 1 + # Knots at the first occurrence of each distinct value, padded + # with the largest one. + steps = torch.arange(max_knots, device=values.device) + positions = torch.searchsorted( + distinct, + steps.minimum(num_distinct - 1), + ) + # With more distinct values than knots, select rows whose + # mid-ranks are closest to normal quantiles spaced evenly + # between the extremes. Mid-ranks never decrease along rows. + lower = torch.special.ndtri(probabilities[..., :1]) + upper = torch.special.ndtri(probabilities.gather(-1, last)) + quantiles = torch.special.ndtr( + lower.lerp( + upper, + steps.to(probabilities.dtype) / (max_knots - 1), + ) + ) + above = torch.searchsorted(probabilities, quantiles).minimum( + last + ) + below = (above - 1).clamp_(min=0) + closer = (quantiles - probabilities.gather(-1, below)) < ( + probabilities.gather(-1, above) - quantiles + ) + spaced = torch.where(closer, below, above) + spaced[..., :1] = 0 + spaced[..., -1:] = last + positions = torch.where( + num_distinct <= max_knots, + positions, + spaced, + ) + + missing = count == 0 + knots.append( + ( + values.gather(-1, positions).masked_fill_( + missing, torch.nan + ), + probabilities.gather(-1, positions).masked_fill_( + missing, + torch.nan, + ), + ) + ) + + values, probabilities = zip(*knots, strict=True) + self._values = torch.cat(values, dim=-2) + self._probabilities = torch.cat(probabilities, dim=-2) + + def _transform(self, table: TableTensor) -> TableTensor: + numerical = table.numerical + columns = numerical.movedim(-1, -2) # [..., C, R] + values = self._values.flatten(end_dim=-2) + probabilities = self._probabilities.flatten(end_dim=-2) + output = torch.empty_like(numerical) + # Rows transform on their own, so chunks of rows bound the + # interpolation temporaries of about a dozen values per cell. + size = split_size( + num_items=columns.size(-1), + item_bytes=columns[..., :1].numel() * 12 * values.itemsize, + device=numerical.device, + ) + for rows, out in zip( + columns.split(size, dim=-1), + output.split(size, dim=-2), + strict=True, + ): + quantiles = _batched_interp( + rows.flatten(end_dim=-2).to(values.dtype).contiguous(), + values, + probabilities, + ) + normal = torch.special.ndtri(quantiles).reshape(rows.shape) + out.copy_(normal.masked_fill_(rows.isnan(), torch.nan).mT) + return table.replace_blocks(numerical=output) + + def __repr__(self, *, indent: int = 0) -> str: + return ( + f"{' ' * indent}{self.__class__.__name__}(" + f"max_knots={self.max_knots})" + ) diff --git a/test/processing/numerical/test_rank_gaussian.py b/test/processing/numerical/test_rank_gaussian.py new file mode 100644 index 000000000..2eedb2cb4 --- /dev/null +++ b/test/processing/numerical/test_rank_gaussian.py @@ -0,0 +1,306 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import pytest +import torch +from torch import Tensor + +from sdm import TableTensor +from sdm.processing import RankGaussian +from sdm.testing import onlyCUDA, withCUDA + + +def edge_query(context: Tensor) -> Tensor: + # [..., 14, C] queries at, between and beyond the fitted values. + values = context.new_tensor( + [ + -torch.inf, + -100.0, + -7.5, + -4.0, + -0.0, + 0.0, + 1.0, + 2.5, + 3.0, + 6.0, + 7.5, + 100.0, + torch.inf, + torch.nan, + ] + ) + return values[:, None].expand(*context.shape[:-2], -1, context.size(-1)) + + +@withCUDA +@pytest.mark.parametrize( + "dtype", [torch.float16, torch.bfloat16, torch.float32, torch.float64] +) +@pytest.mark.parametrize("max_knots", [None, 3]) +def test_mid_ranks_and_interpolated_query( + device: torch.device, dtype: torch.dtype, max_knots: int | None +) -> None: + context = torch.tensor( + [[0.0], [0.0], [1.0], [1.0], [2.0], [2.0]], + dtype=dtype, + device=device, + ) + processor = RankGaussian(max_knots=max_knots).fit( + TableTensor.from_tensor(context) + ) + probabilities = torch.tensor( + [[1 / 6], [1 / 6], [0.5], [0.5], [5 / 6], [5 / 6]], + dtype=torch.float64, + device=device, + ) + torch.testing.assert_close( + processor.transform(TableTensor.from_tensor(context)).numerical, + torch.special.ndtri(probabilities).to(dtype), + ) + query = context.new_tensor([[-100.0], [0.5], [1.5], [100.0], [torch.nan]]) + expected = torch.tensor( + [[1 / 6], [1 / 3], [2 / 3], [5 / 6], [torch.nan]], + dtype=torch.float64, + device=device, + ) + torch.testing.assert_close( + processor.transform(TableTensor.from_tensor(query)).numerical, + torch.special.ndtri(expected).to(dtype), + equal_nan=True, + ) + + +@withCUDA +def test_constants_and_missing_columns(device: torch.device) -> None: + context = torch.tensor([[7.0, torch.nan]], device=device).expand(4, -1) + processor = RankGaussian().fit(TableTensor.from_tensor(context)) + query = context.new_tensor( + [[7.0, 8.0], [torch.nan, torch.nan], [100.0, 3.0]] + ) + expected = context.new_tensor( + [[0.0, torch.nan], [torch.nan, torch.nan], [0.0, torch.nan]] + ) + torch.testing.assert_close( + processor.transform(TableTensor.from_tensor(query)).numerical, + expected, + equal_nan=True, + ) + + +@withCUDA +def test_batched_missing_values_match_independent_columns( + device: torch.device, +) -> None: + context = torch.randn(2, 3, 9, 4, device=device) + context[..., 1, 0] = torch.nan + context[..., :3, 1] = torch.nan + context[..., :, 2] = 2.0 + context[..., :, 3] = torch.nan + query = torch.randn(2, 3, 5, 4, device=device) + query[..., 0, 0] = torch.nan + processor = RankGaussian(max_knots=4).fit(TableTensor.from_tensor(context)) + output = processor.transform(TableTensor.from_tensor(query)).numerical + for train, test, actual in zip( + context.flatten(0, 1), + query.flatten(0, 1), + output.flatten(0, 1), + strict=True, + ): + reference = RankGaussian(max_knots=4).fit( + TableTensor.from_tensor(train) + ) + expected = reference.transform(TableTensor.from_tensor(test)).numerical + torch.testing.assert_close(actual, expected, equal_nan=True) + + +@withCUDA +def test_few_distinct_values_match_all_knots(device: torch.device) -> None: + context = torch.tensor( + [ + [0, 0, 0, torch.nan], + [0, -0.0, 0, torch.nan], + [0, -1, 0, torch.nan], + [1, -1, 0, torch.nan], + [1, -1, 2, torch.nan], + [1, -2, 4, torch.nan], + [2, -2, 6, torch.nan], + [2, -2, 6, torch.nan], + [2, -3, 6, torch.nan], + [3, -3, 6, torch.nan], + ], + dtype=torch.float64, + device=device, + ) + table = TableTensor.from_tensor(context) + query = TableTensor.from_tensor(edge_query(context)) + processor = RankGaussian(max_knots=4) + reference = RankGaussian() + torch.testing.assert_close( + actual=processor.fit_transform(table).numerical, + expected=reference.fit_transform(table).numerical, + rtol=0, + atol=0, + equal_nan=True, + ) + torch.testing.assert_close( + actual=processor.transform(query).numerical, + expected=reference.transform(query).numerical, + rtol=0, + atol=0, + equal_nan=True, + ) + + +@withCUDA +@pytest.mark.parametrize("max_knots", [2, 16]) +def test_many_distinct_values_stay_within_one_knot_spacing( + max_knots: int, device: torch.device +) -> None: + # For columns without ties, these knot counts bound the error by the knot + # spacing whatever the values. + context = torch.randn(2, 300, 2, dtype=torch.float64, device=device) + context[..., 1] = context[..., 1].mul(2).exp() + table = TableTensor.from_tensor(context) + processor = RankGaussian(max_knots=max_knots).fit(table) + reference = RankGaussian().fit(table) + output = processor.transform(table).numerical + expected = reference.transform(table).numerical + + lower = expected.amin(dim=-2, keepdim=True) + upper = expected.amax(dim=-2, keepdim=True) + spacing = (upper - lower) / (max_knots - 1) + assert (output - expected).abs().le(spacing).all() + + # The extreme fitted values keep their scores, and queries beyond clamp: + smallest = context.amin(dim=-2, keepdim=True) + largest = context.amax(dim=-2, keepdim=True) + beyond = torch.cat( + [ + smallest, + largest, + smallest - 1, + largest + 1, + torch.full_like(smallest, -torch.inf), + torch.full_like(largest, torch.inf), + ], + dim=-2, + ) + torch.testing.assert_close( + actual=processor.transform(TableTensor.from_tensor(beyond)).numerical, + expected=torch.cat([lower, upper] * 3, dim=-2), + rtol=0, + atol=0, + equal_nan=True, + ) + + values = context.sort(dim=-2).values + midpoints = values[..., :-1, :].lerp(values[..., 1:, :], 0.5) + queries = torch.cat([values, midpoints, beyond], dim=-2) + scores = processor.transform( + TableTensor.from_tensor(queries.sort(dim=-2).values) + ).numerical + assert scores.diff(dim=-2).ge(0).all() + + +@withCUDA +@pytest.mark.parametrize("max_knots", [16, None]) +def test_nonfinite_context_does_not_change_fitted_ranks( + max_knots: int | None, + device: torch.device, +) -> None: + finite = torch.randn(200, 2, dtype=torch.float64, device=device) + nonfinite = finite.new_tensor([torch.nan, torch.inf, -torch.inf]) + context = torch.cat([finite, nonfinite[:, None].repeat(20, 2)]) + context = context[torch.randperm(context.size(0), device=device)] + query = torch.cat([context, edge_query(context)]) + processor = RankGaussian(max_knots=max_knots).fit( + TableTensor.from_tensor(context) + ) + reference = RankGaussian(max_knots=max_knots).fit( + TableTensor.from_tensor(finite) + ) + output = processor.transform(TableTensor.from_tensor(query)).numerical + torch.testing.assert_close( + actual=output, + expected=reference.transform(TableTensor.from_tensor(query)).numerical, + rtol=0, + atol=0, + equal_nan=True, + ) + assert torch.equal(output.isnan(), query.isnan()) + assert output[~query.isnan()].isfinite().all() + + +@withCUDA +def test_values_next_to_ties_stay_within_one_knot_spacing( + device: torch.device, +) -> None: + # 500 zeros, then 500 distinct values. + context = torch.cat([torch.zeros(500), torch.arange(1.0, 501.0)]).to( + dtype=torch.float64, device=device + )[:, None] + table = TableTensor.from_tensor(context) + output = RankGaussian(max_knots=64).fit_transform(table).numerical + reference = RankGaussian() + expected = reference.fit_transform(table).numerical + spacing = (expected.max() - expected.min()) / 63 + assert (output - expected).abs().le(spacing).all() + + +@onlyCUDA +def test_chunks_match_single_chunk(monkeypatch: pytest.MonkeyPatch) -> None: + context = torch.randn(2, 300, 5, device="cuda") + context[..., 1] = context[..., 1].mul(2).round() + context[:, ::7, 0] = torch.nan + context[:, 3, 1] = torch.inf + context[..., 3] = 1.0 + context[..., 4] = torch.nan + table = TableTensor.from_tensor(context) + query = torch.randn(2, 40, 5, device="cuda") * 3 + query[:, ::6] = torch.nan + query = TableTensor.from_tensor( + torch.cat([query, edge_query(query)], dim=-2) + ) + + # Chunks run without autograd: + with torch.inference_mode(): + processor = RankGaussian(max_knots=16) + expected = processor.fit_transform(table) + expected_query = processor.transform(query) + # Chunks of a single column in fit and a single row in transform: + monkeypatch.setenv("SDM_CHUNK_MEMORY_FRACTION", "1e-12") + processor = RankGaussian(max_knots=16) + output = processor.fit_transform(table) + output_query = processor.transform(query) + + torch.testing.assert_close( + actual=output.numerical, + expected=expected.numerical, + rtol=0, + atol=0, + equal_nan=True, + ) + torch.testing.assert_close( + actual=output_query.numerical, + expected=expected_query.numerical, + rtol=0, + atol=0, + equal_nan=True, + ) + + +def test_loaded_state_keeps_double_precision() -> None: + # Distinct only in double precision, as after the Kumo recipe's cast. + context = 1 + torch.arange(100, dtype=torch.float64).mul(1e-12) + table = TableTensor.from_tensor(context[:, None]) + processor = RankGaussian().fit(table) + loaded = RankGaussian() + loaded.load_state_dict(processor.state_dict(), assign=True) + torch.testing.assert_close( + actual=loaded.transform(table).numerical, + expected=processor.transform(table).numerical, + rtol=0, + atol=0, + equal_nan=True, + ) diff --git a/test/processing/test_contract.py b/test/processing/test_contract.py index 43692e47b..642a3d69a 100644 --- a/test/processing/test_contract.py +++ b/test/processing/test_contract.py @@ -126,6 +126,7 @@ def _make_processor_pair( sp.QuantileTransform(n_quantiles=4, subsample=None), ), ProcessorCase(sp.Standardize()), + ProcessorCase(sp.RankGaussian()), ProcessorCase(sp.RobustScale()), ProcessorCase(sp.FlipSign()), ProcessorCase(sp.DropConstantColumns()),