diff --git a/benchmark/tabular/model.py b/benchmark/tabular/model.py index 89b08811a..0f1ce7164 100644 --- a/benchmark/tabular/model.py +++ b/benchmark/tabular/model.py @@ -222,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()[ @@ -230,14 +230,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, dtypes) if self.problem_type == REGRESSION: return out.numerical.float().mean(dim=-1).cpu().numpy() diff --git a/sdm/_memory.py b/sdm/_memory.py new file mode 100644 index 000000000..68806ea50 --- /dev/null +++ b/sdm/_memory.py @@ -0,0 +1,33 @@ +# 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")) + ) + + +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/models/base.py b/sdm/models/base.py index 75c019738..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, @@ -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, + dtypes=dtypes, + ) def fit( self, @@ -442,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 @@ -465,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): @@ -529,19 +526,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, + dtypes=dtypes, + ) def clear(self) -> None: r"""Clear cached context state created by :meth:`fit`.""" @@ -644,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 @@ -661,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 @@ -716,7 +709,7 @@ def _forward_members( generator=generator, **kwargs, ) - return outs + return outs, [query.x.dtype for query in queries] def _prepare_context( self, @@ -796,7 +789,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/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/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, 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/execution.py b/sdm/processing/execution.py index 3f03f4c77..7da820264 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 @@ -166,7 +170,7 @@ def transform( ) -> tuple[MemberQuery, ...]: """Transform query data.""" 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,25 +258,96 @@ def inverse_transform_target( def transform_output( self, outputs: Sequence[TableTensor], + dtypes: Sequence[torch.dtype], ) -> TableTensor: - """Apply ``recipe.output`` to member outputs.""" + """Apply ``recipe.output`` after restoring each member dtype. + + Numerical targets are inverse-transformed before applying the output + recipe. + """ + # 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, dtypes) + parts = tuple( + self._transform_output(chunk, dtypes) + 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], + dtypes: Sequence[torch.dtype], + ) -> TableTensor: + outputs = tuple( + cast(TableTensor, output.to(dtype)) + for output, dtype in zip(outputs, dtypes, strict=True) + ) + if self._numerical_target: + 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)) +def _transform_rows( + transform: Callable[[EnsembleTable], EnsembleTable], + table: EnsembleTable, +) -> EnsembleTable: + # Member cells are transformed in double precision at most. + # 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 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( + *(group.split(rows_per_pass, 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/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..1d6ec874c 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,21 @@ 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) + if torch.is_grad_enabled() and numerical.requires_grad: + 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) + + 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..c6a77b957 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,42 @@ 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) + 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) - # 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 +94,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/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/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/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..76d69fa63 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,98 @@ 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() + # Small, intermediate and single-pass budgets exercise pass boundaries. + total_memory = torch.cuda.get_device_properties("cuda").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) + + +@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.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) + torch.cuda.synchronize() + return out, torch.cuda.max_memory_allocated() - baseline + + 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 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_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..803792406 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,56 @@ def test_clip_sigma_preserves_nonfinite( ), equal_nan=True, ) + + +@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, +) -> 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, diff --git a/test/processing/test_execution.py b/test/processing/test_execution.py index ef8b55b8a..fa3a7b44f 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,98 @@ 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,) * len(outputs), + ) + 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) + + +@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: + 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)