diff --git a/sdm/models/base.py b/sdm/models/base.py index 83d9bc3f5..74df2dabc 100644 --- a/sdm/models/base.py +++ b/sdm/models/base.py @@ -189,14 +189,6 @@ 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"), @@ -473,14 +465,6 @@ 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"), diff --git a/sdm/models/kumo/tabular/recipe.py b/sdm/models/kumo/tabular/recipe.py index f20cad1f5..54b360f7c 100644 --- a/sdm/models/kumo/tabular/recipe.py +++ b/sdm/models/kumo/tabular/recipe.py @@ -64,6 +64,7 @@ def numerical_processor() -> sp.Sequential: sp.Softmax(), ], regression=[ + sp.InvertTarget(), sp.SortQuantiles(), sp.AverageEstimators(trim_fraction=0.2), ], diff --git a/sdm/models/tabfm/recipe.py b/sdm/models/tabfm/recipe.py index 9ba8a3cca..5214569f5 100644 --- a/sdm/models/tabfm/recipe.py +++ b/sdm/models/tabfm/recipe.py @@ -38,6 +38,7 @@ def default_recipe() -> sp.Recipe: # noqa: D103 numerical=sp.Standardize(), ), output=[ + sp.TaskDispatch(regression=sp.InvertTarget()), sp.AverageEstimators(), sp.TaskDispatch( classification=sp.Softmax(temperature=0.9), diff --git a/sdm/models/tabiclv2/recipe.py b/sdm/models/tabiclv2/recipe.py index 95952efa5..57928ebfa 100644 --- a/sdm/models/tabiclv2/recipe.py +++ b/sdm/models/tabiclv2/recipe.py @@ -39,7 +39,9 @@ def default_recipe() -> sp.Recipe: # noqa: D103 ), ], output=[ - sp.TaskDispatch(regression=sp.SortQuantiles()), + sp.TaskDispatch( + regression=[sp.InvertTarget(), sp.SortQuantiles()], + ), sp.AverageEstimators(), sp.TaskDispatch(classification=sp.Softmax(temperature=0.9)), ], diff --git a/sdm/processing/__init__.py b/sdm/processing/__init__.py index ad88657d6..2bf63f5a6 100644 --- a/sdm/processing/__init__.py +++ b/sdm/processing/__init__.py @@ -47,7 +47,12 @@ AddCategoryCounts, ) from sdm.processing.datetime import AddCalendarFields -from sdm.processing.output import AverageEstimators, Softmax, SortQuantiles +from sdm.processing.output import ( + InvertTarget, + AverageEstimators, + Softmax, + SortQuantiles, +) from sdm.processing.recipe import Recipe __all__ = [ @@ -89,6 +94,7 @@ "ImputeMode", "AddCategoryCounts", "AddCalendarFields", + "InvertTarget", "AverageEstimators", "Softmax", "SortQuantiles", diff --git a/sdm/processing/common/ensemble.py b/sdm/processing/common/ensemble.py index 4f5a2cb5d..77820d21d 100644 --- a/sdm/processing/common/ensemble.py +++ b/sdm/processing/common/ensemble.py @@ -24,6 +24,8 @@ class EnsembleProcessorAdapter(EnsembleProcessor, EnsembleInvertibleMixin): :class:`~sdm.processing.ensemble.EnsembleProcessor`. It copies and fits the processor separately for each group of compatible tables in an :class:`~sdm.EnsembleTable`. + Stateful processors require the same logical member count for fitting, + transformation, and inverse transformation. The wrapped processor must preserve row and leading dimensions as required by the :class:`~sdm.processing.base.Processor` contract. A processor that @@ -142,6 +144,11 @@ def _aligned_processors( self, ensemble_table: EnsembleTable, ) -> tuple[Processor, ...]: + if len(ensemble_table) != len(self._fitted_locations): + raise RuntimeError( + f"Expected {len(self._fitted_locations)} fitted ensemble " + f"members (got {len(ensemble_table)})" + ) if ensemble_table._locations == self._fitted_locations: return tuple(self._processors) diff --git a/sdm/processing/execution.py b/sdm/processing/execution.py index 3f03f4c77..03a7bbc96 100644 --- a/sdm/processing/execution.py +++ b/sdm/processing/execution.py @@ -12,7 +12,7 @@ import sdm.processing as sp from sdm import EnsembleTable, Recipe, RelatedTables, Stype, TableTensor -from sdm.processing import EnsembleInvertibleMixin, EnsembleProcessor +from sdm.processing import EnsembleProcessor class MemberContext(NamedTuple): @@ -68,6 +68,11 @@ def fit_transform( self._num_estimators = num_members self._y_locations = y._locations + for module in self.recipe.output.modules(): + if isinstance(module, sp.InvertTarget): + module._target = self.recipe.target + module._locations = y._locations + module._ndim = y[0].dim() task_dispatchers = tuple( module @@ -215,42 +220,6 @@ def transform( return tuple(members) - def inverse_transform_target( - self, - outputs: Sequence[TableTensor], - ) -> tuple[TableTensor, ...]: - """Invert fitted target transforms on member outputs.""" - assert len(outputs) == self.num_members - - # Reconstruct the group layout of the transformed target: - assert self._y_locations is not None - num_groups = max(group for group, _ in self._y_locations) + 1 - groups: list[list[TableTensor | None]] = [ - [] for _ in range(num_groups) - ] - for group_id, _ in self._y_locations: - groups[group_id].append(None) - for i, (group_id, position) in enumerate(self._y_locations): - groups[group_id][position] = outputs[i] - - table = EnsembleTable( - groups=[ - cast( - TableTensor, - group[0].unsqueeze(0) # type: ignore - if len(group) == 1 - else torch.stack(group, dim=0), # type: ignore - ) - for group in groups - ], - locations=self._y_locations, - ) - - if not isinstance(self.recipe.target, EnsembleInvertibleMixin): - raise RuntimeError("Target recipe is not invertible") - table = self.recipe.target.inverse_transform_ensemble(table) - return tuple(table[i] for i in range(len(table))) - def transform_output( self, outputs: Sequence[TableTensor], diff --git a/sdm/processing/output/__init__.py b/sdm/processing/output/__init__.py index 4107e6eab..525dd0cb5 100644 --- a/sdm/processing/output/__init__.py +++ b/sdm/processing/output/__init__.py @@ -3,11 +3,13 @@ """Output postprocessing transforms.""" +from sdm.processing.output.inverse import InvertTarget from sdm.processing.output.reduce import AverageEstimators from sdm.processing.output.softmax import Softmax from sdm.processing.output.sort import SortQuantiles __all__ = [ + "InvertTarget", "AverageEstimators", "Softmax", "SortQuantiles", diff --git a/sdm/processing/output/inverse.py b/sdm/processing/output/inverse.py new file mode 100644 index 000000000..50ed7428c --- /dev/null +++ b/sdm/processing/output/inverse.py @@ -0,0 +1,70 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from typing import cast + +import torch + +from sdm import EnsembleTable, Stype, TableTensor +from sdm.processing import EnsembleInvertibleMixin, EnsembleProcessor + + +class InvertTarget(EnsembleProcessor): + """Invert the fitted target pipeline. + + Bound to ``Recipe.target`` during model execution. Member assignments must + remain compatible with its fitted state. + """ + + handles_stypes = frozenset({Stype.numerical}) + requires_fit = False + + def __init__(self) -> None: + super().__init__() + self._target: EnsembleProcessor | None = None + self._locations: tuple[tuple[int, int], ...] = () + self._ndim = 0 + + def _transform(self, table: TableTensor) -> TableTensor: + if table.dim() == self._ndim: + return self._transform_ensemble( + EnsembleTable.from_table(table, num_members=1) + )[0] + + outputs = self._transform_ensemble( + EnsembleTable( + groups=(table,), + locations=tuple((0, i) for i in range(table.size(0))), + ) + ) + if outputs.num_groups == 1: + return outputs.expanded_group(0) + return cast( + TableTensor, + torch.stack(tuple(outputs[i] for i in range(len(outputs))), dim=0), + ) + + def _transform_ensemble( + self, + ensemble_table: EnsembleTable, + ) -> EnsembleTable: + if not isinstance(self._target, EnsembleInvertibleMixin): + raise RuntimeError("Target recipe is not invertible") + if ( + len(ensemble_table) == len(self._locations) + and ensemble_table._locations != self._locations + ): + # Restore fitted groups without changing logical member order. + groups: list[list[TableTensor]] = [ + [] for _ in range(max(i for i, _ in self._locations) + 1) + ] + for member_id, (group_id, _) in enumerate(self._locations): + groups[group_id].append(ensemble_table[member_id]) + ensemble_table = EnsembleTable( + groups=tuple( + cast(TableTensor, torch.stack(tuple(group), dim=0)) + for group in groups + ), + locations=self._locations, + ) + return self._target.inverse_transform_ensemble(ensemble_table) diff --git a/sdm/processing/recipe.py b/sdm/processing/recipe.py index 48754d872..23ca73225 100644 --- a/sdm/processing/recipe.py +++ b/sdm/processing/recipe.py @@ -19,11 +19,11 @@ class Recipe: relative to the model: - ``features``: model inputs, transformed before the model. - - ``target``: labels transformed forward before the model. Regression - predictions are inverted through this pipeline; classification outputs - are reconstructed from the fitted target categories instead. - - ``output``: transforms member outputs after they have been mapped to a - common class or target space and stacked as ``[E, ..., R, O]``. An + - ``target``: labels transformed forward before the model. Classification + outputs are reconstructed from the fitted target categories. + - ``output``: transforms member outputs stacked as ``[E, ..., R, O]``. + Include :class:`~sdm.processing.InvertTarget` to restore regression + predictions to the original target space at the chosen position. An explicit dimension-changing step such as :class:`~sdm.processing.AverageEstimators` removes ``E``; without one, the output remains stacked. Steps before the reducer must support @@ -39,7 +39,9 @@ class Recipe: Args: features: Steps applied to model inputs before model execution. target: Steps applied to labels before model execution. - output: Steps applied to stacked model outputs. + output: Steps applied to stacked model outputs. If ``None``, use + :class:`~sdm.processing.Identity`. Target inversion requires an + explicit :class:`~sdm.processing.InvertTarget` step. """ _features: EnsembleProcessor diff --git a/test/models/kumo/tabular/test_model.py b/test/models/kumo/tabular/test_model.py index ca17449f0..552d1e60e 100644 --- a/test/models/kumo/tabular/test_model.py +++ b/test/models/kumo/tabular/test_model.py @@ -87,7 +87,10 @@ def _recipe() -> sp.Recipe: return sp.Recipe( features=[sp.ToNumerical()], target=sp.StypeDispatch(numerical=sp.Standardize()), - output=[sp.AverageEstimators()], + output=[ + sp.TaskDispatch(regression=sp.InvertTarget()), + sp.AverageEstimators(), + ], ) diff --git a/test/models/kumo/tabular/test_recipe.py b/test/models/kumo/tabular/test_recipe.py index a8a455719..ed5ef9b44 100644 --- a/test/models/kumo/tabular/test_recipe.py +++ b/test/models/kumo/tabular/test_recipe.py @@ -4,9 +4,9 @@ import pytest import torch -import sdm.processing as sp from sdm import CategoricalTensor, EnsembleTable, Stype, TableTensor from sdm.models.kumo.tabular import KumoTabular +from sdm.processing.execution import RecipeExecution from sdm.testing import withCUDA @@ -130,23 +130,25 @@ def test_default_recipe_adds_category_counts(cardinality: int) -> None: ) -def test_default_recipe_reduces_outputs_per_task() -> None: - recipe = KumoTabular.default_recipe() - dispatch = next( - module - for module in recipe.output.modules() - if isinstance(module, sp.TaskDispatch) - ) - - dispatch._task = "regression" - output = recipe.output.transform( - TableTensor.from_tensor(torch.randn(8, 5, 9)) +@pytest.mark.parametrize("classification", [False, True]) +def test_default_recipe_reduces_outputs_per_task(classification: bool) -> None: + execution = RecipeExecution(KumoTabular.default_recipe()) + target = torch.arange(5).unsqueeze(-1) + if not classification: + target = target.float() + execution.fit_transform( + x=torch.randn(5, 2), + y=target, + related_tables=None, + num_members=8, ) - assert output.size() == (5, 9) - - dispatch._task = "classification" - output = recipe.output.transform( - TableTensor.from_tensor(torch.randn(8, 5, 3)) + num_outputs = 5 if classification else 9 + output = execution.transform_output( + [ + TableTensor.from_tensor(torch.randn(5, num_outputs)) + for _ in range(8) + ] ) - assert output.size() == (5, 3) - torch.testing.assert_close(output.numerical.sum(dim=-1), torch.ones(5)) + assert output.size() == (5, num_outputs) + if classification: + torch.testing.assert_close(output.numerical.sum(dim=-1), torch.ones(5)) diff --git a/test/models/test_base.py b/test/models/test_base.py index 96c796730..271951358 100644 --- a/test/models/test_base.py +++ b/test/models/test_base.py @@ -728,6 +728,60 @@ def _check(out: TableTensor) -> None: _check(model.predict(x)) +@pytest.mark.parametrize( + ("output", "expected"), + [ + (None, 0.5), + ([sp.InvertTarget(), sp.Clip(-1.0, 1.0)], 1.0), + ([sp.Clip(-1.0, 1.0), sp.InvertTarget()], 11.0), + ], +) +def test_target_inversion_follows_output_order( + output: list[Processor] | None, + expected: float, +) -> None: + model = _RecordingModel() + recipe = sp.Recipe(target=sp.Standardize(), output=output) + x_context = torch.zeros(2, 1) + y_context = torch.tensor([[8.0], [12.0]]) + x_query = torch.tensor([[0.5]]) + + prediction = model( + x_context, + y_context, + x_query, + recipe=recipe, + ) + + torch.testing.assert_close( + prediction.numerical, + torch.tensor([[[expected]]]), + ) + + +@pytest.mark.parametrize("separate", [False, True]) +def test_target_inversion_uses_each_members_fitted_state( + separate: bool, +) -> None: + model = _RecordingModel() + recipe = sp.Recipe(target=sp.Standardize(), output=sp.InvertTarget()) + x_context = torch.zeros(2, 2, 1) + y_context = torch.tensor([[[8.0], [12.0]], [[20.0], [28.0]]]) + if separate: + y_context = EnsembleTable.from_tables( + tables=tuple(TableTensor.from_tensor(y) for y in y_context), + member_table_ids=(0, 1), + ) + x_query = torch.ones(2, 1, 1) + + model.fit(x_context, y_context, recipe=recipe) + prediction = model.predict(x_query) + + torch.testing.assert_close( + prediction.numerical, torch.tensor([[[12.0]], [[28.0]]]) + ) + + def test_estimator_batching_does_not_stack_related_tables() -> None: model = _RecordingModel() x_context = _table([0.0, 2.0], [1, 2], value_column="feature") diff --git a/test/processing/test_contract.py b/test/processing/test_contract.py index 642a3d69a..72cac5f67 100644 --- a/test/processing/test_contract.py +++ b/test/processing/test_contract.py @@ -196,6 +196,7 @@ def test_all_public_processors_have_contract_cases() -> None: covered_processors = {type(case.processor) for case in PROCESSOR_CASES} # Their specialized behavior is covered in their dedicated test modules. specialized_processors = { + sp.InvertTarget, sp.TaskDispatch, sp.TableDispatch, sp.SentenceTransformer,