From a4a75ea255e326274d389f7370a9e375a0031ef1 Mon Sep 17 00:00:00 2001 From: Jingang Qu Date: Mon, 28 Sep 2026 23:26:29 +0200 Subject: [PATCH] Chunk recipe execution and output processing - Keep member outputs in the model output dtype until post-processing, which casts them to the callback-prepared query dtype and inverts numerical targets. - Transform queries and post-process member outputs in passes over rows within the chunk memory limit in `RecipeExecution`. - Require equal row counts across ensemble groups only when passes split them. - Post-process the benchmark adapter's outputs through `transform_output` after freeing the transformed queries. Signed-off-by: Jingang Qu Co-authored-by: Cedric Lorenz --- benchmark/tabular/model.py | 12 +--- sdm/models/base.py | 55 +++++++----------- sdm/processing/execution.py | 92 +++++++++++++++++++++++++++-- sdm/processing/recipe.py | 4 ++ test/models/test_base.py | 49 +++++++++++++++- test/processing/test_execution.py | 97 +++++++++++++++++++++++++++++++ 6 files changed, 261 insertions(+), 48 deletions(-) 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/models/base.py b/sdm/models/base.py index a7533a58d..059c3cd31 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,11 @@ 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) def fit( self, @@ -440,6 +432,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 @@ -463,6 +456,7 @@ def predict( strict=True, ) ] + dtypes.extend(query.x.dtype for query in members) start += len(x_schemas) if i + 1 < len(caches): @@ -527,19 +521,11 @@ 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) def clear(self) -> None: r"""Clear cached context state created by :meth:`fit`.""" @@ -638,13 +624,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 @@ -654,18 +640,19 @@ def _forward_members( 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 contexts = [ self._prepare_context(context, callbacks) for context in contexts @@ -709,7 +696,7 @@ def _forward_members( generator=generator, **kwargs, ) - return outs + return outs, [query.x.dtype for query in queries] def _prepare_context( self, @@ -788,7 +775,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..559af14fc 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,8 +258,46 @@ 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: @@ -268,11 +310,53 @@ def transform_output( "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/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/test_base.py b/test/models/test_base.py index d5cdf7d16..275ec5262 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 = ( @@ -990,7 +1037,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/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)