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
15 changes: 15 additions & 0 deletions sdm/_memory.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
21 changes: 17 additions & 4 deletions sdm/processing/numerical/robust_scale.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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)
Expand Down
33 changes: 29 additions & 4 deletions sdm/processing/numerical/sigma_clip.py
Original file line number Diff line number Diff line change
@@ -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

Expand Down Expand Up @@ -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 (
Expand Down
27 changes: 26 additions & 1 deletion test/processing/numerical/test_robust_scale.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
)
30 changes: 29 additions & 1 deletion test/processing/numerical/test_sigma_clip.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
)
Loading