Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 0 additions & 16 deletions sdm/models/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"),
Expand Down Expand Up @@ -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"),
Expand Down
1 change: 1 addition & 0 deletions sdm/models/kumo/tabular/recipe.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,7 @@ def numerical_processor() -> sp.Sequential:
sp.Softmax(),
],
regression=[
sp.InvertTarget(),
sp.SortQuantiles(),
sp.AverageEstimators(trim_fraction=0.2),
],
Expand Down
1 change: 1 addition & 0 deletions sdm/models/tabfm/recipe.py
Original file line number Diff line number Diff line change
Expand Up @@ -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),
Expand Down
4 changes: 3 additions & 1 deletion sdm/models/tabiclv2/recipe.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)),
],
Expand Down
8 changes: 7 additions & 1 deletion sdm/processing/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__ = [
Expand Down Expand Up @@ -89,6 +94,7 @@
"ImputeMode",
"AddCategoryCounts",
"AddCalendarFields",
"InvertTarget",
"AverageEstimators",
"Softmax",
"SortQuantiles",
Expand Down
7 changes: 7 additions & 0 deletions sdm/processing/common/ensemble.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)

Expand Down
43 changes: 6 additions & 37 deletions sdm/processing/execution.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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],
Expand Down
2 changes: 2 additions & 0 deletions sdm/processing/output/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
70 changes: 70 additions & 0 deletions sdm/processing/output/inverse.py
Original file line number Diff line number Diff line change
@@ -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)
14 changes: 8 additions & 6 deletions sdm/processing/recipe.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand Down
5 changes: 4 additions & 1 deletion test/models/kumo/tabular/test_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(),
],
)


Expand Down
40 changes: 21 additions & 19 deletions test/models/kumo/tabular/test_recipe.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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))
Loading
Loading