From 359c723f4fb7dc8357ba1c3281dd71db5af34cd2 Mon Sep 17 00:00:00 2001 From: Valter Hudovernik Date: Tue, 29 Sep 2026 14:09:28 +0200 Subject: [PATCH] Chunk expensive numerical operations (#1016) Co-authored-by: Jingang Qu Co-authored-by: Cedric Lorenz --- sdm/_memory.py | 15 +++++++++ sdm/processing/numerical/robust_scale.py | 21 +++++++++--- sdm/processing/numerical/sigma_clip.py | 33 ++++++++++++++++--- .../processing/numerical/test_robust_scale.py | 27 ++++++++++++++- test/processing/numerical/test_sigma_clip.py | 30 ++++++++++++++++- 5 files changed, 116 insertions(+), 10 deletions(-) diff --git a/sdm/_memory.py b/sdm/_memory.py index 47a735a3d..af8828c47 100644 --- a/sdm/_memory.py +++ b/sdm/_memory.py @@ -17,3 +17,18 @@ def chunk_memory_limit(device: torch.device) -> int: * torch.cuda.get_per_process_memory_fraction(device) * float(os.getenv("SDM_CHUNK_MEMORY_FRACTION", "0.05")) ) + + +def split_size(num_items: int, item_bytes: int, device: torch.device) -> int: + r"""Split size of balanced chunks of items of ``item_bytes`` bytes each. + + On CUDA devices, chunks fit :func:`chunk_memory_limit`, unless a single + item exceeds the limit. Autograd keeps the memory of every chunk, so one + chunk holds all items elsewhere or while gradients are enabled. + """ + if device.type != "cuda" or torch.is_grad_enabled() or item_bytes == 0: + return max(num_items, 1) + limit = max(chunk_memory_limit(device), 1) + capacity = max(limit // item_bytes, 1) + num_chunks = max(-(-num_items // capacity), 1) + return max(-(-num_items // num_chunks), 1) diff --git a/sdm/processing/numerical/robust_scale.py b/sdm/processing/numerical/robust_scale.py index f2ba25506..2525147db 100644 --- a/sdm/processing/numerical/robust_scale.py +++ b/sdm/processing/numerical/robust_scale.py @@ -4,6 +4,7 @@ import torch from sdm import Stype, TableTensor +from sdm._memory import split_size from sdm.processing import InvertibleMixin, Processor from sdm.processing.numerical._stats import _isfinite @@ -51,10 +52,22 @@ def _fit( quantile_input = finite_or_nan.to( dtype=torch.promote_types(numerical.dtype, torch.float32), ) - lower, median, upper = quantile_input.nanquantile( - quantile_input.new_tensor([q_low, 0.5, q_high]), - dim=-2, - keepdim=True, + q = quantile_input.new_tensor([q_low, 0.5, q_high]) + # 'nanquantile' allocates several copies of its input, including + # int64 sort indices, which chunks of the independent columns bound. + size = split_size( + num_items=quantile_input.size(-1), + item_bytes=quantile_input[..., :1].numel() + * 4 + * torch.int64.itemsize, + device=quantile_input.device, + ) + lower, median, upper = torch.cat( + [ + chunk.nanquantile(q, dim=-2, keepdim=True) + for chunk in quantile_input.split(size, dim=-1) + ], + dim=-1, ) self.median = median.to(dtype=numerical.dtype) scale = torch.where(lower == upper, 1.0, upper - lower) diff --git a/sdm/processing/numerical/sigma_clip.py b/sdm/processing/numerical/sigma_clip.py index 511b10df5..56d951a37 100644 --- a/sdm/processing/numerical/sigma_clip.py +++ b/sdm/processing/numerical/sigma_clip.py @@ -1,9 +1,12 @@ # SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 +import math + 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 @@ -86,10 +89,32 @@ def _fit( def _transform(self, table: TableTensor) -> TableTensor: numerical = table.numerical - log_abs = numerical.abs().log1p() - clipped = torch.maximum(-log_abs + self.lower_bound, numerical) - numerical = torch.minimum(log_abs + self.upper_bound, clipped) - return table.replace_blocks(numerical=numerical) + dtype = torch.promote_types(numerical.dtype, self.lower_bound.dtype) + if torch.is_grad_enabled(): + log_abs = numerical.abs().log1p().to(dtype) + clipped = torch.maximum(self.lower_bound - log_abs, numerical) + clipped = torch.minimum(log_abs + self.upper_bound, clipped) + return table.replace_blocks(numerical=clipped) + + out = torch.empty_like(numerical, dtype=dtype) + # Chunks of rows bound the temporary besides the output. + size = split_size( + num_items=numerical.size(-2), + item_bytes=math.prod(out.size()[:-2]) + * out.size(-1) + * out.element_size(), + device=out.device, + ) + for inp, clipped in zip( + numerical.split(size, dim=-2), + out.split(size, dim=-2), + strict=True, + ): + log_abs = inp.abs().log1p_().to(dtype) + torch.sub(self.lower_bound, log_abs, out=clipped) + torch.maximum(clipped, inp, out=clipped) + torch.minimum(log_abs.add_(self.upper_bound), clipped, out=clipped) + return table.replace_blocks(numerical=out) def __repr__(self, *, indent: int = 0) -> str: return ( diff --git a/test/processing/numerical/test_robust_scale.py b/test/processing/numerical/test_robust_scale.py index 0d92a88cd..191dfd8d6 100644 --- a/test/processing/numerical/test_robust_scale.py +++ b/test/processing/numerical/test_robust_scale.py @@ -6,7 +6,7 @@ from sdm import TableTensor from sdm.processing import RobustScale -from sdm.testing import withCUDA +from sdm.testing import onlyCUDA, withCUDA @withCUDA @@ -194,3 +194,28 @@ def test_robust_scale_fits_half_precision( assert out.numerical.dtype == dtype expected = torch.tensor([[-1.0], [0.0], [1.0]], device=device) torch.testing.assert_close(out.numerical.float(), expected) + + +@onlyCUDA +def test_robust_scale_column_chunks_match_single_chunk( + monkeypatch: pytest.MonkeyPatch, +) -> None: + inp = torch.randn(2, 50, 5, device="cuda") + inp[:, ::7, 0] = float("nan") + inp[:, 3, 1] = float("inf") + table = TableTensor.from_tensor(inp) + + # Chunks run without autograd: + with torch.inference_mode(): + expected = RobustScale().fit_transform(table) + # Chunks of a single column: + monkeypatch.setenv("SDM_CHUNK_MEMORY_FRACTION", "1e-12") + out = RobustScale().fit_transform(table) + + torch.testing.assert_close( + out.numerical, + expected.numerical, + rtol=0, + atol=0, + equal_nan=True, + ) diff --git a/test/processing/numerical/test_sigma_clip.py b/test/processing/numerical/test_sigma_clip.py index d29cd94f5..d48b1df78 100644 --- a/test/processing/numerical/test_sigma_clip.py +++ b/test/processing/numerical/test_sigma_clip.py @@ -6,7 +6,7 @@ from sdm import TableTensor from sdm.processing import ClipSigma -from sdm.testing import withCUDA +from sdm.testing import onlyCUDA, withCUDA @withCUDA @@ -153,3 +153,31 @@ def test_clip_sigma_preserves_nonfinite( ), equal_nan=True, ) + + +@onlyCUDA +def test_clip_sigma_row_chunks_match_single_chunk( + monkeypatch: pytest.MonkeyPatch, +) -> None: + inp = torch.randn(2, 50, 3, device="cuda", dtype=torch.float64) + inp[:, ::7, 0] = float("nan") + inp[:, 3, 1] = float("inf") + processor = ClipSigma(threshold=1.0) + processor.fit(TableTensor.from_tensor(inp)) + query = TableTensor.from_tensor(inp.float() * 3) + + # Passes run without autograd: + with torch.inference_mode(): + expected = processor.transform(query) + # Passes of a single row: + monkeypatch.setenv("SDM_CHUNK_MEMORY_FRACTION", "1e-12") + out = processor.transform(query) + + assert out.numerical.dtype == torch.float64 + torch.testing.assert_close( + out.numerical, + expected.numerical, + rtol=0, + atol=0, + equal_nan=True, + )