From 10aafebba4cf6535c0b98e5acf9ceb05ae2476e0 Mon Sep 17 00:00:00 2001 From: Jingang Qu Date: Sun, 27 Sep 2026 01:28:30 -0700 Subject: [PATCH 1/9] Share the device chunk memory limit - Move the `SDM_CHUNK_MEMORY_FRACTION` budget of attention and TabFM cell embedding chunks into one helper, `sdm._memory.chunk_memory_limit`. - Expose the automatic attention batch size limit as `TransformerBlock.auto_batch_size_limit`, so callers can plan passes that align with its chunks. Signed-off-by: Jingang Qu --- sdm/_memory.py | 19 +++++++++ sdm/models/tabfm/cell_embedding.py | 10 ++--- sdm/nn/attention.py | 62 +++++++++++++++++------------- 3 files changed, 57 insertions(+), 34 deletions(-) create mode 100644 sdm/_memory.py diff --git a/sdm/_memory.py b/sdm/_memory.py new file mode 100644 index 000000000..47a735a3d --- /dev/null +++ b/sdm/_memory.py @@ -0,0 +1,19 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import os + +import torch + + +def chunk_memory_limit(device: torch.device) -> int: + r"""Bytes one chunk of a chunked operation may occupy on a CUDA device. + + The limit is the ``SDM_CHUNK_MEMORY_FRACTION`` (default ``0.05``) share of + the device memory available to this process. + """ + return int( + torch.cuda.get_device_properties(device).total_memory + * torch.cuda.get_per_process_memory_fraction(device) + * float(os.getenv("SDM_CHUNK_MEMORY_FRACTION", "0.05")) + ) diff --git a/sdm/models/tabfm/cell_embedding.py b/sdm/models/tabfm/cell_embedding.py index f2a8d64b7..f2c086408 100644 --- a/sdm/models/tabfm/cell_embedding.py +++ b/sdm/models/tabfm/cell_embedding.py @@ -18,13 +18,14 @@ # ruff: noqa: D101, D102 import math -import os from typing import Any, Literal import torch from torch import Tensor from torch.nn import Linear +from sdm._memory import chunk_memory_limit + class CellEmbedding(torch.nn.Module): def __init__( @@ -132,12 +133,7 @@ def forward( + bias.numel() * bias.element_size() ) - memory_limit = int( - torch.cuda.get_device_properties(x.device).total_memory - * torch.cuda.get_per_process_memory_fraction(x.device) - * float(os.getenv("SDM_CHUNK_MEMORY_FRACTION", "0.05")) - ) - memory_limit -= fixed_bytes + memory_limit = chunk_memory_limit(x.device) - fixed_bytes batch_size_limit = memory_limit // max(bytes_per_example, 1) batch_size_limit = max(batch_size_limit, 1) diff --git a/sdm/nn/attention.py b/sdm/nn/attention.py index b2020d9ff..42190df1f 100644 --- a/sdm/nn/attention.py +++ b/sdm/nn/attention.py @@ -4,7 +4,6 @@ """Attention modules for structured tensor models.""" import math -import os from typing import Any, Literal, overload import torch @@ -12,6 +11,7 @@ from torch import Tensor from torch.nn import Linear +from sdm._memory import chunk_memory_limit from sdm.cache import KVCacheEntry from sdm.nn import QueryScaling @@ -533,32 +533,22 @@ def forward( ) if batch_size_limit == "auto": - batch_size_limit = None - if query.is_cuda: - key_value_length: int | None = None - if isinstance(key_value, Tensor): - key_value_length = key_value.size(-2) - elif isinstance(key_value, KVCacheEntry): - key_value_length = key_value.key.size(-3) - - bytes_per_example = self.peak_bytes_per_example( - element_size=torch.empty( - size=(), - dtype=torch.get_autocast_dtype(query.device.type), - ).element_size() - if torch.is_autocast_enabled(query.device.type) - else query.element_size(), - query_length=query.size(-2), - key_value_length=key_value_length, - ) - - memory_limit = int( - torch.cuda.get_device_properties(query.device).total_memory - * torch.cuda.get_per_process_memory_fraction(query.device) - * float(os.getenv("SDM_CHUNK_MEMORY_FRACTION", "0.05")) - ) - batch_size_limit = memory_limit // max(bytes_per_example, 1) - batch_size_limit = max(batch_size_limit, 1) + key_value_length: int | None = None + if isinstance(key_value, Tensor): + key_value_length = key_value.size(-2) + elif isinstance(key_value, KVCacheEntry): + key_value_length = key_value.key.size(-3) + + batch_size_limit = self.auto_batch_size_limit( + device=query.device, + element_size=torch.get_autocast_dtype( + query.device.type + ).itemsize + if torch.is_autocast_enabled(query.device.type) + else query.element_size(), + query_length=query.size(-2), + key_value_length=key_value_length, + ) batch_size_limit = min(batch_size_limit or 65_535, 65_535) @@ -728,6 +718,24 @@ def peak_bytes_per_example( r""":meta private:""" # noqa: D415 return 0 + def auto_batch_size_limit( + self, + device: torch.device, + element_size: int, + query_length: int, + key_value_length: int | None = None, + ) -> int: + r""":meta private:""" # noqa: D415 + if device.type != "cuda": + return 65_535 + bytes_per_example = self.peak_bytes_per_example( + element_size=element_size, + query_length=query_length, + key_value_length=key_value_length, + ) + limit = chunk_memory_limit(device) // max(bytes_per_example, 1) + return min(max(limit, 1), 65_535) + def _batch_shape( query: Tensor, From 9a21b34c9829bc6a2c37932ed9b0b47290cb5fc0 Mon Sep 17 00:00:00 2001 From: Jingang Qu Date: Sun, 27 Sep 2026 01:28:30 -0700 Subject: [PATCH 2/9] Reduce recipe memory on wide contexts - Find finite values without an `abs()` copy and count them without an int64 copy of the mask, and compute Standardize, ClipSigma and PowerTransform statistics with fewer full-size temporaries. - Find the Yeo-Johnson bounds before allocating the PowerTransform workspaces. - Compute RobustScale quantiles over column chunks and ClipSigma transforms over row chunks within the chunk memory limit. - Transform in Standardize, ClipSoft and RobustScale with fewer temporaries. - Select evenly spaced ensemble members as views, select shared DropConstantColumns groups before re-stacking members, and select Choice members one option at a time. Signed-off-by: Jingang Qu --- sdm/_memory.py | 14 +++++ sdm/ensemble.py | 19 +++--- sdm/processing/common/choice.py | 22 +++---- sdm/processing/numerical/_stats.py | 18 ++++++ sdm/processing/numerical/clip_soft.py | 17 +++--- sdm/processing/numerical/constant.py | 30 +++++----- sdm/processing/numerical/power.py | 42 ++++++++----- sdm/processing/numerical/robust_scale.py | 28 ++++++--- sdm/processing/numerical/sigma_clip.py | 59 ++++++++++++++----- sdm/processing/numerical/standardize.py | 26 ++++---- .../processing/numerical/test_robust_scale.py | 27 ++++++++- test/processing/numerical/test_sigma_clip.py | 30 +++++++++- test/processing/numerical/test_standardize.py | 20 +++++++ 13 files changed, 257 insertions(+), 95 deletions(-) diff --git a/sdm/_memory.py b/sdm/_memory.py index 47a735a3d..68806ea50 100644 --- a/sdm/_memory.py +++ b/sdm/_memory.py @@ -17,3 +17,17 @@ 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`. 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(): + return max(num_items, 1) + limit = max(chunk_memory_limit(device), 1) + num_chunks = max(-(-num_items * item_bytes // limit), 1) + return max(-(-num_items // num_chunks), 1) diff --git a/sdm/ensemble.py b/sdm/ensemble.py index 70e629a99..3b4dc6962 100644 --- a/sdm/ensemble.py +++ b/sdm/ensemble.py @@ -212,15 +212,20 @@ def _select_members(self, member_ids: Sequence[int]) -> Self: group = self._groups[group_id] selected_positions = tuple(positions) group_ids[group_id] = new_group_id + start = selected_positions[0] + step = ( + selected_positions[1] - start + if len(selected_positions) > 1 + else 1 + ) + stop = start + step * len(selected_positions) if selected_positions == tuple(range(group.size(0))): groups.append(group) - elif len(selected_positions) == 1: - groups.append( - cast( - TableTensor, - group.narrow(0, selected_positions[0], 1), - ) - ) + elif step > 0 and selected_positions == tuple( + range(start, stop, step) + ): + # Evenly spaced members are selected as a view, not a copy. + groups.append(group[start:stop:step]) else: groups.append( cast( diff --git a/sdm/processing/common/choice.py b/sdm/processing/common/choice.py index c4660235f..25da3bac8 100644 --- a/sdm/processing/common/choice.py +++ b/sdm/processing/common/choice.py @@ -1,6 +1,7 @@ # SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 +from collections.abc import Iterator from typing import Literal, cast import torch @@ -79,14 +80,13 @@ def _draw_option_ids( def _tables_by_option( self, ensemble_table: EnsembleTable, - ) -> dict[int, EnsembleTable]: + ) -> Iterator[tuple[int, EnsembleTable]]: member_ids_by_option: dict[int, list[int]] = {} for member_id, option_id in enumerate(self._option_ids): member_ids_by_option.setdefault(option_id, []).append(member_id) - return { - option_id: ensemble_table[member_ids] - for option_id, member_ids in sorted(member_ids_by_option.items()) - } + # Select lazily, since selecting members may copy them. + for option_id, member_ids in sorted(member_ids_by_option.items()): + yield option_id, ensemble_table[member_ids] def _gather_outputs( self, @@ -124,7 +124,7 @@ def _fit_ensemble( ensemble_table, generator=generator, ) - for option_id, table in self._tables_by_option(ensemble_table).items(): + for option_id, table in self._tables_by_option(ensemble_table): self.options[option_id].fit_ensemble(table, generator=generator) def _fit_transform_ensemble( @@ -142,9 +142,7 @@ def _fit_transform_ensemble( table, generator=generator, ) - for option_id, table in self._tables_by_option( - ensemble_table - ).items() + for option_id, table in self._tables_by_option(ensemble_table) } return self._gather_outputs(ensemble_table, outputs) @@ -155,9 +153,7 @@ def _transform_ensemble( self._check_num_members(ensemble_table) outputs = { option_id: self.options[option_id].transform_ensemble(table) - for option_id, table in self._tables_by_option( - ensemble_table - ).items() + for option_id, table in self._tables_by_option(ensemble_table) } return self._gather_outputs(ensemble_table, outputs) @@ -167,7 +163,7 @@ def _inverse_transform_ensemble( ) -> EnsembleTable: self._check_num_members(ensemble_table) outputs = {} - for option_id, table in self._tables_by_option(ensemble_table).items(): + for option_id, table in self._tables_by_option(ensemble_table): processor = self.options[option_id] if not isinstance(processor, EnsembleInvertibleMixin): raise TypeError( diff --git a/sdm/processing/numerical/_stats.py b/sdm/processing/numerical/_stats.py index 002d435c7..df3d84adf 100644 --- a/sdm/processing/numerical/_stats.py +++ b/sdm/processing/numerical/_stats.py @@ -1,10 +1,28 @@ # SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 +import math + import torch from torch import Tensor +def _isfinite(x: Tensor) -> Tensor: + # Equal to 'x.isfinite()', which allocates 'x.abs()' on the way. + return x.gt(-math.inf).logical_and_(x.lt(math.inf)) + + +def _count(mask: Tensor) -> Tensor: + # [..., N, C] -> [..., 1, C] int64 number of true values per column. + # Summing bool first casts all of 'mask' to int64. Sum blocks of 255 rows + # as uint8 instead, which cannot overflow. + num_blocks = mask.size(-2) // 255 + blocks = mask[..., : num_blocks * 255, :].unflatten(-2, (num_blocks, 255)) + count = blocks.view(torch.uint8).sum(-2, dtype=torch.uint8) + remainder = mask[..., num_blocks * 255 :, :] + return count.sum(-2, keepdim=True) + remainder.sum(-2, keepdim=True) + + def _constant_feature_mask( var: Tensor, mean: Tensor, diff --git a/sdm/processing/numerical/clip_soft.py b/sdm/processing/numerical/clip_soft.py index 716008655..2e9ea02f9 100644 --- a/sdm/processing/numerical/clip_soft.py +++ b/sdm/processing/numerical/clip_soft.py @@ -7,6 +7,7 @@ from sdm import Stype, TableTensor from sdm.processing import Processor +from sdm.processing.numerical._stats import _isfinite class ClipSoft(Processor): @@ -39,15 +40,13 @@ def __init__(self, max_absolute_value: float) -> None: def _transform(self, table: TableTensor) -> TableTensor: numerical = table.numerical bound = self.max_absolute_value - ratio = (numerical / bound).abs() - squared = 1 + ratio.square() - unit = torch.where( - squared.isfinite(), - ratio / squared.sqrt(), - 1.0, - ) - clipped = numerical.sign() * bound * unit - clipped = torch.where(numerical.isnan(), numerical, clipped) + unit = numerical.div(bound).abs_() + root = unit.square().add_(1).sqrt_() + # 'root' is finite exactly where '1 + unit ** 2' is. + unit.div_(root).masked_fill_(~_isfinite(root), 1.0) + del root + clipped = numerical.sign().mul_(bound).mul_(unit) + torch.where(numerical.isnan(), numerical, clipped, out=clipped) return table.replace_blocks(numerical=clipped) def __repr__(self, *, indent: int = 0) -> str: diff --git a/sdm/processing/numerical/constant.py b/sdm/processing/numerical/constant.py index 4cbaaeb2f..4fe2873cf 100644 --- a/sdm/processing/numerical/constant.py +++ b/sdm/processing/numerical/constant.py @@ -180,21 +180,6 @@ def _transform_ensemble( "ensemble members before transform." ) - if sum( - group.size(0) for group in ensemble_table._iter_groups() - ) == len(ensemble_table): - tables = [ - self._select_columns( - ensemble_table[member_id], - kept_indices, - ) - for member_id, kept_indices in enumerate(self._kept_indices) - ] - return ensemble_table.replace_tables( - tables=tables, - member_table_ids=range(len(ensemble_table)), - ) - member_ids_by_kept_indices: dict[tuple[int, ...], list[int]] = {} for member_id, kept_indices in enumerate(self._kept_indices): member_ids_by_kept_indices.setdefault(kept_indices, []).append( @@ -210,6 +195,21 @@ def _transform_ensemble( ] ) + if sum( + group.size(0) for group in ensemble_table._iter_groups() + ) == len(ensemble_table): + tables = [ + self._select_columns( + ensemble_table[member_id], + kept_indices, + ) + for member_id, kept_indices in enumerate(self._kept_indices) + ] + return ensemble_table.replace_tables( + tables=tables, + member_table_ids=range(len(ensemble_table)), + ) + outputs: dict[tuple[int, ...], EnsembleTable] = {} for kept_indices, member_ids in member_ids_by_kept_indices.items(): selected = ensemble_table[member_ids] diff --git a/sdm/processing/numerical/power.py b/sdm/processing/numerical/power.py index c62424226..2fdffe9b1 100644 --- a/sdm/processing/numerical/power.py +++ b/sdm/processing/numerical/power.py @@ -8,7 +8,11 @@ from sdm import Stype, TableTensor from sdm.processing import InvertibleMixin, Processor -from sdm.processing.numerical._stats import _constant_feature_mask +from sdm.processing.numerical._stats import ( + _constant_feature_mask, + _count, + _isfinite, +) # Keep GPU execution batched; adaptive per-column stopping would resynchronize. # For float32 overflow-safe bounds, 44 golden steps reaches ~1.48e-8. @@ -68,7 +72,7 @@ def _yeojohnson_inverse_transform(inp: Tensor, lambdas: Tensor) -> Tensor: def _yeojohnson_bounds(inp: Tensor) -> tuple[Tensor, Tensor]: missing = inp.isnan() - max_abs = inp.abs().nan_to_num(nan=0.0).amax(dim=-2, keepdim=True) + max_abs = inp.abs().nan_to_num_(nan=0.0).amax(dim=-2, keepdim=True) log1p_max_x = (20 * max_abs).log1p() log1p_max_x = torch.where( max_abs == 0, @@ -124,8 +128,11 @@ def _yeojohnson_log_likelihood( exponents=exponents, out=transformed, ) - mean = transformed.nanmean(dim=-2, keepdim=True) - variance = transformed.sub_(mean).square_().nanmean(dim=-2, keepdim=True) + # 'transformed' is NaN exactly where 'inp' is, so this equals 'nanmean', + # which would copy its input to count values. + mean = transformed.nansum(dim=-2, keepdim=True) / count + variance = transformed.sub_(mean).square_().nansum(dim=-2, keepdim=True) + variance /= count tiny = torch.finfo(inp.dtype).tiny valid = variance.isfinite() & (variance >= tiny) loglike = variance.log_().mul_(-count / 2) @@ -139,6 +146,11 @@ def _optimize_lambdas( *, count: Tensor, ) -> Tensor: + # Find the bounds before allocating the workspaces. + left, right = _yeojohnson_bounds(inp) + left = left.masked_fill(constant_features, 1.0) + right = right.masked_fill(constant_features, 1.0) + # Reuse full-table workspaces throughout the golden-section search. magnitude_log = inp.abs().log1p_() positive = inp >= 0 @@ -146,10 +158,6 @@ def _optimize_lambdas( exponents = torch.empty_like(inp) transformed = torch.empty_like(inp) - left, right = _yeojohnson_bounds(inp) - left = left.masked_fill(constant_features, 1.0) - right = right.masked_fill(constant_features, 1.0) - invphi = (math.sqrt(5) - 1) / 2 span = (right - left).mul_(invphi) c = right - span @@ -252,13 +260,15 @@ def _fit( *, generator: torch.Generator | None = None, ) -> None: - finite = table.numerical.isfinite() + finite = _isfinite(table.numerical) finite_or_nan = table.numerical.masked_fill(~finite, torch.nan) - count = finite.sum(dim=-2, keepdim=True) + count = _count(finite) - mean = finite_or_nan.nanmean(-2, keepdim=True) + # Equal to 'nanmean', which would copy its input to count values: + mean = finite_or_nan.nansum(-2, keepdim=True) / count mean.masked_fill_(mean.isnan(), 0.0) - var = (finite_or_nan - mean).square().nanmean(-2, keepdim=True) + var = finite_or_nan.sub(mean).square_().nansum(-2, keepdim=True) + var /= count var.masked_fill_(var.isnan(), 0.0) constant_features = _constant_feature_mask( var, @@ -283,9 +293,11 @@ def _fit( if self.standardize: transformed = _yeojohnson_transform(finite_or_nan, self.lambdas) - mean = transformed.nanmean(dim=-2, keepdim=True) + del finite_or_nan + mean = transformed.nansum(dim=-2, keepdim=True) / count mean.masked_fill_(mean.isnan(), 0.0) - var = (transformed - mean).square().nanmean(dim=-2, keepdim=True) + var = transformed.sub_(mean).square_().nansum(-2, keepdim=True) + var /= count var.masked_fill_(var.isnan(), 0.0) scale = var.sqrt() scale[_constant_feature_mask(var, mean, count)] = 1.0 @@ -312,7 +324,7 @@ def _inverse_transform(self, table: TableTensor) -> TableTensor: # Above the fitted upper bound the inverse diverges, either to # infinity or, past the asymptote, to NaN. - diverged = ~inverse.isfinite() & ~unscaled.isnan() + diverged = ~_isfinite(inverse) & ~unscaled.isnan() return table.replace_blocks( numerical=torch.where( diverged, diff --git a/sdm/processing/numerical/robust_scale.py b/sdm/processing/numerical/robust_scale.py index f9dd58b4d..75ca08e97 100644 --- a/sdm/processing/numerical/robust_scale.py +++ b/sdm/processing/numerical/robust_scale.py @@ -4,7 +4,9 @@ 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 class RobustScale(Processor, InvertibleMixin): @@ -44,27 +46,39 @@ def _fit( generator: torch.Generator | None = None, ) -> None: numerical = table.numerical - finite_or_nan = numerical.masked_fill(~numerical.isfinite(), torch.nan) + finite_or_nan = numerical.masked_fill(~_isfinite(numerical), torch.nan) q_low, q_high = (value / 100.0 for value in self.quantile_range) # 'nanquantile' requires single or double precision input. quantile_input = finite_or_nan.to( 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) self.scale = scale.to(dtype=numerical.dtype) def _transform(self, table: TableTensor) -> TableTensor: - numerical = (table.numerical - self.median) / self.scale + numerical = table.numerical.sub(self.median).div_(self.scale) return table.replace_blocks(numerical=numerical) def _inverse_transform(self, table: TableTensor) -> TableTensor: - numerical = table.numerical * self.scale + self.median + numerical = table.numerical.mul(self.scale).add_(self.median) return table.replace_blocks(numerical=numerical) def __repr__(self, *, indent: int = 0) -> str: diff --git a/sdm/processing/numerical/sigma_clip.py b/sdm/processing/numerical/sigma_clip.py index 81444cffa..137c2e799 100644 --- a/sdm/processing/numerical/sigma_clip.py +++ b/sdm/processing/numerical/sigma_clip.py @@ -1,10 +1,14 @@ # 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 _count, _isfinite class ClipSigma(Processor): @@ -43,30 +47,37 @@ def _fit( generator: torch.Generator | None = None, ) -> None: - finite = table.numerical.isfinite() - finite_or_nan = table.numerical.masked_fill(~finite, torch.nan) + numerical = table.numerical + finite = _isfinite(numerical) + count_finite = _count(finite) + finite_or_nan = numerical.masked_fill(~finite, torch.nan) - # Compute finite mean and standard deviation: - mean = finite_or_nan.nanmean(-2, keepdim=True) + # Compute finite mean and standard deviation (equal to 'nanmean', + # which would copy its input to count values): + mean = finite_or_nan.nansum(-2, keepdim=True) / count_finite mean.masked_fill_(mean.isnan(), 0.0) - var = (finite_or_nan - mean).square().nansum(-2, keepdim=True) - var /= (finite.sum(-2, keepdim=True) - 1).clamp_(min=1) + var = finite_or_nan.sub_(mean).square_().nansum(-2, keepdim=True) + del finite_or_nan + var /= (count_finite - 1).clamp_(min=1) std = var.sqrt().clamp(min=1e-6) - # Find values within range: + # Find values within range (non-finite values are never kept): lower = mean - self.threshold * std upper = mean + self.threshold * std - keep = finite & (finite_or_nan >= lower) & (finite_or_nan <= upper) - count = keep.sum(-2, keepdim=True) + keep = (numerical >= lower).logical_and_(numerical <= upper) + keep.logical_and_(finite) + del finite + count = _count(keep) # Compute mean and standard deviation of kept values: - kept_mean = torch.where(keep, finite_or_nan, 0.0).sum(-2, keepdim=True) + kept_mean = torch.where(keep, numerical, 0.0).sum(-2, keepdim=True) kept_mean /= count.clamp(min=1) - centered = torch.where(keep, finite_or_nan - kept_mean, 0.0) + centered = numerical.sub(kept_mean).masked_fill_(~keep, 0.0) denominator = (count - 1).clamp(min=1) - kept_var = centered.square().sum(-2, keepdim=True) / denominator + kept_var = centered.square_().sum(-2, keepdim=True) / denominator + del centered kept_std = kept_var.sqrt().clamp(min=1e-6) has_kept = count > 0 @@ -78,10 +89,26 @@ 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) + 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/sdm/processing/numerical/standardize.py b/sdm/processing/numerical/standardize.py index 8326d7db6..8beb10fb0 100644 --- a/sdm/processing/numerical/standardize.py +++ b/sdm/processing/numerical/standardize.py @@ -5,7 +5,11 @@ from sdm import Stype, TableTensor from sdm.processing import InvertibleMixin, Processor -from sdm.processing.numerical._stats import _constant_feature_mask +from sdm.processing.numerical._stats import ( + _constant_feature_mask, + _count, + _isfinite, +) class Standardize(Processor, InvertibleMixin): @@ -37,35 +41,35 @@ def _fit( generator: torch.Generator | None = None, ) -> None: - finite = table.numerical.isfinite() + finite = _isfinite(table.numerical) + count = _count(finite) finite_or_nan = table.numerical.masked_fill(~finite, torch.nan) finite_or_nan = finite_or_nan.double() # Ensure high precision. - self.mean = finite_or_nan.nanmean(-2, keepdim=True) + # Equal to 'nanmean', which would copy its input to count values. + self.mean = finite_or_nan.nansum(-2, keepdim=True) / count self.mean.masked_fill_(self.mean.isnan(), 0.0) - var = (finite_or_nan - self.mean).square().nanmean(-2, keepdim=True) + # 'finite_or_nan' is a private copy, so it can be modified in place. + var = finite_or_nan.sub_(self.mean).square_().nansum(-2, keepdim=True) + var /= count var.masked_fill_(var.isnan(), 0.0) self.scale = var.sqrt() if self.eps == 0: - mask = _constant_feature_mask( - var, - self.mean, - num_samples=finite.sum(-2, keepdim=True), - ) + mask = _constant_feature_mask(var, self.mean, num_samples=count) self.scale[mask] = 1.0 else: self.scale += self.eps def _transform(self, table: TableTensor) -> TableTensor: dtype = table.numerical.dtype - numerical = (table.numerical - self.mean) / self.scale + numerical = table.numerical.sub(self.mean).div_(self.scale) return table.replace_blocks(numerical=numerical.to(dtype=dtype)) def _inverse_transform(self, table: TableTensor) -> TableTensor: dtype = table.numerical.dtype - numerical = table.numerical * self.scale + self.mean + numerical = table.numerical.mul(self.scale).add_(self.mean) return table.replace_blocks(numerical=numerical.to(dtype=dtype)) def __repr__(self, *, indent: int = 0) -> str: 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..2ae704ff2 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_transform_passes_match_single_pass( + 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, + ) diff --git a/test/processing/numerical/test_standardize.py b/test/processing/numerical/test_standardize.py index 375c736d4..8c709942d 100644 --- a/test/processing/numerical/test_standardize.py +++ b/test/processing/numerical/test_standardize.py @@ -76,6 +76,26 @@ def test_standardize(device: torch.device) -> None: ) +@withCUDA +def test_standardize_ignores_non_finite_values_in_many_rows( + device: torch.device, +) -> None: + inp = torch.randn(2, 600, 3, dtype=torch.float64, device=device) + inp[inp > 1.0] = float("nan") + inp[inp < -1.5] = float("inf") + + out = Standardize().fit_transform(TableTensor.from_tensor(inp)) + + finite = inp.masked_fill(~inp.isfinite(), float("nan")) + mean = finite.nanmean(-2, keepdim=True) + std = (finite - mean).square().nanmean(-2, keepdim=True).sqrt() + torch.testing.assert_close( + out.numerical, + (inp - mean) / std, + equal_nan=True, + ) + + @withCUDA def test_standardize_single_sample_uses_unit_scale( device: torch.device, From 41c20d57fd7076cb1bc8c1d6350614044323a15a Mon Sep 17 00:00:00 2001 From: Jingang Qu Date: Sun, 27 Sep 2026 01:28:31 -0700 Subject: [PATCH 3/9] Embed KumoTabular query rows in passes within the chunk memory limit - Without gradients on CUDA, embed the context rows once and the query rows in balanced passes that replay the recorded context state, when the query cells would exceed the chunk memory limit. - Align passes with the chunks of the row attention, so every row runs in a chunk of the same size as in a single pass. Passes then match a single pass up to rare rounding differences in small passes. - Free the label embedding and each layer's full key/value early in the ICL block. Signed-off-by: Jingang Qu --- sdm/models/kumo/tabular/icl.py | 11 +- sdm/models/kumo/tabular/row_embedding.py | 105 +++++++++++++++++- test/models/kumo/tabular/test_model.py | 48 ++++++++ .../models/kumo/tabular/test_row_embedding.py | 96 +++++++++++++++- 4 files changed, 253 insertions(+), 7 deletions(-) diff --git a/sdm/models/kumo/tabular/icl.py b/sdm/models/kumo/tabular/icl.py index be431353d..067d5f619 100644 --- a/sdm/models/kumo/tabular/icl.py +++ b/sdm/models/kumo/tabular/icl.py @@ -74,6 +74,7 @@ def forward( y_emb = self.y_lin(y.unsqueeze(-1)) # [..., R_train, D] x[..., :R_train, :] += y_emb.to(x.dtype) + del y_emb for i, layer in enumerate(self.layers): cache_key = f"icl_block.layer{i}" @@ -117,12 +118,14 @@ def forward( if last_layer else x[..., :R_train, :], ) + key_value = KVCacheEntry( + key=key[..., : self.kv_heads, :].contiguous(), + value=value[..., : self.kv_heads, :].contiguous(), + ) + del key, value x_query = layer( query=x[..., R_train:, :], - key_value=KVCacheEntry( - key=key[..., : self.kv_heads, :].contiguous(), - value=value[..., : self.kv_heads, :].contiguous(), - ), + key_value=key_value, out=None if torch.is_grad_enabled() else x[..., R_train:, :], ) if last_layer: diff --git a/sdm/models/kumo/tabular/row_embedding.py b/sdm/models/kumo/tabular/row_embedding.py index f69ff3e0e..ccbb0afff 100644 --- a/sdm/models/kumo/tabular/row_embedding.py +++ b/sdm/models/kumo/tabular/row_embedding.py @@ -3,12 +3,14 @@ # ruff: noqa: D101, D102 +import math from typing import Any, cast import torch from torch import Tensor from torch.nn import Embedding, Linear, ModuleList, Parameter +from sdm._memory import chunk_memory_limit from sdm.cache import Cache, KVCacheEntry from sdm.models.kumo.tabular.block import KumoTabularTransformerBlock from sdm.models.kumo.tabular.cell_embedding import CellEmbedding @@ -66,7 +68,7 @@ def __init__( **factory_kwargs, ) - self.col_blocks = ModuleList( + self.col_blocks: ModuleList[InducedTransformerBlock] = ModuleList( InducedTransformerBlock( channels=channels, num_inducing_points=num_inducing_points, @@ -89,7 +91,7 @@ def __init__( ) for _ in range(num_layers) ) - self.row_blocks = ModuleList( + self.row_blocks: ModuleList[KumoTabularTransformerBlock] = ModuleList( KumoTabularTransformerBlock( channels=channels, num_heads=num_heads, @@ -114,6 +116,105 @@ def forward( *, cache: Cache | None = None, ) -> Tensor: # [..., R, K * D] + starts = self._pass_starts(x, train_size=y.size(-1), cache=cache) + if len(starts) == 1: + return self._forward(x, y, categorical_mask, cache=cache) + + # Query rows only read context state, which the first pass records + # for replay in later passes. + cache = Cache() if cache is None else cache + first = self._forward( + x=x[..., : starts[1], :], + y=y, + categorical_mask=categorical_mask, + cache=cache, + ) # [..., starts[1], K * D] + cache.freeze() + out = first.new_empty((*first.shape[:-2], x.size(-2), first.size(-1))) + out[..., : starts[1], :] = first + del first + ends = [*starts[2:], x.size(-2)] + for start, end in zip(starts[1:], ends, strict=True): + out[..., start:end, :] = self._forward( + x=x[..., start:end, :], + y=y[..., :0], + categorical_mask=categorical_mask, + cache=cache, + ) + return out + + def _pass_starts( + self, + x: Tensor, # [..., R, C] + train_size: int, + cache: Cache | None, + ) -> list[int]: + # First rows of the passes that embed the rows of `x`. Passes run + # without gradients on CUDA and replay context state from a cache. + if ( + torch.is_grad_enabled() + or not x.is_cuda + or (cache is not None and cache.is_recording) + ): + return [0] + *B, R, C = x.size() + N = math.prod(B) + K, D = self.readout_token.size(-2), self.channels + G = self.cell_embedding.group_size + M = self.col_blocks[0].inducing_points.size(-2) + s = ( + torch.get_autocast_dtype(x.device.type).itemsize + if torch.is_autocast_enabled(x.device.type) + else x.element_size() + ) + budget = chunk_memory_limit(x.device) + # Bytes per row: the cell buffer, plus the missingness mask, imputed + # values and their feature groups while embedding cells. + row_bytes = N * ( + (K + C) * D * s + (G + 1) * (x.element_size() + 1) * C + ) + # Without a cache, the context pass records the key/value projections + # of all column blocks for the query passes. Query rows that fit the + # chunk memory budget plus these projections run with the context. + state_bytes = 0 + if cache is None: + state_bytes = 2 * N * C * M * D * s * len(self.col_blocks) + if (R - train_size) * row_bytes <= budget + state_bytes: + return [0] + + # Row blocks run the rows of all batch entries in chunks of `chunk` + # rows. In passes starting on multiples of `grid` rows, every row runs + # in a chunk of the same size as in a single pass, since the last pass + # holds the partial last chunk of a single pass, the last + # `N * R % chunk` rows of the last batch entry. Attention over long + # rows rounds differently in small chunks, so this keeps passes equal + # to a single pass up to rare rounding differences in small passes. + chunk = self.row_blocks[0].auto_batch_size_limit( + device=x.device, + element_size=s, + query_length=K + C, + key_value_length=K + C, + ) + grid = chunk // math.gcd(N, chunk) + context = -(-train_size // grid) * grid + last = R - max(N * R % chunk, 1) + if context > last: + return [0] + # Balanced query passes need no more memory than the context pass, + # the budget or one grid of rows, whichever is more. + grids = max(max(train_size, budget // row_bytes) // grid, 1) + num_passes = -(-(R - context) // (grids * grid)) + step = -(-(R - context) // (num_passes * grid)) * grid + return [0, *range(context or step, last + 1, step)] + + def _forward( + self, + x: Tensor, # [..., R, C] + y: Tensor, # [..., R_train] + categorical_mask: Tensor, # [..., C] + *, + cache: Cache | None = None, + ) -> Tensor: # [..., R, K * D] *B, R, C = x.size() R_train = y.size(-1) diff --git a/test/models/kumo/tabular/test_model.py b/test/models/kumo/tabular/test_model.py index 71e1782bc..baca474a1 100644 --- a/test/models/kumo/tabular/test_model.py +++ b/test/models/kumo/tabular/test_model.py @@ -9,6 +9,8 @@ import sdm.processing as sp from sdm import CategoricalTensor, Stype, TableTensor from sdm.models import KumoTabular +from sdm.models.kumo.tabular.row_embedding import RowEmbedding +from sdm.testing import onlyCUDA def _build( @@ -376,3 +378,49 @@ def test_estimator_batching_many_classes_query_chunks( x_query=x_query, estimator_batch_size="auto", ) + + +@onlyCUDA +@pytest.mark.parametrize("task", ["classification", "regression"]) +def test_query_passes_match_single_pass( + task: Literal["classification", "regression"], + monkeypatch: pytest.MonkeyPatch, +) -> None: + model = _build(task, "small").cuda().eval() + x = torch.randn(1008, 4, device="cuda") + x[::3, 0] = float("nan") + x_context, x_query = TableTensor.from_tensor(x).split([8, 1000], dim=0) + # 12 classes run through ECOC. + target = ( + TableTensor( + columns={Stype.categorical: ("target",)}, + categorical=CategoricalTensor( + code=torch.arange(8, device="cuda").unsqueeze(-1), + categories=(torch.arange(12, device="cuda"),), + ), + ) + if task == "classification" + else TableTensor.from_tensor(torch.randn(8, 1, device="cuda")) + ) + + def forward() -> TableTensor: + return model( + x_context=x_context, + y_context=target, + x_query=x_query, + num_estimators=2, + generator=torch.Generator("cuda").manual_seed(0), + ) + + # A chunk memory limit of 1 MiB embeds the query rows in passes. + total_memory = torch.cuda.get_device_properties("cuda").total_memory + monkeypatch.setenv("SDM_CHUNK_MEMORY_FRACTION", str(2**20 / total_memory)) + actual = forward() + monkeypatch.setattr( + RowEmbedding, + "_pass_starts", + lambda self, x, train_size, cache: [0], + ) + expected = forward() + + torch.testing.assert_close(actual.numerical, expected.numerical) diff --git a/test/models/kumo/tabular/test_row_embedding.py b/test/models/kumo/tabular/test_row_embedding.py index 29fa0380e..b7e5e7963 100644 --- a/test/models/kumo/tabular/test_row_embedding.py +++ b/test/models/kumo/tabular/test_row_embedding.py @@ -1,11 +1,12 @@ # SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 +import pytest import torch from sdm.cache import Cache from sdm.models.kumo.tabular.row_embedding import RowEmbedding -from sdm.testing import withCUDA +from sdm.testing import onlyCUDA, withCUDA @withCUDA @@ -46,3 +47,96 @@ def test_row_embedding(device: torch.device) -> None: assert out.size() == (2, 2, 32) assert out.device == device torch.testing.assert_close(out, expected[:, 3:], atol=1e-5, rtol=1e-5) + + +@onlyCUDA +@pytest.mark.parametrize("cached", [False, True]) +def test_row_embedding_passes( + cached: bool, + monkeypatch: pytest.MonkeyPatch, +) -> None: + encoder = RowEmbedding( + num_classes=0, + channels=16, + num_layers=2, + num_heads=2, + group_size=3, + num_frequencies=8, + num_inducing_points=4, + num_readout_tokens=2, + device="cuda", + ) + # With many batch entries, e.g. estimators or ECOC tasks, a chunk of the + # row attention spans more rows than the chunk memory limit allows. + x = torch.randn(17, 64, 3, device="cuda") + x[x > 1.0] = torch.nan + y = torch.randn(17, 4, device="cuda") + categorical_mask = torch.tensor([True, False, False], device="cuda") + + @torch.inference_mode() + def embed() -> torch.Tensor: + if not cached: + return encoder(x, y, categorical_mask) + cache = Cache() + encoder(x[:, :4], y, categorical_mask, cache=cache) + return encoder( + x[:, 4:], + y[:, :0], + categorical_mask, + cache=cache.freeze(), + ) + + expected = embed() + # Chunk memory limits from 1 KiB to 512 KiB embed the rows in passes of + # a few rows up to a single pass. + total_memory = torch.cuda.get_device_properties("cuda").total_memory + for exponent in range(40, 77): + fraction = 2 ** (exponent / 4) / total_memory + monkeypatch.setenv("SDM_CHUNK_MEMORY_FRACTION", str(fraction)) + torch.testing.assert_close(embed(), expected) + + +@onlyCUDA +def test_row_embedding_passes_on_wide_tables( + monkeypatch: pytest.MonkeyPatch, +) -> None: + # Attention over more than 256 row tokens splits keys and values for + # small batches of rows, so passes must keep the row chunks of one pass. + encoder = RowEmbedding( + num_classes=10, + channels=64, + num_layers=2, + num_heads=2, + group_size=3, + num_frequencies=8, + num_inducing_points=8, + num_readout_tokens=2, + device="cuda", + ) + # Randomize the zero-initialized residual branches, so that attention + # outputs reach the embedding. + for parameter in encoder.parameters(): + if not parameter.any(): + torch.nn.init.normal_(parameter, std=0.02) + x = torch.randn(3, 1169, 300, device="cuda") + x[x > 2.0] = torch.nan + y = torch.randint(10, (3, 200), device="cuda") + categorical_mask = torch.arange(300, device="cuda") % 4 == 0 + + def embed() -> tuple[torch.Tensor, int]: + torch.cuda.reset_peak_memory_stats() + with torch.no_grad(), torch.autocast("cuda", dtype=torch.float16): + out = encoder(x, y, categorical_mask) + return out, torch.cuda.max_memory_allocated() + + monkeypatch.setenv("SDM_CHUNK_MEMORY_FRACTION", "0.001") + actual, peak = embed() + monkeypatch.setattr( + RowEmbedding, + "_pass_starts", + lambda self, x, train_size, cache: [0], + ) + expected, expected_peak = embed() + + assert torch.equal(actual, expected) + assert peak < expected_peak From 1f17b3548b2df836e4581f39f52dd6e164fdc32f Mon Sep 17 00:00:00 2001 From: Jingang Qu Date: Sun, 27 Sep 2026 01:28:31 -0700 Subject: [PATCH 4/9] Transform query rows and model outputs in passes - Transform queries and post-process member outputs in passes over rows within the chunk memory limit in `RecipeExecution`. - Keep member outputs in the model output dtype until post-processing, which casts them to the query dtype and inverts numerical targets pass by pass. - Post-process the benchmark adapter's outputs through `transform_output` after freeing the transformed queries. Signed-off-by: Jingang Qu --- benchmark/tabular/model.py | 11 ++-- sdm/models/base.py | 28 ++++------- sdm/processing/execution.py | 84 +++++++++++++++++++++++++++++-- sdm/processing/recipe.py | 4 ++ test/processing/test_execution.py | 66 ++++++++++++++++++++++++ 5 files changed, 162 insertions(+), 31 deletions(-) diff --git a/benchmark/tabular/model.py b/benchmark/tabular/model.py index 89b08811a..8e86da0b9 100644 --- a/benchmark/tabular/model.py +++ b/benchmark/tabular/model.py @@ -214,6 +214,7 @@ def _predict_proba( x=x_query, related_tables=None, ) + dtype = queries[0].x.dtype generator = torch.Generator(self._device).set_state( self._rng_state ) @@ -230,14 +231,8 @@ def _predict_proba( ], generator=generator, ) - - if self.problem_type == REGRESSION: - outputs = list( - self._recipe_execution.inverse_transform_target( - outputs - ) - ) - out = self._recipe_execution.transform_output(outputs) + del queries + out = self._recipe_execution.transform_output(outputs, dtype) if self.problem_type == REGRESSION: return out.numerical.float().mean(dim=-1).cpu().numpy() diff --git a/sdm/models/base.py b/sdm/models/base.py index 75c019738..2e4e04add 100644 --- a/sdm/models/base.py +++ b/sdm/models/base.py @@ -204,19 +204,14 @@ def forward( **kwargs, ) - # Regression: invert target before stacking estimator outputs. - if contexts[0].y.numerical.size(-1) > 0: - with ( - torch.amp.autocast(x_query.device.type, enabled=False), - inference_mode("grad" if requires_grad else "inference"), - ): - outs = list(recipe_execution.inverse_transform_target(outs)) - with ( torch.amp.autocast(x_query.device.type, enabled=False), inference_mode("grad" if requires_grad else "inference"), ): - return recipe_execution.transform_output(outs) + return recipe_execution.transform_output( + outs, + dtype=queries[0].x.dtype, + ) def fit( self, @@ -529,19 +524,14 @@ def predict( transfer_stream.synchronize() raise - # Regression: invert target before stacking estimator outputs. - if cast(Cache, self._cache[0])["classes"] is None: - with ( - torch.amp.autocast(x.device.type, enabled=False), - inference_mode("grad" if requires_grad else "inference"), - ): - outs = list(recipe_execution.inverse_transform_target(outs)) - with ( torch.amp.autocast(x.device.type, enabled=False), inference_mode("grad" if requires_grad else "inference"), ): - return recipe_execution.transform_output(outs) + return recipe_execution.transform_output( + outs, + dtype=queries[0].x.dtype, + ) def clear(self) -> None: r"""Clear cached context state created by :meth:`fit`.""" @@ -796,7 +786,7 @@ def _forward_batch( for i in range(len(outs)): for callback in callbacks: outs[i] = callback.on_model_forward_end(self, outs[i]) - return [cast(TableTensor, out.to(query.x.dtype)) for out in outs] + return outs def _validate_context( self, diff --git a/sdm/processing/execution.py b/sdm/processing/execution.py index 3f03f4c77..5680ce87f 100644 --- a/sdm/processing/execution.py +++ b/sdm/processing/execution.py @@ -4,7 +4,8 @@ from __future__ import annotations import copy -from collections.abc import Mapping, Sequence +import math +from collections.abc import Callable, Mapping, Sequence from typing import NamedTuple, cast import torch @@ -12,6 +13,7 @@ import sdm.processing as sp from sdm import EnsembleTable, Recipe, RelatedTables, Stype, TableTensor +from sdm._memory import split_size from sdm.processing import EnsembleInvertibleMixin, EnsembleProcessor @@ -44,6 +46,7 @@ def __init__(self, recipe: Recipe) -> None: ) = None self._num_estimators: int | None = None self._y_locations: tuple[tuple[int, int], ...] | None = None + self._numerical_target = False @property def num_members(self) -> int: @@ -68,6 +71,7 @@ def fit_transform( self._num_estimators = num_members self._y_locations = y._locations + self._numerical_target = y[0].numerical.size(-1) > 0 task_dispatchers = tuple( module @@ -164,9 +168,13 @@ def transform( x: Tensor | TableTensor | EnsembleTable, related_tables: RelatedTables | None, ) -> tuple[MemberQuery, ...]: - """Transform query data.""" + """Transform query data. + + Fitted processors transform query rows independently, so large + queries are transformed in passes over their rows. + """ x = _to_ensemble_table(x, self._num_estimators) - x = self.recipe.features.transform_ensemble(x) + x = _transform_rows(self.recipe.features.transform_ensemble, x) if len(x) != self.num_members: raise ValueError( "Expected inputs to map to the same number of ensemble members" @@ -254,8 +262,44 @@ def inverse_transform_target( def transform_output( self, outputs: Sequence[TableTensor], + dtype: torch.dtype, + ) -> TableTensor: + """Apply ``recipe.output`` to member outputs in ``dtype``. + + Outputs of numerical targets first pass through the inverted target + transforms. Output rows are processed independently, so large outputs + are processed in passes over their rows. + """ + # Output cells are processed in double precision at most. + size = split_size( + num_items=outputs[0].size(-2), + item_bytes=sum( + math.prod(output.size()[:-2]) * output.size(-1) + for output in outputs + ) + * torch.float64.itemsize, + device=outputs[0].device, + ) + if size >= outputs[0].size(-2): + return self._transform_output(outputs, dtype) + parts = tuple( + self._transform_output(chunk, dtype) + for chunk in zip( + *(output.split(size, dim=-2) for output in outputs), + strict=True, + ) + ) + return cast(TableTensor, torch.cat(parts, dim=-2)) + + def _transform_output( + self, + outputs: Sequence[TableTensor], + dtype: torch.dtype, ) -> TableTensor: - """Apply ``recipe.output`` to member outputs.""" + outputs = [cast(TableTensor, output.to(dtype)) for output in outputs] + if self._numerical_target: + outputs = list(self.inverse_transform_target(outputs)) + if len(outputs) == 1: out = outputs[0].unsqueeze(0) else: @@ -273,6 +317,38 @@ def transform_output( return self.recipe.output.transform(cast(TableTensor, out)) +def _transform_rows( + transform: Callable[[EnsembleTable], EnsembleTable], + table: EnsembleTable, +) -> EnsembleTable: + # Member cells are transformed in double precision at most. + first = table._groups[0] # [members, ..., rows, columns] + size = split_size( + num_items=first.size(-2), + item_bytes=len(table) + * math.prod(first.size()[1:-2]) + * first.size(-1) + * torch.float64.itemsize, + device=first.device, + ) + if size >= first.size(-2): + return transform(table) + parts = [ + transform(table.replace_groups(groups)) + for groups in zip( + *(group.split(size, dim=-2) for group in table._groups), + strict=True, + ) + ] + # Passes share the fitted state and thus the member layout: + return parts[0].replace_groups( + [ + cast(TableTensor, torch.cat(groups, dim=-2)) + for groups in zip(*(part._groups for part in parts), strict=True) + ] + ) + + def _align_to_fitted_groups( ensemble_table: EnsembleTable, fitted_locations: Sequence[tuple[int, int]], diff --git a/sdm/processing/recipe.py b/sdm/processing/recipe.py index 48754d872..9a705dac2 100644 --- a/sdm/processing/recipe.py +++ b/sdm/processing/recipe.py @@ -29,6 +29,10 @@ class Recipe: the output remains stacked. Steps before the reducer must support stacked outputs, while steps after it receive already-reduced outputs. + Fitted ``features`` and ``output`` steps, and inverse ``target`` steps, + transform each row on its own, so models may apply them to large query + sets in passes over rows. + Each pipeline exposes ``fit``/``transform``/``fit_transform`` and, when its steps are invertible, ``inverse_transform``. Call them directly, e.g. ``recipe.features.transform(table)`` or diff --git a/test/processing/test_execution.py b/test/processing/test_execution.py index ef8b55b8a..8f9aee6c6 100644 --- a/test/processing/test_execution.py +++ b/test/processing/test_execution.py @@ -13,7 +13,9 @@ Stype, TableTensor, ) +from sdm.models import KumoTabular from sdm.processing.execution import RecipeExecution +from sdm.testing import onlyCUDA def test_sequence_uses_batched_fit_states_for_shared_query() -> None: @@ -380,3 +382,67 @@ def test_member_context_exposes_input_stypes() -> None: assert context.input_stypes == x.stypes assert context.x.columns[Stype.numerical] == ("n", "c") + + +@onlyCUDA +@pytest.mark.parametrize("task", ["classification", "regression"]) +def test_row_passes_match_single_pass( + task: str, + monkeypatch: pytest.MonkeyPatch, +) -> None: + device = torch.device("cuda:0") + x = torch.randn(28, 4, device=device) + x[::3, 0] = float("nan") + x_context, x_query = TableTensor.from_tensor(x).split([8, 20], dim=0) + if task == "classification": + y = TableTensor( + columns={Stype.categorical: ("target",)}, + categorical=CategoricalTensor( + code=torch.arange(8, device=device).unsqueeze(-1) % 3, + categories=(torch.arange(3, device=device),), + ), + ) + num_outputs = 3 + else: + y = TableTensor.from_tensor(torch.randn(8, 1, device=device)) + num_outputs = 5 + execution = RecipeExecution(KumoTabular.default_recipe()) + execution.fit_transform( + x=x_context, + y=y, + related_tables=None, + num_members=4, + ) + columns = [str(i) for i in range(num_outputs)] + outputs = [ + TableTensor( + columns={Stype.numerical: columns}, + numerical=torch.randn(20, num_outputs, device=device).half(), + ) + for _ in range(4) + ] + + # Passes run without autograd, like model inference. + @torch.inference_mode() + def run() -> tuple[tuple[TableTensor, ...], TableTensor]: + queries = execution.transform(x=x_query, related_tables=None) + output = execution.transform_output(outputs, torch.float32) + return tuple(query.x for query in queries), output + + expected_queries, expected_output = run() + # Passes of a single row: + monkeypatch.setenv("SDM_CHUNK_MEMORY_FRACTION", "1e-12") + queries, output = run() + + for query, expected_query in zip(queries, expected_queries, strict=True): + assert query.columns == expected_query.columns + torch.testing.assert_close( + query.numerical, + expected_query.numerical, + rtol=0, + atol=0, + equal_nan=True, + ) + assert output.columns == expected_output.columns + assert output.numerical.dtype == torch.float32 + assert torch.equal(output.numerical, expected_output.numerical) From 1adb0358f847b6b687bd5a2f8f8ef2ce4c28d633 Mon Sep 17 00:00:00 2001 From: Cedric Lorenz Date: Mon, 28 Sep 2026 15:39:00 +0200 Subject: [PATCH 5/9] Fix heterogeneous ensemble row preprocessing --- sdm/processing/execution.py | 26 ++++++++++++++++---------- test/processing/test_execution.py | 29 ++++++++++++++++++++++++++++- 2 files changed, 44 insertions(+), 11 deletions(-) diff --git a/sdm/processing/execution.py b/sdm/processing/execution.py index 5680ce87f..eb4e24dda 100644 --- a/sdm/processing/execution.py +++ b/sdm/processing/execution.py @@ -322,21 +322,27 @@ def _transform_rows( table: EnsembleTable, ) -> EnsembleTable: # Member cells are transformed in double precision at most. - first = table._groups[0] # [members, ..., rows, columns] - size = split_size( - num_items=first.size(-2), - item_bytes=len(table) - * math.prod(first.size()[1:-2]) - * first.size(-1) - * torch.float64.itemsize, - device=first.device, + # Groups have shape [stored members, ..., rows, columns]. + num_rows = table._groups[0].size(-2) + row_bytes_by_device: dict[torch.device, int] = {} + # Count logical members because several members may share one stored table. + for group_id, _ in table._locations: + group = table._groups[group_id] + batch_size = math.prod(group.size()[1:-2]) + row_bytes = batch_size * group.size(-1) * torch.float64.itemsize + row_bytes_by_device[group.device] = ( + row_bytes_by_device.get(group.device, 0) + row_bytes + ) + rows_per_pass = min( + split_size(num_rows, row_bytes, device) + for device, row_bytes in row_bytes_by_device.items() ) - if size >= first.size(-2): + if rows_per_pass >= num_rows: return transform(table) parts = [ transform(table.replace_groups(groups)) for groups in zip( - *(group.split(size, dim=-2) for group in table._groups), + *(group.split(rows_per_pass, dim=-2) for group in table._groups), strict=True, ) ] diff --git a/test/processing/test_execution.py b/test/processing/test_execution.py index 8f9aee6c6..624476d33 100644 --- a/test/processing/test_execution.py +++ b/test/processing/test_execution.py @@ -14,7 +14,7 @@ TableTensor, ) from sdm.models import KumoTabular -from sdm.processing.execution import RecipeExecution +from sdm.processing.execution import RecipeExecution, _transform_rows from sdm.testing import onlyCUDA @@ -446,3 +446,30 @@ def run() -> tuple[tuple[TableTensor, ...], TableTensor]: assert output.columns == expected_output.columns assert output.numerical.dtype == torch.float32 assert torch.equal(output.numerical, expected_output.numerical) + + +def test_transform_rows_sizes_every_member( + monkeypatch: pytest.MonkeyPatch, +) -> None: + item_bytes: list[int] = [] + + def record_size( + num_items: int, + bytes_per_item: int, + device: torch.device, + ) -> int: + del device + item_bytes.append(bytes_per_item) + return num_items + + monkeypatch.setattr("sdm.processing.execution.split_size", record_size) + narrow = TableTensor.from_tensor(torch.zeros(3, 1)) + wide = TableTensor.from_tensor(torch.zeros(3, 4)) + for tables, member_ids in ( + ((narrow, wide), (0, 1, 1)), + ((wide, narrow), (1, 0, 0)), + ): + table = EnsembleTable.from_tables(tables, member_ids) + _transform_rows(lambda value: value, table) + + assert item_bytes == [9 * torch.float64.itemsize] * 2 From 64b23ea1a5e504301844361239b4a260ff7b1b86 Mon Sep 17 00:00:00 2001 From: Cedric Lorenz Date: Mon, 28 Sep 2026 07:01:34 -0700 Subject: [PATCH 6/9] Preserve gradients through clipping --- sdm/processing/numerical/clip_soft.py | 8 ++++++++ sdm/processing/numerical/sigma_clip.py | 6 ++++++ 2 files changed, 14 insertions(+) diff --git a/sdm/processing/numerical/clip_soft.py b/sdm/processing/numerical/clip_soft.py index 2e9ea02f9..5a84ee7b3 100644 --- a/sdm/processing/numerical/clip_soft.py +++ b/sdm/processing/numerical/clip_soft.py @@ -40,6 +40,14 @@ def __init__(self, max_absolute_value: float) -> None: def _transform(self, table: TableTensor) -> TableTensor: numerical = table.numerical bound = self.max_absolute_value + if torch.is_grad_enabled() and numerical.requires_grad: + unit = numerical.div(bound).abs() + root = unit.square().add(1).sqrt() + unit = unit.div(root).masked_fill(~_isfinite(root), 1.0) + clipped = numerical.sign().mul(bound).mul(unit) + clipped = torch.where(numerical.isnan(), numerical, clipped) + return table.replace_blocks(numerical=clipped) + unit = numerical.div(bound).abs_() root = unit.square().add_(1).sqrt_() # 'root' is finite exactly where '1 + unit ** 2' is. diff --git a/sdm/processing/numerical/sigma_clip.py b/sdm/processing/numerical/sigma_clip.py index 137c2e799..b04b456fc 100644 --- a/sdm/processing/numerical/sigma_clip.py +++ b/sdm/processing/numerical/sigma_clip.py @@ -90,6 +90,12 @@ def _fit( def _transform(self, table: TableTensor) -> TableTensor: numerical = table.numerical dtype = torch.promote_types(numerical.dtype, self.lower_bound.dtype) + if torch.is_grad_enabled() and numerical.requires_grad: + 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( From c71f0c36f755d01bfc9d28963790343d0f661bf3 Mon Sep 17 00:00:00 2001 From: Cedric Lorenz Date: Mon, 28 Sep 2026 07:17:42 -0700 Subject: [PATCH 7/9] Preserve estimator output dtypes --- benchmark/tabular/model.py | 4 ++-- sdm/models/base.py | 4 ++-- sdm/processing/execution.py | 15 +++++++++------ test/processing/test_execution.py | 5 ++++- 4 files changed, 17 insertions(+), 11 deletions(-) diff --git a/benchmark/tabular/model.py b/benchmark/tabular/model.py index 8e86da0b9..3ea9f3b84 100644 --- a/benchmark/tabular/model.py +++ b/benchmark/tabular/model.py @@ -214,7 +214,7 @@ def _predict_proba( x=x_query, related_tables=None, ) - dtype = queries[0].x.dtype + dtypes = tuple(query.x.dtype for query in queries) generator = torch.Generator(self._device).set_state( self._rng_state ) @@ -232,7 +232,7 @@ def _predict_proba( generator=generator, ) del queries - out = self._recipe_execution.transform_output(outputs, dtype) + out = self._recipe_execution.transform_output(outputs, dtypes) if self.problem_type == REGRESSION: return out.numerical.float().mean(dim=-1).cpu().numpy() diff --git a/sdm/models/base.py b/sdm/models/base.py index 2e4e04add..1421fc4ce 100644 --- a/sdm/models/base.py +++ b/sdm/models/base.py @@ -210,7 +210,7 @@ def forward( ): return recipe_execution.transform_output( outs, - dtype=queries[0].x.dtype, + dtypes=tuple(query.x.dtype for query in queries), ) def fit( @@ -530,7 +530,7 @@ def predict( ): return recipe_execution.transform_output( outs, - dtype=queries[0].x.dtype, + dtypes=tuple(query.x.dtype for query in queries), ) def clear(self) -> None: diff --git a/sdm/processing/execution.py b/sdm/processing/execution.py index eb4e24dda..3e3df24d0 100644 --- a/sdm/processing/execution.py +++ b/sdm/processing/execution.py @@ -262,9 +262,9 @@ def inverse_transform_target( def transform_output( self, outputs: Sequence[TableTensor], - dtype: torch.dtype, + dtypes: Sequence[torch.dtype], ) -> TableTensor: - """Apply ``recipe.output`` to member outputs in ``dtype``. + """Apply ``recipe.output`` after restoring each member dtype. Outputs of numerical targets first pass through the inverted target transforms. Output rows are processed independently, so large outputs @@ -281,9 +281,9 @@ def transform_output( device=outputs[0].device, ) if size >= outputs[0].size(-2): - return self._transform_output(outputs, dtype) + return self._transform_output(outputs, dtypes) parts = tuple( - self._transform_output(chunk, dtype) + self._transform_output(chunk, dtypes) for chunk in zip( *(output.split(size, dim=-2) for output in outputs), strict=True, @@ -294,9 +294,12 @@ def transform_output( def _transform_output( self, outputs: Sequence[TableTensor], - dtype: torch.dtype, + dtypes: Sequence[torch.dtype], ) -> TableTensor: - outputs = [cast(TableTensor, output.to(dtype)) for output in outputs] + outputs = [ + cast(TableTensor, output.to(dtype)) + for output, dtype in zip(outputs, dtypes, strict=True) + ] if self._numerical_target: outputs = list(self.inverse_transform_target(outputs)) diff --git a/test/processing/test_execution.py b/test/processing/test_execution.py index 624476d33..7034e84f8 100644 --- a/test/processing/test_execution.py +++ b/test/processing/test_execution.py @@ -426,7 +426,10 @@ def test_row_passes_match_single_pass( @torch.inference_mode() def run() -> tuple[tuple[TableTensor, ...], TableTensor]: queries = execution.transform(x=x_query, related_tables=None) - output = execution.transform_output(outputs, torch.float32) + output = execution.transform_output( + outputs, + (torch.float32,) * len(outputs), + ) return tuple(query.x for query in queries), output expected_queries, expected_output = run() From 53e6ae6808fcf2f4f6221f807ad31015e6cdb56d Mon Sep 17 00:00:00 2001 From: Cedric Lorenz Date: Mon, 28 Sep 2026 07:53:52 -0700 Subject: [PATCH 8/9] Reject mismatched ensemble row counts --- sdm/processing/execution.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/sdm/processing/execution.py b/sdm/processing/execution.py index 3e3df24d0..89b111f98 100644 --- a/sdm/processing/execution.py +++ b/sdm/processing/execution.py @@ -327,6 +327,8 @@ def _transform_rows( # Member cells are transformed in double precision at most. # Groups have shape [stored members, ..., rows, columns]. num_rows = table._groups[0].size(-2) + if any(group.size(-2) != num_rows for group in table._groups[1:]): + raise ValueError("Expected all ensemble groups to have the same row count") row_bytes_by_device: dict[torch.device, int] = {} # Count logical members because several members may share one stored table. for group_id, _ in table._locations: From 4e58b39a28a4b01b7932bd321fec043508fc4091 Mon Sep 17 00:00:00 2001 From: Valter Hudovernik Date: Mon, 28 Sep 2026 21:15:46 +0200 Subject: [PATCH 9/9] cleanup after rebase --- benchmark/tabular/model.py | 3 +- sdm/models/base.py | 37 +++++++------- sdm/processing/execution.py | 34 ++++--------- sdm/processing/numerical/clip_soft.py | 8 +-- sdm/processing/numerical/sigma_clip.py | 11 ++-- .../models/kumo/tabular/test_row_embedding.py | 12 +++-- test/models/test_base.py | 51 ++++++++++++++++++- test/models/test_ecoc.py | 26 +++------- test/processing/numerical/test_sigma_clip.py | 25 +++++++++ test/processing/test_execution.py | 49 +++++++++--------- 10 files changed, 156 insertions(+), 100 deletions(-) diff --git a/benchmark/tabular/model.py b/benchmark/tabular/model.py index 3ea9f3b84..0f1ce7164 100644 --- a/benchmark/tabular/model.py +++ b/benchmark/tabular/model.py @@ -214,7 +214,6 @@ def _predict_proba( x=x_query, related_tables=None, ) - dtypes = tuple(query.x.dtype for query in queries) generator = torch.Generator(self._device).set_state( self._rng_state ) @@ -223,7 +222,7 @@ def _predict_proba( self.autocast_dtype, enabled=x_query.is_cuda, ): - outputs = self.model._forward_members( + outputs, dtypes = self.model._forward_members( contexts=self._contexts, queries=queries, estimator_batch_size=self._get_model_params()[ diff --git a/sdm/models/base.py b/sdm/models/base.py index 1421fc4ce..d206e5455 100644 --- a/sdm/models/base.py +++ b/sdm/models/base.py @@ -195,7 +195,7 @@ def forward( related_tables=related_query_tables, ) - outs = self._forward_members( + outs, dtypes = self._forward_members( contexts=contexts, queries=queries, estimator_batch_size=estimator_batch_size, @@ -210,7 +210,7 @@ def forward( ): return recipe_execution.transform_output( outs, - dtypes=tuple(query.x.dtype for query in queries), + dtypes=dtypes, ) def fit( @@ -437,6 +437,7 @@ def predict( compute_stream.wait_stream(transfer_stream) outs: list[TableTensor] = [] + dtypes: list[torch.dtype] = [] start = 0 for i in range(len(caches)): cache, next_cache = next_cache, None @@ -460,6 +461,7 @@ def predict( strict=True, ) ] + dtypes.extend(member.x.dtype for member in members) start += len(x_schemas) if i + 1 < len(caches): @@ -530,7 +532,7 @@ def predict( ): return recipe_execution.transform_output( outs, - dtypes=tuple(query.x.dtype for query in queries), + dtypes=dtypes, ) def clear(self) -> None: @@ -634,13 +636,13 @@ def _forward_members( callbacks: Sequence[Callback] | None = None, generator: torch.Generator | None = None, **kwargs: Any, - ) -> list[TableTensor]: + ) -> tuple[list[TableTensor], list[torch.dtype]]: r"""Run recipe-transformed members on the device of their queries. Context features and targets may live on another device; each batch is moved when it runs. - Returns one output per member before target inversion and - ``recipe.output``. + Returns member outputs before target inversion and ``recipe.output``, + and their callback-prepared query dtypes. """ callbacks = () if callbacks is None else callbacks requires_grad = self.training @@ -651,18 +653,19 @@ def _forward_members( # Complete each sequential member before preparing the next member. if estimator_batch_size == 1 and len(contexts) > 1: outs: list[TableTensor] = [] + dtypes: list[torch.dtype] = [] for context, query in zip(contexts, queries, strict=True): - outs.extend( - self._forward_members( - contexts=(context,), - queries=(query,), - estimator_batch_size=1, - callbacks=callbacks, - generator=generator, - **kwargs, - ) + member_outs, member_dtypes = self._forward_members( + contexts=(context,), + queries=(query,), + estimator_batch_size=1, + callbacks=callbacks, + generator=generator, + **kwargs, ) - return outs + outs.extend(member_outs) + dtypes.extend(member_dtypes) + return outs, dtypes members = [ self._prepare_context(context, callbacks) for context in contexts @@ -706,7 +709,7 @@ def _forward_members( generator=generator, **kwargs, ) - return outs + return outs, [query.x.dtype for query in queries] def _prepare_context( self, diff --git a/sdm/processing/execution.py b/sdm/processing/execution.py index 89b111f98..7da820264 100644 --- a/sdm/processing/execution.py +++ b/sdm/processing/execution.py @@ -168,11 +168,7 @@ def transform( x: Tensor | TableTensor | EnsembleTable, related_tables: RelatedTables | None, ) -> tuple[MemberQuery, ...]: - """Transform query data. - - Fitted processors transform query rows independently, so large - queries are transformed in passes over their rows. - """ + """Transform query data.""" x = _to_ensemble_table(x, self._num_estimators) x = _transform_rows(self.recipe.features.transform_ensemble, x) if len(x) != self.num_members: @@ -266,9 +262,8 @@ def transform_output( ) -> TableTensor: """Apply ``recipe.output`` after restoring each member dtype. - Outputs of numerical targets first pass through the inverted target - transforms. Output rows are processed independently, so large outputs - are processed in passes over their rows. + Numerical targets are inverse-transformed before applying the output + recipe. """ # Output cells are processed in double precision at most. size = split_size( @@ -296,26 +291,17 @@ def _transform_output( outputs: Sequence[TableTensor], dtypes: Sequence[torch.dtype], ) -> TableTensor: - outputs = [ + outputs = tuple( cast(TableTensor, output.to(dtype)) for output, dtype in zip(outputs, dtypes, strict=True) - ] + ) if self._numerical_target: - outputs = list(self.inverse_transform_target(outputs)) + outputs = self.inverse_transform_target(outputs) if len(outputs) == 1: out = outputs[0].unsqueeze(0) else: - expected = set(outputs[0].columns[Stype.numerical]) - for output in outputs[1:]: - if set(output.columns[Stype.numerical]) != expected: - raise ValueError( - "Expected all model outputs to have the same columns " - "before applying 'Recipe.output'. Ensure every target " - "contains the same set of classes." - ) - - out = torch.stack(list(outputs), dim=0) + out = torch.stack(outputs, dim=0) return self.recipe.output.transform(cast(TableTensor, out)) @@ -327,8 +313,6 @@ def _transform_rows( # Member cells are transformed in double precision at most. # Groups have shape [stored members, ..., rows, columns]. num_rows = table._groups[0].size(-2) - if any(group.size(-2) != num_rows for group in table._groups[1:]): - raise ValueError("Expected all ensemble groups to have the same row count") row_bytes_by_device: dict[torch.device, int] = {} # Count logical members because several members may share one stored table. for group_id, _ in table._locations: @@ -344,6 +328,10 @@ def _transform_rows( ) if rows_per_pass >= num_rows: return transform(table) + if any(group.size(-2) != num_rows for group in table._groups[1:]): + raise ValueError( + "Expected all ensemble groups to have the same row count" + ) parts = [ transform(table.replace_groups(groups)) for groups in zip( diff --git a/sdm/processing/numerical/clip_soft.py b/sdm/processing/numerical/clip_soft.py index 5a84ee7b3..1d6ec874c 100644 --- a/sdm/processing/numerical/clip_soft.py +++ b/sdm/processing/numerical/clip_soft.py @@ -41,10 +41,10 @@ def _transform(self, table: TableTensor) -> TableTensor: numerical = table.numerical bound = self.max_absolute_value if torch.is_grad_enabled() and numerical.requires_grad: - unit = numerical.div(bound).abs() - root = unit.square().add(1).sqrt() - unit = unit.div(root).masked_fill(~_isfinite(root), 1.0) - clipped = numerical.sign().mul(bound).mul(unit) + ratio = (numerical / bound).abs() + squared = 1 + ratio.square() + unit = torch.where(squared.isfinite(), ratio / squared.sqrt(), 1.0) + clipped = numerical.sign() * bound * unit clipped = torch.where(numerical.isnan(), numerical, clipped) return table.replace_blocks(numerical=clipped) diff --git a/sdm/processing/numerical/sigma_clip.py b/sdm/processing/numerical/sigma_clip.py index b04b456fc..c6a77b957 100644 --- a/sdm/processing/numerical/sigma_clip.py +++ b/sdm/processing/numerical/sigma_clip.py @@ -57,8 +57,13 @@ def _fit( mean = finite_or_nan.nansum(-2, keepdim=True) / count_finite mean.masked_fill_(mean.isnan(), 0.0) - var = finite_or_nan.sub_(mean).square_().nansum(-2, keepdim=True) - del finite_or_nan + centered = ( + finite_or_nan.sub(mean) + if torch.is_grad_enabled() + else finite_or_nan.sub_(mean) + ) + var = centered.square_().nansum(-2, keepdim=True) + del centered, finite_or_nan var /= (count_finite - 1).clamp_(min=1) std = var.sqrt().clamp(min=1e-6) @@ -90,7 +95,7 @@ def _fit( def _transform(self, table: TableTensor) -> TableTensor: numerical = table.numerical dtype = torch.promote_types(numerical.dtype, self.lower_bound.dtype) - if torch.is_grad_enabled() and numerical.requires_grad: + 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) diff --git a/test/models/kumo/tabular/test_row_embedding.py b/test/models/kumo/tabular/test_row_embedding.py index b7e5e7963..76d69fa63 100644 --- a/test/models/kumo/tabular/test_row_embedding.py +++ b/test/models/kumo/tabular/test_row_embedding.py @@ -87,11 +87,10 @@ def embed() -> torch.Tensor: ) expected = embed() - # Chunk memory limits from 1 KiB to 512 KiB embed the rows in passes of - # a few rows up to a single pass. + # Small, intermediate and single-pass budgets exercise pass boundaries. total_memory = torch.cuda.get_device_properties("cuda").total_memory - for exponent in range(40, 77): - fraction = 2 ** (exponent / 4) / total_memory + for budget in (2**10, 2**14, 2**19): + fraction = budget / total_memory monkeypatch.setenv("SDM_CHUNK_MEMORY_FRACTION", str(fraction)) torch.testing.assert_close(embed(), expected) @@ -124,10 +123,13 @@ def test_row_embedding_passes_on_wide_tables( categorical_mask = torch.arange(300, device="cuda") % 4 == 0 def embed() -> tuple[torch.Tensor, int]: + torch.cuda.synchronize() + baseline = torch.cuda.memory_allocated() torch.cuda.reset_peak_memory_stats() with torch.no_grad(), torch.autocast("cuda", dtype=torch.float16): out = encoder(x, y, categorical_mask) - return out, torch.cuda.max_memory_allocated() + torch.cuda.synchronize() + return out, torch.cuda.max_memory_allocated() - baseline monkeypatch.setenv("SDM_CHUNK_MEMORY_FRACTION", "0.001") actual, peak = embed() diff --git a/test/models/test_base.py b/test/models/test_base.py index 6c5171507..db05e933b 100644 --- a/test/models/test_base.py +++ b/test/models/test_base.py @@ -21,6 +21,7 @@ from sdm.models.callback import Callback from sdm.processing import InvertibleMixin, Processor from sdm.processing.execution import RecipeExecution +from sdm.testing import withCUDA @dataclass @@ -320,6 +321,52 @@ def test_model_recipe_generator_does_not_advance_global_rng( assert torch.equal(torch.get_rng_state(), state) +@withCUDA +@pytest.mark.parametrize("cached", [False, True]) +@pytest.mark.parametrize("estimator_batch_size", [1, "auto"]) +def test_callback_query_dtype_is_preserved( + device: torch.device, + cached: bool, + estimator_batch_size: int | Literal["auto"], +) -> None: + class DoubleQuery(Callback): + def on_query_preprocessing_end( + self, + model: torch.nn.Module, + x: TableTensor, + related_tables: RelatedTables[TableTensor] | None, + ) -> tuple[TableTensor, RelatedTables[TableTensor] | None]: + return ( + x.replace_blocks(numerical=x.numerical.double() + 2**-40), + related_tables, + ) + + model = _RecordingModel().to(device) + context = torch.ones(4, 1, device=device) + target = torch.arange(4, device=device).float().unsqueeze(-1) + query = torch.ones(2, 1, device=device) + if cached: + model.fit( + x=context, + y=target, + num_estimators=2, + estimator_batch_size=estimator_batch_size, + ) + output = model.predict(query, callbacks=[DoubleQuery()]) + else: + output = model( + x_context=context, + y_context=target, + x_query=query, + num_estimators=2, + estimator_batch_size=estimator_batch_size, + callbacks=[DoubleQuery()], + ) + + expected = (query.double() + 2**-40).expand(2, -1, -1) + torch.testing.assert_close(output.numerical, expected, rtol=0, atol=0) + + def test_callback() -> None: events: list[str] = [] callbacks = ( @@ -892,7 +939,7 @@ def test_estimator_batching_fails_like_sequential_on_class_mismatch() -> None: member_table_ids=(0, 1), ) for estimator_batch_size in (1, None): - with pytest.raises(ValueError, match="same set of classes"): + with pytest.raises(ValueError, match="column names"): model( x, y, @@ -1079,7 +1126,7 @@ def test_forward_members_moves_contexts_to_query_device() -> None: for query in recipe.transform(x=x_query, related_tables=None) ] - outs = model._forward_members(contexts=contexts, queries=queries) + outs, _ = model._forward_members(contexts=contexts, queries=queries) assert all( cast(TableTensor, call.x_context).is_cuda for call in model.calls diff --git a/test/models/test_ecoc.py b/test/models/test_ecoc.py index 8563fd806..f0006c18d 100644 --- a/test/models/test_ecoc.py +++ b/test/models/test_ecoc.py @@ -161,24 +161,10 @@ def test_ecoc_members_draw_codebooks_in_order(device: torch.device) -> None: torch.testing.assert_close(replayed, expected) -@pytest.mark.parametrize("num_classes", [3, 10, 11, 100, 201]) -def test_ecoc_num_tasks(num_classes: int) -> None: +@pytest.mark.parametrize( + ("num_classes", "expected"), + [(3, 1), (10, 1), (11, 8), (100, 12), (201, 23)], +) +def test_ecoc_num_tasks(num_classes: int, expected: int) -> None: ecoc = ECOC(max_classes=10) - x = torch.randn(num_classes + 2, num_classes, dtype=torch.float64) - y = torch.arange(num_classes) - cache = Cache() - - ecoc( - model=MyModel(10), - x=x[:num_classes], - y=y, - num_classes=num_classes, - cache=cache, - ) - - num_tasks = ( - cast(Tensor, cache["ecoc_codebook"]).size(-2) - if num_classes > 10 - else 1 - ) - assert ecoc.num_tasks(num_classes) == num_tasks + assert ecoc.num_tasks(num_classes) == expected diff --git a/test/processing/numerical/test_sigma_clip.py b/test/processing/numerical/test_sigma_clip.py index 2ae704ff2..803792406 100644 --- a/test/processing/numerical/test_sigma_clip.py +++ b/test/processing/numerical/test_sigma_clip.py @@ -155,6 +155,31 @@ def test_clip_sigma_preserves_nonfinite( ) +@withCUDA +@pytest.mark.parametrize("fit_transform", [False, True]) +def test_clip_sigma_gradients( + device: torch.device, + fit_transform: bool, +) -> None: + context = torch.tensor( + [[1.0], [2.0], [4.0], [100.0]], + dtype=torch.float64, + device=device, + requires_grad=True, + ) + query = TableTensor.from_tensor(context.detach().clone()) + + def transform(values: torch.Tensor) -> torch.Tensor: + processor = ClipSigma(threshold=1.0) + table = TableTensor.from_tensor(values) + if fit_transform: + return processor.fit_transform(table).numerical + processor.fit(table) + return processor.transform(query).numerical + + assert torch.autograd.gradcheck(transform, (context,)) + + @onlyCUDA def test_clip_sigma_transform_passes_match_single_pass( monkeypatch: pytest.MonkeyPatch, diff --git a/test/processing/test_execution.py b/test/processing/test_execution.py index 7034e84f8..fa3a7b44f 100644 --- a/test/processing/test_execution.py +++ b/test/processing/test_execution.py @@ -14,7 +14,7 @@ TableTensor, ) from sdm.models import KumoTabular -from sdm.processing.execution import RecipeExecution, _transform_rows +from sdm.processing.execution import RecipeExecution from sdm.testing import onlyCUDA @@ -451,28 +451,29 @@ def run() -> tuple[tuple[TableTensor, ...], TableTensor]: assert torch.equal(output.numerical, expected_output.numerical) -def test_transform_rows_sizes_every_member( +@onlyCUDA +@pytest.mark.parametrize("member_ids", [(0, 1, 1), (1, 0, 0)]) +def test_row_passes_preserve_heterogeneous_members( + member_ids: tuple[int, ...], monkeypatch: pytest.MonkeyPatch, ) -> None: - item_bytes: list[int] = [] - - def record_size( - num_items: int, - bytes_per_item: int, - device: torch.device, - ) -> int: - del device - item_bytes.append(bytes_per_item) - return num_items - - monkeypatch.setattr("sdm.processing.execution.split_size", record_size) - narrow = TableTensor.from_tensor(torch.zeros(3, 1)) - wide = TableTensor.from_tensor(torch.zeros(3, 4)) - for tables, member_ids in ( - ((narrow, wide), (0, 1, 1)), - ((wide, narrow), (1, 0, 0)), - ): - table = EnsembleTable.from_tables(tables, member_ids) - _transform_rows(lambda value: value, table) - - assert item_bytes == [9 * torch.float64.itemsize] * 2 + tables = [ + TableTensor.from_tensor(torch.randn(20, columns, device="cuda")) + for columns in (1, 4) + ] + table = EnsembleTable.from_tables(tables, member_ids) + execution = RecipeExecution(Recipe(features=sp.Standardize())) + with torch.inference_mode(): + execution.fit_transform( + x=table, + y=tables[0], + related_tables=None, + num_members=3, + ) + expected = execution.transform(x=table, related_tables=None) + monkeypatch.setenv("SDM_CHUNK_MEMORY_FRACTION", "1e-12") + actual = execution.transform(x=table, related_tables=None) + + for query, reference in zip(actual, expected, strict=True): + assert query.x.columns == reference.x.columns + torch.testing.assert_close(query.x.numerical, reference.x.numerical)