From f120b2cb36f6ba98af6075fab6ea52e0e3f9a6b7 Mon Sep 17 00:00:00 2001 From: Valter Hudovernik Date: Mon, 28 Sep 2026 15:39:34 +0200 Subject: [PATCH 1/6] Add rank-Gaussian views to Kumo recipes --- sdm/models/kumo/tabular/recipe.py | 1 + sdm/processing/__init__.py | 2 + sdm/processing/numerical/__init__.py | 2 + sdm/processing/numerical/rank_gaussian.py | 72 ++++++++++ .../numerical/test_rank_gaussian.py | 132 ++++++++++++++++++ test/processing/test_contract.py | 1 + 6 files changed, 210 insertions(+) create mode 100644 sdm/processing/numerical/rank_gaussian.py create mode 100644 test/processing/numerical/test_rank_gaussian.py diff --git a/sdm/models/kumo/tabular/recipe.py b/sdm/models/kumo/tabular/recipe.py index 38173f3b4..29c85cad0 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(), 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..34bbb4ce4 --- /dev/null +++ b/sdm/processing/numerical/rank_gaussian.py @@ -0,0 +1,72 @@ +# 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.processing import Processor +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 fitted values and clamp to the + endpoint probabilities, keeping finite outputs even outside the range. + + 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. All fitted values are retained. + """ + + handles_stypes = frozenset({Stype.numerical}) + requires_fit = True + + def __init__(self) -> None: + super().__init__() + 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: + # [..., N, C] -> [..., C, N]; double precision keeps tail ranks open. + numerical = table.numerical.double().movedim(-1, -2) + values = numerical.masked_fill(~numerical.isfinite(), torch.inf) + values = values.sort(dim=-1).values.contiguous() + finite = values.isfinite() + count = finite.sum(dim=-1, keepdim=True) + left = torch.searchsorted(values, values, right=False) + right = torch.searchsorted(values, values, right=True) + probabilities = (left + right).double() / (2 * count.clamp_min(1)) + + # Pad missing observations with the last finite knot and its rank. + last = (count - 1).clamp_min(0) + self._values = torch.where(finite, values, values.gather(-1, last)) + self._probabilities = torch.where( + finite, probabilities, probabilities.gather(-1, last) + ) + self._values.masked_fill_(count == 0, torch.nan) + self._probabilities.masked_fill_(count == 0, torch.nan) + + def _transform(self, table: TableTensor) -> TableTensor: + numerical = table.numerical + columns = numerical.movedim(-1, -2) + n_rows = columns.size(-1) + n_fitted = self._values.size(-1) + probabilities = _batched_interp( + columns.reshape(-1, n_rows).to(self._values.dtype).contiguous(), + self._values.reshape(-1, n_fitted), + self._probabilities.reshape(-1, n_fitted), + ) + output = torch.special.ndtri(probabilities).reshape(columns.shape) + output = output.masked_fill(columns.isnan(), torch.nan) + return table.replace_blocks( + numerical=output.movedim(-1, -2).to(numerical.dtype) + ) diff --git a/test/processing/numerical/test_rank_gaussian.py b/test/processing/numerical/test_rank_gaussian.py new file mode 100644 index 000000000..f71fafcb3 --- /dev/null +++ b/test/processing/numerical/test_rank_gaussian.py @@ -0,0 +1,132 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import pytest +import torch + +from sdm import TableTensor +from sdm.processing import RankGaussian +from sdm.testing import withCUDA + + +@withCUDA +@pytest.mark.parametrize("dtype", [torch.float32, torch.float64]) +def test_mid_ranks_and_interpolated_query( + dtype: torch.dtype, device: torch.device +) -> None: + context = torch.tensor( + [[0.0], [1.0], [1.0], [2.0]], dtype=dtype, device=device + ) + processor = RankGaussian().fit(TableTensor.from_tensor(context)) + probabilities = context.new_tensor([[0.125], [0.5], [0.5], [0.875]]) + torch.testing.assert_close( + processor.transform(TableTensor.from_tensor(context)).numerical, + torch.special.ndtri(probabilities), + ) + query = context.new_tensor([[-100.0], [0.5], [1.5], [100.0], [torch.nan]]) + expected = context.new_tensor( + [[0.125], [0.3125], [0.6875], [0.875], [torch.nan]] + ) + torch.testing.assert_close( + processor.transform(TableTensor.from_tensor(query)).numerical, + torch.special.ndtri(expected), + equal_nan=True, + ) + + +@withCUDA +def test_ties_at_endpoints_use_mid_ranks(device: torch.device) -> None: + context = torch.tensor([[0.0], [0.0], [2.0], [2.0]], device=device) + output = RankGaussian().fit_transform(TableTensor.from_tensor(context)) + probabilities = context.new_tensor([[0.25], [0.25], [0.75], [0.75]]) + torch.testing.assert_close( + output.numerical, torch.special.ndtri(probabilities) + ) + + +@withCUDA +@pytest.mark.parametrize("n_rows", [1, 4]) +def test_constants_and_missing_columns( + n_rows: int, device: torch.device +) -> None: + context = torch.tensor([[7.0, torch.nan]], device=device).expand( + n_rows, -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_nonfinite_context_does_not_change_observed_ranks( + device: torch.device, +) -> None: + context = torch.tensor( + [[0.0], [1.0], [torch.nan], [1.0], [torch.inf], [-torch.inf], [2.0]], + device=device, + ) + processor = RankGaussian().fit(TableTensor.from_tensor(context)) + reference = RankGaussian().fit( + TableTensor.from_tensor(context[[0, 1, 3, 6]]) + ) + query = TableTensor.from_tensor(context) + output = processor.transform(query).numerical + torch.testing.assert_close( + output, reference.transform(query).numerical, equal_nan=True + ) + assert torch.equal(output.isnan(), context.isnan()) + assert output[~context.isnan()].isfinite().all() + + +@withCUDA +def test_query_batch_does_not_change_fitted_ranks( + device: torch.device, +) -> None: + context = torch.arange(10.0, device=device).unsqueeze(-1) + processor = RankGaussian().fit(TableTensor.from_tensor(context)) + query = context.new_tensor([[1.5], [torch.nan], [5.5]]) + expected = processor.transform(TableTensor.from_tensor(query)).numerical + extended = torch.cat([query, context.new_tensor([[-1e12], [1e12]])]) + actual = processor.transform(TableTensor.from_tensor(extended)).numerical + torch.testing.assert_close(actual[:3], expected, equal_nan=True) + + +@withCUDA +@pytest.mark.parametrize("batch_shape", [(), (2,), (2, 3)]) +def test_batched_missing_values_match_independent_columns( + batch_shape: tuple[int, ...], device: torch.device +) -> None: + context = torch.randn(*batch_shape, 17, 4, device=device) + context[..., 1, 0] = torch.nan + context[..., :3, 1] = torch.nan + context[..., :, 2] = 2.0 + context[..., :, 3] = torch.nan + query = torch.randn(*batch_shape, 7, 4, device=device) + query[..., 0, 0] = torch.nan + processor = RankGaussian().fit(TableTensor.from_tensor(context)) + output = processor.transform(TableTensor.from_tensor(query)).numerical + for train, test, actual in zip( + context.reshape(-1, 17, 4), + query.reshape(-1, 7, 4), + output.reshape(-1, 7, 4), + strict=True, + ): + for column in range(4): + reference = RankGaussian().fit( + TableTensor.from_tensor(train[:, column : column + 1]) + ) + expected = reference.transform( + TableTensor.from_tensor(test[:, column : column + 1]) + ).numerical + torch.testing.assert_close( + actual[:, column : column + 1], expected, 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()), From 6046278b8a887205f4de0cc41ffc7da5c55fe696 Mon Sep 17 00:00:00 2001 From: Jingang Qu Date: Sun, 27 Sep 2026 09:51:43 -0700 Subject: [PATCH 2/6] Keep at most max_knots knots per column in RankGaussian - Keep every fitted value as a knot while a column has at most max_knots rows, and every distinct value while it has at most max_knots of them, so outputs are unchanged in both cases. - Otherwise keep the fitted values whose mid-ranks come closest to normal quantiles evenly spaced between the extremes, which keeps outputs within about two knot spacings in normal scores. - Fit columns and transform rows in chunks within the chunk memory limit. - Keep loaded state in double precision, and support contexts and queries without rows. Signed-off-by: Jingang Qu --- sdm/processing/numerical/rank_gaussian.py | 169 ++++++++-- .../numerical/test_rank_gaussian.py | 296 +++++++++++++++++- 2 files changed, 432 insertions(+), 33 deletions(-) diff --git a/sdm/processing/numerical/rank_gaussian.py b/sdm/processing/numerical/rank_gaussian.py index 34bbb4ce4..3875f5341 100644 --- a/sdm/processing/numerical/rank_gaussian.py +++ b/sdm/processing/numerical/rank_gaussian.py @@ -2,9 +2,12 @@ # SPDX-License-Identifier: Apache-2.0 import torch +from torch import Tensor 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 @@ -14,21 +17,40 @@ class RankGaussian(Processor): 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 fitted values and clamp to the - endpoint probabilities, keeping finite outputs even outside the range. + Query probabilities interpolate between the fitted values kept as knots + and clamp to the endpoint probabilities, keeping finite outputs even + outside the range. + + Each column keeps at most ``max_knots`` knots. A column with at most + ``max_knots`` distinct finite fitted values keeps all of them, which + matches interpolating between all fitted values. A column with more + distinct values keeps the fitted values whose mid-ranks come closest to + ``max_knots`` normal quantiles evenly spaced between its extremes. + Knots then stay dense in the tails, where the normal quantile function + magnifies rank errors, and outputs stay within about two knot spacings + in normal scores of interpolating between all fitted values. 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. All fitted values are retained. + finite fitted values produce NaNs. + + Args: + max_knots: Maximum number of knots per column, at least ``2``. """ handles_stypes = frozenset({Stype.numerical}) requires_fit = True - def __init__(self) -> None: + def __init__(self, *, max_knots: int = 8192) -> None: super().__init__() - self.register_buffer("_values", torch.empty(0)) - self.register_buffer("_probabilities", torch.empty(0)) + if max_knots < 2: + raise ValueError("max_knots must be at least 2.") + self.max_knots = max_knots + self.register_buffer("_values", torch.empty(0, dtype=torch.float64)) + self.register_buffer( + "_probabilities", + torch.empty(0, dtype=torch.float64), + ) def _fit( self, @@ -36,37 +58,126 @@ def _fit( *, generator: torch.Generator | None = None, ) -> None: - # [..., N, C] -> [..., C, N]; double precision keeps tail ranks open. - numerical = table.numerical.double().movedim(-1, -2) - values = numerical.masked_fill(~numerical.isfinite(), torch.inf) + numerical = table.numerical + # Columns are fitted on their own, so chunks of columns bound the + # sorting and ranking temporaries of about ten doubles per cell. + size = split_size( + num_items=numerical.size(-1), + item_bytes=numerical[..., :1].numel() * 10 * 8, + device=numerical.device, + ) + knots = [ + self._fit_columns(chunk) for chunk in numerical.split(size, dim=-1) + ] + values, probabilities = zip(*knots, strict=True) + self._values = torch.cat(values, dim=-2) + self._probabilities = torch.cat(probabilities, dim=-2) + + def _fit_columns( + self, + numerical: Tensor, # [..., N, C] + ) -> tuple[Tensor, Tensor]: # [..., C, K] knot values and probabilities + *batch, num_rows, num_columns = numerical.shape + if num_rows == 0: + missing = numerical.new_full( + (*batch, num_columns, 1), + torch.nan, + dtype=torch.float64, + ) + return missing, missing.clone() + + # Double precision keeps tail ranks open. + columns = numerical.double().movedim(-1, -2) # [..., C, N] + finite = _isfinite(columns) + count = finite.sum(dim=-1, keepdim=True) # [..., C, 1] + values = columns.masked_fill(~finite, torch.inf) + del columns, finite values = values.sort(dim=-1).values.contiguous() - finite = values.isfinite() - count = finite.sum(dim=-1, keepdim=True) left = torch.searchsorted(values, values, right=False) right = torch.searchsorted(values, values, right=True) probabilities = (left + right).double() / (2 * count.clamp_min(1)) - - # Pad missing observations with the last finite knot and its rank. + del right last = (count - 1).clamp_min(0) - self._values = torch.where(finite, values, values.gather(-1, last)) - self._probabilities = torch.where( - finite, probabilities, probabilities.gather(-1, last) + rows = torch.arange(num_rows, device=values.device) + + if num_rows <= self.max_knots: + # Every fitted value is a knot, padded with the largest one. + positions = rows.minimum(last) + else: + # Index of each sorted value among the column's distinct 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(self.max_knots, device=values.device) + positions = torch.searchsorted( + distinct, + steps.minimum(num_distinct - 1), + ) + del distinct + # With more distinct values than knots, the rows whose mid-ranks + # come closest to normal quantiles evenly spaced from the + # smallest to the largest value. Mid-ranks never decrease along + # the sorted rows, padding included. + lower = torch.special.ndtri(probabilities[..., :1]) + upper = torch.special.ndtri(probabilities.gather(-1, last)) + quantiles = torch.special.ndtr( + lower.lerp(upper, steps.double() / (self.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 <= self.max_knots, + positions, + spaced, + ) + del left + + missing = count == 0 + return ( + values.gather(-1, positions).masked_fill_(missing, torch.nan), + probabilities.gather(-1, positions).masked_fill_( + missing, + torch.nan, + ), ) - self._values.masked_fill_(count == 0, torch.nan) - self._probabilities.masked_fill_(count == 0, torch.nan) def _transform(self, table: TableTensor) -> TableTensor: numerical = table.numerical - columns = numerical.movedim(-1, -2) - n_rows = columns.size(-1) - n_fitted = self._values.size(-1) - probabilities = _batched_interp( - columns.reshape(-1, n_rows).to(self._values.dtype).contiguous(), - self._values.reshape(-1, n_fitted), - self._probabilities.reshape(-1, n_fitted), + 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 doubles per cell. + size = split_size( + num_items=columns.size(-1), + item_bytes=columns[..., :1].numel() * 12 * 8, + device=numerical.device, ) - output = torch.special.ndtri(probabilities).reshape(columns.shape) - output = output.masked_fill(columns.isnan(), torch.nan) - return table.replace_blocks( - numerical=output.movedim(-1, -2).to(numerical.dtype) + 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 index f71fafcb3..d9d812d6d 100644 --- a/test/processing/numerical/test_rank_gaussian.py +++ b/test/processing/numerical/test_rank_gaussian.py @@ -3,10 +3,63 @@ import pytest import torch +from torch import Tensor from sdm import TableTensor from sdm.processing import RankGaussian -from sdm.testing import withCUDA +from sdm.testing import onlyCUDA, withCUDA + +# Enough knots to keep every fitted value. +ALL_KNOTS = 2**20 + + +def assert_same(actual: Tensor, expected: Tensor) -> None: + torch.testing.assert_close( + actual, expected, rtol=0, atol=0, equal_nan=True + ) + + +def few_distinct_context( + *shape: int, dtype: torch.dtype, device: torch.device +) -> Tensor: + # [*shape, 5] columns with at most four distinct finite values: ties with + # non-finite values, ties with signed zeros, two rare values between two + # large ties, a constant and a missing column. + context = torch.randint(4, (*shape, 5), device=device).to(dtype) * 2.5 + context[..., ::4, 0] = torch.nan + context[..., 1::5, 0] = torch.inf + context[..., 2::7, 0] = -torch.inf + context[..., 1] = -context[..., 1] + rows = torch.arange(shape[-1], device=device) + context[..., 2] = rows.ge(shape[-1] // 2).to(dtype) * 7.5 + context[..., 1:2, 2] = 2.5 + context[..., 2:3, 2] = 5.0 + context[..., 3] = 3.0 + context[..., 4] = torch.nan + return context + + +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 @@ -101,9 +154,10 @@ def test_query_batch_does_not_change_fitted_ranks( @withCUDA +@pytest.mark.parametrize("max_knots", [4, 8192]) @pytest.mark.parametrize("batch_shape", [(), (2,), (2, 3)]) def test_batched_missing_values_match_independent_columns( - batch_shape: tuple[int, ...], device: torch.device + batch_shape: tuple[int, ...], max_knots: int, device: torch.device ) -> None: context = torch.randn(*batch_shape, 17, 4, device=device) context[..., 1, 0] = torch.nan @@ -112,7 +166,9 @@ def test_batched_missing_values_match_independent_columns( context[..., :, 3] = torch.nan query = torch.randn(*batch_shape, 7, 4, device=device) query[..., 0, 0] = torch.nan - processor = RankGaussian().fit(TableTensor.from_tensor(context)) + processor = RankGaussian(max_knots=max_knots).fit( + TableTensor.from_tensor(context) + ) output = processor.transform(TableTensor.from_tensor(query)).numerical for train, test, actual in zip( context.reshape(-1, 17, 4), @@ -121,7 +177,7 @@ def test_batched_missing_values_match_independent_columns( strict=True, ): for column in range(4): - reference = RankGaussian().fit( + reference = RankGaussian(max_knots=max_knots).fit( TableTensor.from_tensor(train[:, column : column + 1]) ) expected = reference.transform( @@ -130,3 +186,235 @@ def test_batched_missing_values_match_independent_columns( torch.testing.assert_close( actual[:, column : column + 1], expected, equal_nan=True ) + + +@pytest.mark.parametrize("max_knots", [0, 1]) +def test_rejects_fewer_than_two_knots(max_knots: int) -> None: + with pytest.raises(ValueError, match="max_knots"): + RankGaussian(max_knots=max_knots) + + +@withCUDA +@pytest.mark.parametrize("dtype", [torch.float32, torch.float64]) +@pytest.mark.parametrize("shape", [(1,), (30,), (3, 30), (2, 3, 30)]) +def test_few_distinct_values_match_all_knots( + shape: tuple[int, ...], dtype: torch.dtype, device: torch.device +) -> None: + context = few_distinct_context(*shape, dtype=dtype, device=device) + table = TableTensor.from_tensor(context) + query = TableTensor.from_tensor(edge_query(context)) + processor = RankGaussian(max_knots=4) + reference = RankGaussian(max_knots=ALL_KNOTS) + assert_same( + actual=processor.fit_transform(table).numerical, + expected=reference.fit_transform(table).numerical, + ) + assert_same( + actual=processor.transform(query).numerical, + expected=reference.transform(query).numerical, + ) + + +@withCUDA +def test_strided_members_match_all_knots(device: torch.device) -> None: + members = few_distinct_context(6, 30, dtype=torch.float32, device=device) + context = members[::2] + query = TableTensor.from_tensor(edge_query(context)) + processor = RankGaussian(max_knots=4) + output = processor.fit_transform(TableTensor.from_tensor(context)) + reference = RankGaussian(max_knots=ALL_KNOTS) + expected = reference.fit_transform( + TableTensor.from_tensor(context.contiguous()) + ) + assert_same(actual=output.numerical, expected=expected.numerical) + assert_same( + actual=processor.transform(query).numerical, + expected=reference.transform(query).numerical, + ) + + +@withCUDA +@pytest.mark.parametrize("max_knots", [2, 16, 64]) +def test_many_distinct_values_stay_within_one_knot_spacing( + max_knots: int, device: torch.device +) -> None: + # Symmetric and skewed columns without ties, and a two-valued column. For + # 300 values without ties, these knot counts bound the error by the knot + # spacing whatever the values. + context = torch.randn(2, 300, 3, dtype=torch.float64, device=device) + context[..., 1] = context[..., 1].mul(2).exp() + context[..., 2] = context[..., 2].sign() + table = TableTensor.from_tensor(context) + processor = RankGaussian(max_knots=max_knots).fit(table) + reference = RankGaussian(max_knots=ALL_KNOTS).fit(table) + output = processor.transform(table).numerical + expected = reference.transform(table).numerical + assert_same(actual=output[..., 2], expected=expected[..., 2]) + + 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, + ) + assert_same( + actual=processor.transform(TableTensor.from_tensor(beyond)).numerical, + expected=torch.cat([lower, upper] * 3, dim=-2), + ) + + 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 +def test_nonfinite_context_does_not_change_coarse_knots( + 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=16).fit( + TableTensor.from_tensor(context) + ) + reference = RankGaussian(max_knots=16).fit(TableTensor.from_tensor(finite)) + assert_same( + actual=processor.transform(TableTensor.from_tensor(query)).numerical, + expected=reference.transform(TableTensor.from_tensor(query)).numerical, + ) + + +@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(max_knots=ALL_KNOTS) + expected = reference.fit_transform(table).numerical + spacing = (expected.max() - expected.min()) / 63 + assert (output - expected).abs().le(spacing).all() + + +@onlyCUDA +@pytest.mark.parametrize("max_knots", [16, 8192]) +def test_chunks_match_single_chunk( + max_knots: int, 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=max_knots) + 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=max_knots) + output = processor.fit_transform(table) + output_query = processor.transform(query) + + assert_same(actual=output.numerical, expected=expected.numerical) + assert_same( + actual=output_query.numerical, expected=expected_query.numerical + ) + + +@withCUDA +@pytest.mark.parametrize( + "dtype", [torch.float16, torch.bfloat16, torch.float32] +) +def test_keeps_input_dtype(dtype: torch.dtype, device: torch.device) -> None: + context = torch.randn(40, 3, device=device).to(dtype) + context[::5, 0] = torch.nan + context[:, 2] = context[:, 2].round() + query = torch.cat( + [torch.randn(8, 3, device=device).to(dtype) * 3, edge_query(context)] + ) + processor = RankGaussian(max_knots=8) + output = processor.fit_transform(TableTensor.from_tensor(context)) + transformed = processor.transform(TableTensor.from_tensor(query)) + assert output.numerical.dtype == dtype + assert transformed.numerical.dtype == dtype + + reference = RankGaussian(max_knots=8).fit( + TableTensor.from_tensor(context.double()) + ) + for inputs, actual in ((context, output), (query, transformed)): + expected = reference.transform( + TableTensor.from_tensor(inputs.double()) + ) + torch.testing.assert_close( + actual=actual.numerical, + expected=expected.numerical.to(dtype), + equal_nan=True, + ) + + +@withCUDA +def test_empty_query_gives_empty_output(device: torch.device) -> None: + context = TableTensor.from_tensor(torch.randn(10, 2, device=device)) + processor = RankGaussian(max_knots=4).fit(context) + query = TableTensor.from_tensor(torch.empty(0, 2, device=device)) + assert processor.transform(query).numerical.shape == (0, 2) + + +@withCUDA +def test_context_without_rows_gives_missing_columns( + device: torch.device, +) -> None: + context = TableTensor.from_tensor(torch.empty(0, 2, device=device)) + processor = RankGaussian().fit(context) + query = torch.tensor([[1.0, torch.nan]], device=device) + output = processor.transform(TableTensor.from_tensor(query)).numerical + assert output.isnan().all() + + +@withCUDA +def test_loaded_state_keeps_double_precision(device: torch.device) -> None: + # Distinct only in double precision, as after the Kumo recipe's cast. + context = 1 + torch.arange(100, dtype=torch.float64, device=device).mul( + 1e-12 + ) + table = TableTensor.from_tensor(context[:, None]) + processor = RankGaussian().fit(table) + loaded = RankGaussian().to(device) + loaded.load_state_dict(processor.state_dict()) + assert_same( + actual=loaded.transform(table).numerical, + expected=processor.transform(table).numerical, + ) From 35c497aff6493979805260a7b988462c6187e51b Mon Sep 17 00:00:00 2001 From: Valter Hudovernik Date: Mon, 28 Sep 2026 16:49:37 +0200 Subject: [PATCH 3/6] Trim RankGaussian tests and simplify documentation --- sdm/processing/numerical/rank_gaussian.py | 26 +- .../numerical/test_rank_gaussian.py | 255 ++++-------------- 2 files changed, 64 insertions(+), 217 deletions(-) diff --git a/sdm/processing/numerical/rank_gaussian.py b/sdm/processing/numerical/rank_gaussian.py index 3875f5341..5d654a533 100644 --- a/sdm/processing/numerical/rank_gaussian.py +++ b/sdm/processing/numerical/rank_gaussian.py @@ -17,18 +17,11 @@ class RankGaussian(Processor): 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 the fitted values kept as knots - and clamp to the endpoint probabilities, keeping finite outputs even - outside the range. - - Each column keeps at most ``max_knots`` knots. A column with at most - ``max_knots`` distinct finite fitted values keeps all of them, which - matches interpolating between all fitted values. A column with more - distinct values keeps the fitted values whose mid-ranks come closest to - ``max_knots`` normal quantiles evenly spaced between its extremes. - Knots then stay dense in the tails, where the normal quantile function - magnifies rank errors, and outputs stay within about two knot spacings - in normal scores of interpolating between all fitted values. + Query probabilities interpolate between retained fitted value-rank pairs + and clamp to the endpoint probabilities. + + Each column retains at most ``max_knots`` pairs. When needed, knots are + 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 @@ -77,14 +70,7 @@ def _fit_columns( self, numerical: Tensor, # [..., N, C] ) -> tuple[Tensor, Tensor]: # [..., C, K] knot values and probabilities - *batch, num_rows, num_columns = numerical.shape - if num_rows == 0: - missing = numerical.new_full( - (*batch, num_columns, 1), - torch.nan, - dtype=torch.float64, - ) - return missing, missing.clone() + num_rows = numerical.size(-2) # Double precision keeps tail ranks open. columns = numerical.double().movedim(-1, -2) # [..., C, N] diff --git a/test/processing/numerical/test_rank_gaussian.py b/test/processing/numerical/test_rank_gaussian.py index d9d812d6d..2ddb6d4bb 100644 --- a/test/processing/numerical/test_rank_gaussian.py +++ b/test/processing/numerical/test_rank_gaussian.py @@ -19,26 +19,6 @@ def assert_same(actual: Tensor, expected: Tensor) -> None: ) -def few_distinct_context( - *shape: int, dtype: torch.dtype, device: torch.device -) -> Tensor: - # [*shape, 5] columns with at most four distinct finite values: ties with - # non-finite values, ties with signed zeros, two rare values between two - # large ties, a constant and a missing column. - context = torch.randint(4, (*shape, 5), device=device).to(dtype) * 2.5 - context[..., ::4, 0] = torch.nan - context[..., 1::5, 0] = torch.inf - context[..., 2::7, 0] = -torch.inf - context[..., 1] = -context[..., 1] - rows = torch.arange(shape[-1], device=device) - context[..., 2] = rows.ge(shape[-1] // 2).to(dtype) * 7.5 - context[..., 1:2, 2] = 2.5 - context[..., 2:3, 2] = 5.0 - context[..., 3] = 3.0 - context[..., 4] = torch.nan - return context - - def edge_query(context: Tensor) -> Tensor: # [..., 14, C] queries at, between and beyond the fitted values. values = context.new_tensor( @@ -68,17 +48,21 @@ def test_mid_ranks_and_interpolated_query( dtype: torch.dtype, device: torch.device ) -> None: context = torch.tensor( - [[0.0], [1.0], [1.0], [2.0]], dtype=dtype, device=device + [[0.0], [0.0], [1.0], [1.0], [2.0], [2.0]], + dtype=dtype, + device=device, ) processor = RankGaussian().fit(TableTensor.from_tensor(context)) - probabilities = context.new_tensor([[0.125], [0.5], [0.5], [0.875]]) + probabilities = context.new_tensor( + [[1 / 6], [1 / 6], [0.5], [0.5], [5 / 6], [5 / 6]] + ) torch.testing.assert_close( processor.transform(TableTensor.from_tensor(context)).numerical, torch.special.ndtri(probabilities), ) query = context.new_tensor([[-100.0], [0.5], [1.5], [100.0], [torch.nan]]) expected = context.new_tensor( - [[0.125], [0.3125], [0.6875], [0.875], [torch.nan]] + [[1 / 6], [1 / 3], [2 / 3], [5 / 6], [torch.nan]] ) torch.testing.assert_close( processor.transform(TableTensor.from_tensor(query)).numerical, @@ -88,23 +72,8 @@ def test_mid_ranks_and_interpolated_query( @withCUDA -def test_ties_at_endpoints_use_mid_ranks(device: torch.device) -> None: - context = torch.tensor([[0.0], [0.0], [2.0], [2.0]], device=device) - output = RankGaussian().fit_transform(TableTensor.from_tensor(context)) - probabilities = context.new_tensor([[0.25], [0.25], [0.75], [0.75]]) - torch.testing.assert_close( - output.numerical, torch.special.ndtri(probabilities) - ) - - -@withCUDA -@pytest.mark.parametrize("n_rows", [1, 4]) -def test_constants_and_missing_columns( - n_rows: int, device: torch.device -) -> None: - context = torch.tensor([[7.0, torch.nan]], device=device).expand( - n_rows, -1 - ) +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]] @@ -120,87 +89,49 @@ def test_constants_and_missing_columns( @withCUDA -def test_nonfinite_context_does_not_change_observed_ranks( - device: torch.device, -) -> None: - context = torch.tensor( - [[0.0], [1.0], [torch.nan], [1.0], [torch.inf], [-torch.inf], [2.0]], - device=device, - ) - processor = RankGaussian().fit(TableTensor.from_tensor(context)) - reference = RankGaussian().fit( - TableTensor.from_tensor(context[[0, 1, 3, 6]]) - ) - query = TableTensor.from_tensor(context) - output = processor.transform(query).numerical - torch.testing.assert_close( - output, reference.transform(query).numerical, equal_nan=True - ) - assert torch.equal(output.isnan(), context.isnan()) - assert output[~context.isnan()].isfinite().all() - - -@withCUDA -def test_query_batch_does_not_change_fitted_ranks( - device: torch.device, -) -> None: - context = torch.arange(10.0, device=device).unsqueeze(-1) - processor = RankGaussian().fit(TableTensor.from_tensor(context)) - query = context.new_tensor([[1.5], [torch.nan], [5.5]]) - expected = processor.transform(TableTensor.from_tensor(query)).numerical - extended = torch.cat([query, context.new_tensor([[-1e12], [1e12]])]) - actual = processor.transform(TableTensor.from_tensor(extended)).numerical - torch.testing.assert_close(actual[:3], expected, equal_nan=True) - - -@withCUDA -@pytest.mark.parametrize("max_knots", [4, 8192]) -@pytest.mark.parametrize("batch_shape", [(), (2,), (2, 3)]) def test_batched_missing_values_match_independent_columns( - batch_shape: tuple[int, ...], max_knots: int, device: torch.device + device: torch.device, ) -> None: - context = torch.randn(*batch_shape, 17, 4, device=device) + 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(*batch_shape, 7, 4, device=device) + query = torch.randn(2, 3, 5, 4, device=device) query[..., 0, 0] = torch.nan - processor = RankGaussian(max_knots=max_knots).fit( - TableTensor.from_tensor(context) - ) + 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.reshape(-1, 17, 4), - query.reshape(-1, 7, 4), - output.reshape(-1, 7, 4), + context.flatten(0, 1), + query.flatten(0, 1), + output.flatten(0, 1), strict=True, ): - for column in range(4): - reference = RankGaussian(max_knots=max_knots).fit( - TableTensor.from_tensor(train[:, column : column + 1]) - ) - expected = reference.transform( - TableTensor.from_tensor(test[:, column : column + 1]) - ).numerical - torch.testing.assert_close( - actual[:, column : column + 1], expected, equal_nan=True - ) - - -@pytest.mark.parametrize("max_knots", [0, 1]) -def test_rejects_fewer_than_two_knots(max_knots: int) -> None: - with pytest.raises(ValueError, match="max_knots"): - RankGaussian(max_knots=max_knots) + 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 -@pytest.mark.parametrize("dtype", [torch.float32, torch.float64]) -@pytest.mark.parametrize("shape", [(1,), (30,), (3, 30), (2, 3, 30)]) -def test_few_distinct_values_match_all_knots( - shape: tuple[int, ...], dtype: torch.dtype, device: torch.device -) -> None: - context = few_distinct_context(*shape, dtype=dtype, device=device) +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) @@ -216,40 +147,19 @@ def test_few_distinct_values_match_all_knots( @withCUDA -def test_strided_members_match_all_knots(device: torch.device) -> None: - members = few_distinct_context(6, 30, dtype=torch.float32, device=device) - context = members[::2] - query = TableTensor.from_tensor(edge_query(context)) - processor = RankGaussian(max_knots=4) - output = processor.fit_transform(TableTensor.from_tensor(context)) - reference = RankGaussian(max_knots=ALL_KNOTS) - expected = reference.fit_transform( - TableTensor.from_tensor(context.contiguous()) - ) - assert_same(actual=output.numerical, expected=expected.numerical) - assert_same( - actual=processor.transform(query).numerical, - expected=reference.transform(query).numerical, - ) - - -@withCUDA -@pytest.mark.parametrize("max_knots", [2, 16, 64]) +@pytest.mark.parametrize("max_knots", [2, 16]) def test_many_distinct_values_stay_within_one_knot_spacing( max_knots: int, device: torch.device ) -> None: - # Symmetric and skewed columns without ties, and a two-valued column. For - # 300 values without ties, these knot counts bound the error by the knot + # For columns without ties, these knot counts bound the error by the knot # spacing whatever the values. - context = torch.randn(2, 300, 3, dtype=torch.float64, device=device) + context = torch.randn(2, 300, 2, dtype=torch.float64, device=device) context[..., 1] = context[..., 1].mul(2).exp() - context[..., 2] = context[..., 2].sign() table = TableTensor.from_tensor(context) processor = RankGaussian(max_knots=max_knots).fit(table) reference = RankGaussian(max_knots=ALL_KNOTS).fit(table) output = processor.transform(table).numerical expected = reference.transform(table).numerical - assert_same(actual=output[..., 2], expected=expected[..., 2]) lower = expected.amin(dim=-2, keepdim=True) upper = expected.amax(dim=-2, keepdim=True) @@ -285,7 +195,9 @@ def test_many_distinct_values_stay_within_one_knot_spacing( @withCUDA -def test_nonfinite_context_does_not_change_coarse_knots( +@pytest.mark.parametrize("max_knots", [16, ALL_KNOTS]) +def test_nonfinite_context_does_not_change_fitted_ranks( + max_knots: int, device: torch.device, ) -> None: finite = torch.randn(200, 2, dtype=torch.float64, device=device) @@ -293,14 +205,19 @@ def test_nonfinite_context_does_not_change_coarse_knots( 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=16).fit( + processor = RankGaussian(max_knots=max_knots).fit( TableTensor.from_tensor(context) ) - reference = RankGaussian(max_knots=16).fit(TableTensor.from_tensor(finite)) + reference = RankGaussian(max_knots=max_knots).fit( + TableTensor.from_tensor(finite) + ) + output = processor.transform(TableTensor.from_tensor(query)).numerical assert_same( - actual=processor.transform(TableTensor.from_tensor(query)).numerical, + actual=output, expected=reference.transform(TableTensor.from_tensor(query)).numerical, ) + assert torch.equal(output.isnan(), query.isnan()) + assert output[~query.isnan()].isfinite().all() @withCUDA @@ -320,10 +237,7 @@ def test_values_next_to_ties_stay_within_one_knot_spacing( @onlyCUDA -@pytest.mark.parametrize("max_knots", [16, 8192]) -def test_chunks_match_single_chunk( - max_knots: int, monkeypatch: pytest.MonkeyPatch -) -> None: +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 @@ -339,12 +253,12 @@ def test_chunks_match_single_chunk( # Chunks run without autograd: with torch.inference_mode(): - processor = RankGaussian(max_knots=max_knots) + 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=max_knots) + processor = RankGaussian(max_knots=16) output = processor.fit_transform(table) output_query = processor.transform(query) @@ -354,65 +268,12 @@ def test_chunks_match_single_chunk( ) -@withCUDA -@pytest.mark.parametrize( - "dtype", [torch.float16, torch.bfloat16, torch.float32] -) -def test_keeps_input_dtype(dtype: torch.dtype, device: torch.device) -> None: - context = torch.randn(40, 3, device=device).to(dtype) - context[::5, 0] = torch.nan - context[:, 2] = context[:, 2].round() - query = torch.cat( - [torch.randn(8, 3, device=device).to(dtype) * 3, edge_query(context)] - ) - processor = RankGaussian(max_knots=8) - output = processor.fit_transform(TableTensor.from_tensor(context)) - transformed = processor.transform(TableTensor.from_tensor(query)) - assert output.numerical.dtype == dtype - assert transformed.numerical.dtype == dtype - - reference = RankGaussian(max_knots=8).fit( - TableTensor.from_tensor(context.double()) - ) - for inputs, actual in ((context, output), (query, transformed)): - expected = reference.transform( - TableTensor.from_tensor(inputs.double()) - ) - torch.testing.assert_close( - actual=actual.numerical, - expected=expected.numerical.to(dtype), - equal_nan=True, - ) - - -@withCUDA -def test_empty_query_gives_empty_output(device: torch.device) -> None: - context = TableTensor.from_tensor(torch.randn(10, 2, device=device)) - processor = RankGaussian(max_knots=4).fit(context) - query = TableTensor.from_tensor(torch.empty(0, 2, device=device)) - assert processor.transform(query).numerical.shape == (0, 2) - - -@withCUDA -def test_context_without_rows_gives_missing_columns( - device: torch.device, -) -> None: - context = TableTensor.from_tensor(torch.empty(0, 2, device=device)) - processor = RankGaussian().fit(context) - query = torch.tensor([[1.0, torch.nan]], device=device) - output = processor.transform(TableTensor.from_tensor(query)).numerical - assert output.isnan().all() - - -@withCUDA -def test_loaded_state_keeps_double_precision(device: torch.device) -> None: +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, device=device).mul( - 1e-12 - ) + context = 1 + torch.arange(100, dtype=torch.float64).mul(1e-12) table = TableTensor.from_tensor(context[:, None]) processor = RankGaussian().fit(table) - loaded = RankGaussian().to(device) + loaded = RankGaussian() loaded.load_state_dict(processor.state_dict()) assert_same( actual=loaded.transform(table).numerical, From 5ddabb04b8585d914ef192e64f9d776666e855da Mon Sep 17 00:00:00 2001 From: RBendias Date: Mon, 28 Sep 2026 15:40:26 +0000 Subject: [PATCH 4/6] Address RankGaussian review feedback --- sdm/models/kumo/tabular/recipe.py | 2 +- sdm/processing/numerical/rank_gaussian.py | 153 +++++++++--------- .../numerical/test_rank_gaussian.py | 67 ++++---- 3 files changed, 117 insertions(+), 105 deletions(-) diff --git a/sdm/models/kumo/tabular/recipe.py b/sdm/models/kumo/tabular/recipe.py index 29c85cad0..05ac30d5c 100644 --- a/sdm/models/kumo/tabular/recipe.py +++ b/sdm/models/kumo/tabular/recipe.py @@ -22,7 +22,7 @@ def numerical_processor() -> sp.Sequential: sp.RobustScale(), sp.ClipSoft(3.0), ], - sp.RankGaussian(), + sp.RankGaussian(max_knots=8192), method="round_robin", ), sp.ClipSigma(threshold=4.0), diff --git a/sdm/processing/numerical/rank_gaussian.py b/sdm/processing/numerical/rank_gaussian.py index 5d654a533..b49cb563a 100644 --- a/sdm/processing/numerical/rank_gaussian.py +++ b/sdm/processing/numerical/rank_gaussian.py @@ -2,7 +2,6 @@ # SPDX-License-Identifier: Apache-2.0 import torch -from torch import Tensor from sdm import Stype, TableTensor from sdm._memory import split_size @@ -20,7 +19,7 @@ class RankGaussian(Processor): Query probabilities interpolate between retained fitted value-rank pairs and clamp to the endpoint probabilities. - Each column retains at most ``max_knots`` pairs. When needed, knots are + 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 @@ -28,15 +27,16 @@ class RankGaussian(Processor): finite fitted values produce NaNs. Args: - max_knots: Maximum number of knots per column, at least ``2``. + 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 = 8192) -> None: + def __init__(self, *, max_knots: int | None = None) -> None: super().__init__() - if max_knots < 2: + 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, dtype=torch.float64)) @@ -59,81 +59,80 @@ def _fit( item_bytes=numerical[..., :1].numel() * 10 * 8, device=numerical.device, ) - knots = [ - self._fit_columns(chunk) for chunk in numerical.split(size, dim=-1) - ] - values, probabilities = zip(*knots, strict=True) - self._values = torch.cat(values, dim=-2) - self._probabilities = torch.cat(probabilities, dim=-2) - - def _fit_columns( - self, - numerical: Tensor, # [..., N, C] - ) -> tuple[Tensor, Tensor]: # [..., C, K] knot values and probabilities - num_rows = numerical.size(-2) - - # Double precision keeps tail ranks open. - columns = numerical.double().movedim(-1, -2) # [..., C, N] - finite = _isfinite(columns) - count = finite.sum(dim=-1, keepdim=True) # [..., C, 1] - values = columns.masked_fill(~finite, torch.inf) - del columns, finite - values = values.sort(dim=-1).values.contiguous() - left = torch.searchsorted(values, values, right=False) - right = torch.searchsorted(values, values, right=True) - probabilities = (left + right).double() / (2 * count.clamp_min(1)) - del right - last = (count - 1).clamp_min(0) - rows = torch.arange(num_rows, device=values.device) - - if num_rows <= self.max_knots: - # Every fitted value is a knot, padded with the largest one. - positions = rows.minimum(last) - else: - # Index of each sorted value among the column's distinct 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(self.max_knots, device=values.device) - positions = torch.searchsorted( - distinct, - steps.minimum(num_distinct - 1), - ) - del distinct - # With more distinct values than knots, the rows whose mid-ranks - # come closest to normal quantiles evenly spaced from the - # smallest to the largest value. Mid-ranks never decrease along - # the sorted rows, padding included. - lower = torch.special.ndtri(probabilities[..., :1]) - upper = torch.special.ndtri(probabilities.gather(-1, last)) - quantiles = torch.special.ndtr( - lower.lerp(upper, steps.double() / (self.max_knots - 1)) + 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.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) ) - above = torch.searchsorted(probabilities, quantiles).minimum(last) - below = (above - 1).clamp_(min=0) - closer = (quantiles - probabilities.gather(-1, below)) < ( - probabilities.gather(-1, above) - quantiles + 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, + ), + ) ) - spaced = torch.where(closer, below, above) - spaced[..., :1] = 0 - spaced[..., -1:] = last - positions = torch.where( - num_distinct <= self.max_knots, - positions, - spaced, - ) - del left - missing = count == 0 - return ( - 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 diff --git a/test/processing/numerical/test_rank_gaussian.py b/test/processing/numerical/test_rank_gaussian.py index 2ddb6d4bb..75f26dbee 100644 --- a/test/processing/numerical/test_rank_gaussian.py +++ b/test/processing/numerical/test_rank_gaussian.py @@ -9,15 +9,6 @@ from sdm.processing import RankGaussian from sdm.testing import onlyCUDA, withCUDA -# Enough knots to keep every fitted value. -ALL_KNOTS = 2**20 - - -def assert_same(actual: Tensor, expected: Tensor) -> None: - torch.testing.assert_close( - actual, expected, rtol=0, atol=0, equal_nan=True - ) - def edge_query(context: Tensor) -> Tensor: # [..., 14, C] queries at, between and beyond the fitted values. @@ -43,13 +34,10 @@ def edge_query(context: Tensor) -> Tensor: @withCUDA -@pytest.mark.parametrize("dtype", [torch.float32, torch.float64]) -def test_mid_ranks_and_interpolated_query( - dtype: torch.dtype, device: torch.device -) -> None: +def test_mid_ranks_and_interpolated_query(device: torch.device) -> None: context = torch.tensor( [[0.0], [0.0], [1.0], [1.0], [2.0], [2.0]], - dtype=dtype, + dtype=torch.float64, device=device, ) processor = RankGaussian().fit(TableTensor.from_tensor(context)) @@ -135,14 +123,20 @@ def test_few_distinct_values_match_all_knots(device: torch.device) -> None: table = TableTensor.from_tensor(context) query = TableTensor.from_tensor(edge_query(context)) processor = RankGaussian(max_knots=4) - reference = RankGaussian(max_knots=ALL_KNOTS) - assert_same( + 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, ) - assert_same( + torch.testing.assert_close( actual=processor.transform(query).numerical, expected=reference.transform(query).numerical, + rtol=0, + atol=0, + equal_nan=True, ) @@ -157,7 +151,7 @@ def test_many_distinct_values_stay_within_one_knot_spacing( context[..., 1] = context[..., 1].mul(2).exp() table = TableTensor.from_tensor(context) processor = RankGaussian(max_knots=max_knots).fit(table) - reference = RankGaussian(max_knots=ALL_KNOTS).fit(table) + reference = RankGaussian().fit(table) output = processor.transform(table).numerical expected = reference.transform(table).numerical @@ -180,9 +174,12 @@ def test_many_distinct_values_stay_within_one_knot_spacing( ], dim=-2, ) - assert_same( + 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 @@ -195,9 +192,9 @@ def test_many_distinct_values_stay_within_one_knot_spacing( @withCUDA -@pytest.mark.parametrize("max_knots", [16, ALL_KNOTS]) +@pytest.mark.parametrize("max_knots", [16, None]) def test_nonfinite_context_does_not_change_fitted_ranks( - max_knots: int, + max_knots: int | None, device: torch.device, ) -> None: finite = torch.randn(200, 2, dtype=torch.float64, device=device) @@ -212,9 +209,12 @@ def test_nonfinite_context_does_not_change_fitted_ranks( TableTensor.from_tensor(finite) ) output = processor.transform(TableTensor.from_tensor(query)).numerical - assert_same( + 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() @@ -230,7 +230,7 @@ def test_values_next_to_ties_stay_within_one_knot_spacing( )[:, None] table = TableTensor.from_tensor(context) output = RankGaussian(max_knots=64).fit_transform(table).numerical - reference = RankGaussian(max_knots=ALL_KNOTS) + reference = RankGaussian() expected = reference.fit_transform(table).numerical spacing = (expected.max() - expected.min()) / 63 assert (output - expected).abs().le(spacing).all() @@ -262,9 +262,19 @@ def test_chunks_match_single_chunk(monkeypatch: pytest.MonkeyPatch) -> None: output = processor.fit_transform(table) output_query = processor.transform(query) - assert_same(actual=output.numerical, expected=expected.numerical) - assert_same( - actual=output_query.numerical, expected=expected_query.numerical + 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, ) @@ -275,7 +285,10 @@ def test_loaded_state_keeps_double_precision() -> None: processor = RankGaussian().fit(table) loaded = RankGaussian() loaded.load_state_dict(processor.state_dict()) - assert_same( + torch.testing.assert_close( actual=loaded.transform(table).numerical, expected=processor.transform(table).numerical, + rtol=0, + atol=0, + equal_nan=True, ) From c4a7248377e6fb3f1ac70876c2af1f87133f036b Mon Sep 17 00:00:00 2001 From: Valter Hudovernik Date: Tue, 29 Sep 2026 14:52:13 +0200 Subject: [PATCH 5/6] address comments --- sdm/processing/numerical/rank_gaussian.py | 15 ++++++--------- test/processing/numerical/test_rank_gaussian.py | 2 +- 2 files changed, 7 insertions(+), 10 deletions(-) diff --git a/sdm/processing/numerical/rank_gaussian.py b/sdm/processing/numerical/rank_gaussian.py index b49cb563a..53d71b3f4 100644 --- a/sdm/processing/numerical/rank_gaussian.py +++ b/sdm/processing/numerical/rank_gaussian.py @@ -39,11 +39,8 @@ def __init__(self, *, max_knots: int | None = None) -> None: 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, dtype=torch.float64)) - self.register_buffer( - "_probabilities", - torch.empty(0, dtype=torch.float64), - ) + self.register_buffer("_values", torch.empty(0)) + self.register_buffer("_probabilities", torch.empty(0)) def _fit( self, @@ -53,10 +50,10 @@ def _fit( ) -> None: numerical = table.numerical # Columns are fitted on their own, so chunks of columns bound the - # sorting and ranking temporaries of about ten doubles per cell. + # sorting and ranking temporaries of about ten values per cell. size = split_size( num_items=numerical.size(-1), - item_bytes=numerical[..., :1].numel() * 10 * 8, + item_bytes=numerical[..., :1].numel() * 10 * numerical.itemsize, device=numerical.device, ) knots = [] @@ -141,10 +138,10 @@ def _transform(self, table: TableTensor) -> TableTensor: 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 doubles per cell. + # interpolation temporaries of about a dozen values per cell. size = split_size( num_items=columns.size(-1), - item_bytes=columns[..., :1].numel() * 12 * 8, + item_bytes=columns[..., :1].numel() * 12 * values.itemsize, device=numerical.device, ) for rows, out in zip( diff --git a/test/processing/numerical/test_rank_gaussian.py b/test/processing/numerical/test_rank_gaussian.py index 75f26dbee..8f1472fd0 100644 --- a/test/processing/numerical/test_rank_gaussian.py +++ b/test/processing/numerical/test_rank_gaussian.py @@ -284,7 +284,7 @@ def test_loaded_state_keeps_double_precision() -> None: table = TableTensor.from_tensor(context[:, None]) processor = RankGaussian().fit(table) loaded = RankGaussian() - loaded.load_state_dict(processor.state_dict()) + loaded.load_state_dict(processor.state_dict(), assign=True) torch.testing.assert_close( actual=loaded.transform(table).numerical, expected=processor.transform(table).numerical, From 3dcc9c70b257666662fa89300632c91493965b46 Mon Sep 17 00:00:00 2001 From: Valter Hudovernik Date: Tue, 29 Sep 2026 15:07:09 +0200 Subject: [PATCH 6/6] remove precision issues --- sdm/processing/numerical/rank_gaussian.py | 7 +++-- .../numerical/test_rank_gaussian.py | 30 +++++++++++++------ 2 files changed, 26 insertions(+), 11 deletions(-) diff --git a/sdm/processing/numerical/rank_gaussian.py b/sdm/processing/numerical/rank_gaussian.py index 53d71b3f4..1907570ff 100644 --- a/sdm/processing/numerical/rank_gaussian.py +++ b/sdm/processing/numerical/rank_gaussian.py @@ -49,18 +49,19 @@ def _fit( 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 * numerical.itemsize, + 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.movedim(-1, -2) # [..., C, N] + 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) @@ -70,6 +71,8 @@ def _fit( 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) diff --git a/test/processing/numerical/test_rank_gaussian.py b/test/processing/numerical/test_rank_gaussian.py index 8f1472fd0..2eedb2cb4 100644 --- a/test/processing/numerical/test_rank_gaussian.py +++ b/test/processing/numerical/test_rank_gaussian.py @@ -34,27 +34,39 @@ def edge_query(context: Tensor) -> Tensor: @withCUDA -def test_mid_ranks_and_interpolated_query(device: torch.device) -> None: +@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=torch.float64, + dtype=dtype, device=device, ) - processor = RankGaussian().fit(TableTensor.from_tensor(context)) - probabilities = context.new_tensor( - [[1 / 6], [1 / 6], [0.5], [0.5], [5 / 6], [5 / 6]] + 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), + torch.special.ndtri(probabilities).to(dtype), ) query = context.new_tensor([[-100.0], [0.5], [1.5], [100.0], [torch.nan]]) - expected = context.new_tensor( - [[1 / 6], [1 / 3], [2 / 3], [5 / 6], [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), + torch.special.ndtri(expected).to(dtype), equal_nan=True, )