From 64e9eb94c8b1e6af6f17d9ff011ec264853fb8a9 Mon Sep 17 00:00:00 2001 From: Akihiro Nitta Date: Fri, 25 Sep 2026 20:57:41 +0000 Subject: [PATCH] Batch inference across estimators --- benchmark/tabular/model.py | 57 ++- sdm/explain/gradient.py | 1 + sdm/models/base.py | 499 +++++++++++++++++++------ sdm/models/kumo/tabular/model.py | 35 +- sdm/models/tabfm/model.py | 35 +- sdm/processing/execution.py | 10 +- test/explain/test_gradient.py | 34 ++ test/models/kumo/tabular/test_model.py | 54 +++ test/models/tabfm/test_model.py | 33 +- test/models/tabiclv2/test_model.py | 37 +- test/models/test_base.py | 355 +++++++++++++++++- test/processing/test_execution.py | 28 +- 12 files changed, 980 insertions(+), 198 deletions(-) diff --git a/benchmark/tabular/model.py b/benchmark/tabular/model.py index 5a5caead5..923815410 100644 --- a/benchmark/tabular/model.py +++ b/benchmark/tabular/model.py @@ -62,6 +62,7 @@ def _set_default_params(self) -> None: self._set_default_param_value("max_context_size", None) self._set_default_param_value("max_columns", None) self._set_default_param_value("kv_cache", False) + self._set_default_param_value("estimator_batch_size", "auto") def _fit( self, @@ -127,6 +128,7 @@ def _fit( y_context = y_context[perm].unflatten(0, shape) num_estimators = None self._expand_query = num_estimators is None + self._context_shape = x_context.shape[-2:] recipe = self._create_recipe() if params["max_columns"] is not None: @@ -146,13 +148,12 @@ def _fit( y=y_context, recipe=recipe, num_estimators=num_estimators, + estimator_batch_size=self._estimator_batch_size(x_context), generator=generator, ) return self._recipe_execution = RecipeExecution(recipe) - # KumoTabular and TabFM need the column types before preprocessing. - self._schema = x_context.schema with ( torch.inference_mode(), torch.amp.autocast(self._device.type, enabled=False), @@ -215,24 +216,32 @@ def _predict_proba( generator = torch.Generator(self._device).set_state( self._rng_state ) - outputs = [] - for context, query in zip(self._contexts, queries): + outputs: list[sdm.TableTensor] = [] + size = self._estimator_batch_size(x_query) + size = len(self._contexts) if size is None else size + for start in range(0, len(self._contexts), size): + batch = [ + context._replace( + x=cast( + sdm.TableTensor, context.x.to(self._device) + ), + y=cast( + sdm.TableTensor, context.y.to(self._device) + ), + ) + for context in self._contexts[start : start + size] + ] with torch.amp.autocast( self._device.type, self.autocast_dtype, enabled=x_query.is_cuda, ): - out = self.model._forward( - x_context=context.x.to(self._device), - y_context=context.y.to(self._device), - x_query=query.x, - related_context_tables=None, - related_query_tables=None, - cache=None, + outputs += self.model._forward_members( + contexts=batch, + queries=queries[start : start + size], + estimator_batch_size=None, generator=generator, - _schema=self._schema, ) - outputs.append(out.to(query.x.dtype)) if self.problem_type == REGRESSION: outputs = list( @@ -251,6 +260,28 @@ def _predict_proba( probabilities = out.numerical[..., indices].float().cpu().numpy() return self._convert_proba_to_unified_form(probabilities) + def _estimator_batch_size(self, x: torch.Tensor) -> int | None: + params = self._get_model_params() + estimator_batch_size = params["estimator_batch_size"] + if estimator_batch_size != "auto": + return estimator_batch_size + # Subsampled contexts can produce different cache shapes per estimator. + if not x.is_cuda or self._expand_query: + return 1 + + num_rows, num_cols = self._context_shape + if not params["kv_cache"]: + # Uncached inference processes context and query rows together. + num_rows += x.size(-2) + if num_rows > 3_000 or num_rows * num_cols > 50_000: + return 1 + return self._num_estimators + + num_rows = max(num_rows, x.size(-2)) + if num_rows > 2_000 or num_rows * num_cols >= 50_000: + return 1 + return self._num_estimators + def get_device(self) -> str: return str(next(self.model.parameters()).device) diff --git a/sdm/explain/gradient.py b/sdm/explain/gradient.py index 60e11f4c4..8171b5b96 100644 --- a/sdm/explain/gradient.py +++ b/sdm/explain/gradient.py @@ -78,6 +78,7 @@ def on_model_forward_end( objective, [numerical for _, _, numerical in self._inputs], allow_unused=True, + retain_graph=True, ) grad_tables: dict[str | None, TableTensor] = {} for (table_name, columns, numerical), grad in zip( diff --git a/sdm/models/base.py b/sdm/models/base.py index 43293d194..43228a0b9 100644 --- a/sdm/models/base.py +++ b/sdm/models/base.py @@ -3,7 +3,7 @@ import abc import copy -from collections.abc import Iterable, Sequence +from collections.abc import Iterable, Mapping, Sequence from typing import Any, ClassVar, cast import torch @@ -14,6 +14,7 @@ Recipe, RelatedTables, Stype, + StypeLike, TableTensor, Task, TaskLike, @@ -102,6 +103,7 @@ def forward( *, recipe: Recipe | None = None, num_estimators: int | None = None, + estimator_batch_size: int | None = 1, callbacks: Sequence[Callback] | None = None, generator: torch.Generator | None = None, **kwargs: Any, @@ -124,6 +126,16 @@ def forward( used as the estimator dimension, allowing input data to be customized per estimator (*e.g.*, different in-context examples per estimator). + estimator_batch_size: Maximum number of estimators run through the + model in one call. ``1`` (default) runs estimators one by one; + ``None`` runs all of them together. Estimators in one batch + must share column names, table shapes, category counts and + classes, and related tables require ``1``. Device memory grows + with the batch size. + Model-side randomness drawn per call (*e.g.*, the ECOC codebook + of :class:`~sdm.models.KumoTabular` for more than 10 classes) + is shared within a batch, so batched and sequential predictions + differ numerically there. callbacks: Callbacks applied in sequence to this model call. generator: Pseudorandom number generator used for sampling during pre-processing and model execution. @@ -172,48 +184,14 @@ def forward( related_tables=related_query_tables, ) - outs: list[TableTensor] = [] - for context, query in zip(contexts, queries): - for callback in callbacks: - context = MemberContext( - *callback.on_context_preprocessing_end(self, *context) - ) - self._validate_context( - x=context.x, - y=context.y, - related_tables=context.related_tables, - ) - - for callback in callbacks: - query = MemberQuery( - *callback.on_query_preprocessing_end(self, *query) - ) - self._validate_query( - x_context=context.x.schema, - x_query=query.x, - related_context_tables=context.related_tables.schema - if context.related_tables is not None - else None, - related_query_tables=query.related_tables, - ) - - with inference_mode("grad" if requires_grad else "inference"): - out = self._forward( - x_context=context.x, - y_context=context.y, - x_query=query.x, - related_context_tables=context.related_tables, - related_query_tables=query.related_tables, - cache=None, - generator=generator, - **kwargs, - ) - - for callback in callbacks: - out = callback.on_model_forward_end(self, out) - - out = cast(TableTensor, out.to(query.x.dtype)) - outs.append(out) + outs = self._forward_members( + contexts=contexts, + queries=queries, + estimator_batch_size=estimator_batch_size, + callbacks=callbacks, + generator=generator, + **kwargs, + ) # Regression: invert target before stacking estimator outputs. if contexts[0].y.numerical.size(-1) > 0: @@ -237,6 +215,7 @@ def fit( *, recipe: Recipe | None = None, num_estimators: int | None = None, + estimator_batch_size: int | None = 1, callbacks: Sequence[Callback] | None = None, generator: torch.Generator | None = None, **kwargs: Any, @@ -258,6 +237,17 @@ def fit( used as the estimator dimension, allowing input data to be customized per estimator (*e.g.*, different in-context examples per estimator). + estimator_batch_size: Maximum number of estimators run through the + model in one call. ``1`` (default) runs estimators one by one; + ``None`` runs all of them together. Estimators in one batch + must share column names, table shapes, category counts and + classes, and related tables require ``1``. Device memory grows + with the batch size. Estimators fitted together are predicted + together. + Model-side randomness drawn per call (*e.g.*, the ECOC codebook + of :class:`~sdm.models.KumoTabular` for more than 10 classes) + is shared within a batch, so batched and sequential predictions + differ numerically there. callbacks: Callbacks applied in sequence to this model call. generator: Pseudorandom number generator used for sampling during pre-processing and model execution. @@ -282,48 +272,52 @@ def fit( generator=generator, ) + if estimator_batch_size is None: + estimator_batch_size = len(contexts) + batches = range(0, len(contexts), estimator_batch_size) cache = Cache( recipe_execution=recipe_execution, kwargs=kwargs, + estimator_batch_size=estimator_batch_size, ) - for i, context in enumerate(contexts): - for callback in callbacks: - context = MemberContext( - *callback.on_context_preprocessing_end(self, *context) - ) - self._validate_context( - x=context.x, - y=context.y, - related_tables=context.related_tables, - ) - estimator_cache = Cache( - x_schema=context.x.schema, - y_schema=context.y.schema, - related_tables_schema=context.related_tables.schema - if context.related_tables is not None - else None, - classes=( - context.y.categorical.categories[0] - if context.y.categorical.size(-1) > 0 - else None - ), - ) - + for i, start in enumerate(batches): + members = [ + self._prepare_context(context, callbacks) + for context in contexts[start : start + estimator_batch_size] + ] with inference_mode("no_grad"): + class_values = _class_values(members) + context = _stack_context(members, class_values) + categorical_mask = _categorical_mask(members) + batch_cache = Cache( + x_schemas=tuple(member.x.schema for member in members), + y_schema=context.y.schema, + related_tables_schema=context.related_tables.schema + if context.related_tables is not None + else None, + classes=( + context.y.categorical.categories[0] + if context.y.categorical.size(-1) > 0 + else None + ), + class_values=class_values, + categorical_mask=categorical_mask, + ) self._forward( x_context=context.x, y_context=context.y, x_query=None, related_context_tables=context.related_tables, related_query_tables=None, - cache=estimator_cache, + cache=batch_cache, generator=generator, + categorical_mask=categorical_mask, **kwargs, ) if x.is_cuda and len(contexts) > 1: try: # Copy to pinned CPU memory: - estimator_cache = estimator_cache._apply_tensor( + batch_cache = batch_cache._apply_tensor( lambda tensor: torch.ops.aten._to_copy.default( tensor, device="cpu", @@ -334,7 +328,7 @@ def fit( finally: torch.cuda.current_stream(x.device).synchronize() - cache[i] = estimator_cache + cache[i] = batch_cache self._cache = cache.freeze() @@ -393,10 +387,9 @@ def predict( RecipeExecution, self._cache["recipe_execution"], ) - caches = [ - cast(Cache, self._cache[i]) - for i in range(recipe_execution.num_members) - ] + size = cast(int, self._cache["estimator_batch_size"]) + num_batches = len(range(0, recipe_execution.num_members, size)) + caches = [cast(Cache, self._cache[i]) for i in range(num_batches)] next_cache = caches[0] compute_stream: torch.cuda.Stream | None = None @@ -424,23 +417,29 @@ def predict( compute_stream.wait_stream(transfer_stream) outs: list[TableTensor] = [] - for i, query in enumerate(queries): + start = 0 + for i in range(len(caches)): cache, next_cache = next_cache, None assert cache is not None - for callback in callbacks: - query = MemberQuery( - *callback.on_query_preprocessing_end(self, *query) + x_schemas = cast(tuple[TableSchema, ...], cache["x_schemas"]) + members = [ + self._prepare_query( + query=query, + x_schema=x_schema, + related_tables_schema=cast( + RelatedTablesSchema | None, + cache["related_tables_schema"], + ), + callbacks=callbacks, ) - self._validate_query( - x_context=cast(TableSchema, cache["x_schema"]), - x_query=query.x, - related_context_tables=cast( - RelatedTablesSchema, - cache["related_tables_schema"], - ), - related_query_tables=query.related_tables, - ) + for query, x_schema in zip( + queries[start : start + len(x_schemas)], + x_schemas, + strict=True, + ) + ] + start += len(x_schemas) if i + 1 < len(caches): next_cache = caches[i + 1] @@ -449,29 +448,26 @@ def predict( with torch.cuda.stream(transfer_stream): next_cache = next_cache.to(x.device, non_blocking=True) - with inference_mode("grad" if requires_grad else "inference"): - out = self._forward( - x_context=None, - y_context=None, - x_query=query.x, - related_context_tables=None, - related_query_tables=query.related_tables, - cache=cache, - generator=None, - **cast(dict[str, Any], self._cache["kwargs"]), - ) - - for callback in callbacks: - out = callback.on_model_forward_end(self, out) + outs += self._forward_batch( + contexts=None, + queries=members, + cache=cache, + categorical_mask=cast(Tensor, cache["categorical_mask"]), + class_values=cast( + tuple[tuple[Any, ...], ...] | None, + cache["class_values"], + ), + callbacks=callbacks, + requires_grad=requires_grad, + generator=None, + **cast(dict[str, Any], self._cache["kwargs"]), + ) if x.is_cuda: assert compute_stream is not None for tensor in cache._tensors(): tensor.record_stream(compute_stream) - out = cast(TableTensor, out.to(query.x.dtype)) - outs.append(out) - if x.is_cuda and next_cache is not None: assert compute_stream is not None assert transfer_stream is not None @@ -528,9 +524,44 @@ def _forward( related_query_tables: RelatedTables[TableTensor] | None, cache: Cache | None, generator: torch.Generator | None, + *, + categorical_mask: Tensor, **kwargs: Any, ) -> TableTensor: # [..., R_query, *] - pass + r"""Run the model on preprocessed tables of one estimator batch. + + Tables carry shape ``[E, ..., R, D]`` when ``E > 1`` estimators run + together and ``[..., R, D]`` otherwise. Column names and category + vocabularies are those of the first estimator; per-estimator column + and class order is not observable from the tables. + + Args: + x_context: The feature tensor of in-context examples, or ``None`` + when replaying a cache. + y_context: The targets of in-context examples, or ``None`` when + replaying a cache. + x_query: The feature tensor of query examples, or ``None`` when + recording a cache. + related_context_tables: Related context for in-context examples. + related_query_tables: Related context for query examples. + cache: The cache to record into or replay from, or ``None``. + Recorded tensors are replayed with queries of the same + estimator batch. + generator: Pseudorandom number generator for model execution. + categorical_mask: Boolean ``[C]`` or ``[E, 1, ..., C]`` tensor + broadcastable to ``x.numerical.size()[:-2] + (C,)`` marking + numerical feature columns that were categorical before + preprocessing. Passed on recording, replaying and uncached + calls alike. + kwargs: Additional keyword arguments passed by the caller. + + Returns: + Predictions of shape ``[..., R_query, *]``. Classification outputs + hold exactly one column per class in the order of + ``y_context.categorical.categories[0]`` (``cache["classes"]`` on + replay), each named by its class value; :class:`ICLModel` relabels + them per estimator when stacked. + """ @classmethod @abc.abstractmethod @@ -539,6 +570,143 @@ def default_recipe(cls) -> Recipe: # Helpers ################################################################# + def _forward_members( + self, + contexts: Sequence[MemberContext], + queries: Sequence[MemberQuery], + *, + estimator_batch_size: int | None = 1, + callbacks: Sequence[Callback] | None = None, + generator: torch.Generator | None = None, + **kwargs: Any, + ) -> list[TableTensor]: + r"""Run recipe-transformed members that live on the model device. + + Returns one output per member before target inversion and + ``recipe.output``. + """ + callbacks = () if callbacks is None else callbacks + requires_grad = self.training + requires_grad |= any(callback.requires_grad for callback in callbacks) + if estimator_batch_size is None: + estimator_batch_size = len(contexts) + outs: list[TableTensor] = [] + for start in range(0, len(contexts), estimator_batch_size): + members = [ + self._prepare_context(context, callbacks) + for context in contexts[start : start + estimator_batch_size] + ] + query_members = [ + self._prepare_query( + query=query, + x_schema=member.x.schema, + related_tables_schema=member.related_tables.schema + if member.related_tables is not None + else None, + callbacks=callbacks, + ) + for member, query in zip( + members, + queries[start : start + estimator_batch_size], + strict=True, + ) + ] + outs += self._forward_batch( + contexts=members, + queries=query_members, + cache=None, + categorical_mask=None, + class_values=_class_values(members), + callbacks=callbacks, + requires_grad=requires_grad, + generator=generator, + **kwargs, + ) + return outs + + def _prepare_context( + self, + context: MemberContext, + callbacks: Sequence[Callback], + ) -> MemberContext: + for callback in callbacks: + x, y, related_tables = callback.on_context_preprocessing_end( + self, + context.x, + context.y, + context.related_tables, + ) + context = context._replace(x=x, y=y, related_tables=related_tables) + self._validate_context( + x=context.x, + y=context.y, + related_tables=context.related_tables, + ) + return context + + def _prepare_query( + self, + query: MemberQuery, + x_schema: TableSchema, + related_tables_schema: RelatedTablesSchema | None, + callbacks: Sequence[Callback], + ) -> MemberQuery: + for callback in callbacks: + query = MemberQuery( + *callback.on_query_preprocessing_end(self, *query) + ) + self._validate_query( + x_context=x_schema, + x_query=query.x, + related_context_tables=related_tables_schema, + related_query_tables=query.related_tables, + ) + return query + + def _forward_batch( + self, + contexts: Sequence[MemberContext] | None, + queries: Sequence[MemberQuery], + *, + cache: Cache | None, + categorical_mask: Tensor | None, + class_values: tuple[tuple[Any, ...], ...] | None, + callbacks: Sequence[Callback], + requires_grad: bool, + generator: torch.Generator | None, + **kwargs: Any, + ) -> list[TableTensor]: + # Stacking inside the autograd region keeps callback-captured leaves + # attached to the graph. + with inference_mode("grad" if requires_grad else "inference"): + context = ( + None + if contexts is None + else _stack_context(contexts, class_values) + ) + if categorical_mask is None: + assert contexts is not None + categorical_mask = _categorical_mask(contexts) + query = _stack_query(queries) + out = self._forward( + x_context=None if context is None else context.x, + y_context=None if context is None else context.y, + x_query=query.x, + related_context_tables=( + None if context is None else context.related_tables + ), + related_query_tables=query.related_tables, + cache=cache, + generator=generator, + categorical_mask=categorical_mask, + **kwargs, + ) + outs = _unstack(out, class_values, len(queries)) + 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] + def _validate_context( self, x: TableTensor, @@ -632,3 +800,126 @@ def _validate_query( "Expected related context and query tables to share the " "same schema" ) + + +def _stack(tables: Sequence[TableTensor]) -> TableTensor: + ref = tables[0] + if len(tables) == 1: + return ref + # torch.stack aligns columns by name, which would undo per-estimator column + # shuffles; renaming to the first member's names stacks blocks by position. + columns = cast(Mapping[StypeLike, Sequence[str]], ref.columns) + renamed: list[Tensor] = [ + ref, + *( + table.__class__(columns=columns, **dict(table.items())) + for table in tables[1:] + ), + ] + return cast(TableTensor, torch.stack(renamed)) + + +_INCOMPATIBLE_ESTIMATORS = ( + "Estimators in one batch must share column names, category counts and " + "classes; use 'estimator_batch_size=1'" +) + + +def _check_compatible(tables: Sequence[TableTensor]) -> None: + ref = tables[0] + names = {stype: frozenset(names) for stype, names in ref.columns.items()} + counts = tuple(c.numel() for c in ref.categorical.categories) + for table in tables[1:]: + if { + stype: frozenset(names) for stype, names in table.columns.items() + } != names or ( + tuple(c.numel() for c in table.categorical.categories) != counts + ): + raise ValueError(_INCOMPATIBLE_ESTIMATORS) + + +def _stack_context( + members: Sequence[MemberContext], + class_values: tuple[tuple[Any, ...], ...] | None, +) -> MemberContext: + if len(members) == 1: + return members[0] + if members[0].related_tables is not None: + raise ValueError("Related tables require 'estimator_batch_size=1'") + xs = [member.x for member in members] + ys = [member.y for member in members] + _check_compatible(xs) + _check_compatible(ys) + if class_values is not None and any( + set(values) != set(class_values[0]) for values in class_values[1:] + ): + raise ValueError(_INCOMPATIBLE_ESTIMATORS) + return MemberContext( + x=_stack(xs), + y=_stack(ys), + related_tables=None, + input_stypes=members[0].input_stypes, + ) + + +def _stack_query(members: Sequence[MemberQuery]) -> MemberQuery: + if len(members) == 1: + return members[0] + xs = [member.x for member in members] + _check_compatible(xs) + return MemberQuery(x=_stack(xs), related_tables=None) + + +def _categorical_mask(members: Sequence[MemberContext]) -> Tensor: + x = members[0].x + mask = torch.tensor( + [ + [ + member.input_stypes.get(column) == Stype.categorical + for column in member.x.columns[Stype.numerical] + ] + for member in members + ], + dtype=torch.bool, + device=x.device, + ) # [E, C] + if len(members) == 1: + return mask[0] + # Insert the member's batch dimensions so the mask broadcasts over them: + return mask.view(len(members), *(1,) * (x.dim() - 2), -1) # [E, 1, ..., C] + + +def _class_values( + members: Sequence[MemberContext], +) -> tuple[tuple[Any, ...], ...] | None: + if len(members) == 1 or members[0].y.categorical.size(-1) == 0: + return None + return tuple( + tuple(member.y.categorical.categories[0].tolist()) + for member in members + ) + + +def _unstack( + out: TableTensor, + class_values: tuple[tuple[Any, ...], ...] | None, + num_members: int, +) -> list[TableTensor]: + outs = ( + [out] + if num_members == 1 + else list(cast(tuple[TableTensor, ...], out.unbind(0))) + ) + if class_values is None: + return outs + # The model labels columns in the first member's class order; member `e`'s + # column `j` holds class `class_values[e][j]`. + labels = out.columns[Stype.numerical] + index = {value: i for i, value in enumerate(class_values[0])} + return [ + TableTensor( + columns={Stype.numerical: [labels[index[v]] for v in values]}, + numerical=member_out.numerical, + ) + for member_out, values in zip(outs, class_values, strict=True) + ] diff --git a/sdm/models/kumo/tabular/model.py b/sdm/models/kumo/tabular/model.py index 34f617cca..814bbc9ab 100644 --- a/sdm/models/kumo/tabular/model.py +++ b/sdm/models/kumo/tabular/model.py @@ -18,7 +18,6 @@ from sdm.models.kumo.tabular.icl import ICLBlock from sdm.models.kumo.tabular.recipe import default_recipe from sdm.models.kumo.tabular.row_embedding import RowEmbedding -from sdm.tensor.table import TableSchema MODEL_KWARGS: dict[str, dict[str, Any]] = { "small": { @@ -166,22 +165,6 @@ def _load_from_pretrained( return model - def forward(self, *args: Any, **kwargs: Any) -> TableTensor: - r""":meta private:""" # noqa: D415 - x_context = kwargs["x_context"] if "x_context" in kwargs else args[0] - if not isinstance(x_context, TableTensor): - x_context = TableTensor.from_tensor(x_context) - kwargs["_schema"] = x_context.schema - return super().forward(*args, **kwargs) - - def fit(self, *args: Any, **kwargs: Any) -> None: - r""":meta private:""" # noqa: D415 - x = kwargs["x"] if "x" in kwargs else args[0] - if not isinstance(x, TableTensor): - x = TableTensor.from_tensor(x) - kwargs["_schema"] = x.schema - return super().fit(*args, **kwargs) - def _forward( self, x_context: TableTensor | None, # [..., R_context, D] @@ -191,6 +174,8 @@ def _forward( related_query_tables: RelatedTables[TableTensor] | None, cache: Cache | None, generator: torch.Generator | None, + *, + categorical_mask: Tensor, **kwargs: Any, ) -> TableTensor: # [..., R_query, num_classes or 999] @@ -217,22 +202,6 @@ def _forward( dtype=torch.int64 if classes is not None else x.dtype, ) - if cache is None or cache.is_recording: - assert x_context is not None - schema: TableSchema = kwargs["_schema"] - categorical_columns = set(schema.columns[Stype.categorical]) - categorical_mask = torch.tensor( - [ - column in categorical_columns - for column in x_context.columns[Stype.numerical] - ], - device=x.device, - dtype=torch.bool, - ) - if cache is not None: - cache["categorical_mask"] = categorical_mask - else: - categorical_mask = cast(Tensor, cache["categorical_mask"]) categorical_mask = categorical_mask.expand(*x.size()[:-2], -1) if classes is None: diff --git a/sdm/models/tabfm/model.py b/sdm/models/tabfm/model.py index 8b53228bf..39af1437b 100644 --- a/sdm/models/tabfm/model.py +++ b/sdm/models/tabfm/model.py @@ -34,7 +34,6 @@ from sdm.models.tabfm.icl import ICLBlock from sdm.models.tabfm.recipe import default_recipe from sdm.models.tabfm.row_embedding import RowEmbedding -from sdm.tensor.table import TableSchema class TabFM(ICLModel): @@ -147,22 +146,6 @@ def _load_from_pretrained( return self - def forward(self, *args: Any, **kwargs: Any) -> TableTensor: - r""":meta private:""" # noqa: D415 - x_context = kwargs["x_context"] if "x_context" in kwargs else args[0] - if not isinstance(x_context, TableTensor): - x_context = TableTensor.from_tensor(x_context) - kwargs["_schema"] = x_context.schema - return super().forward(*args, **kwargs) - - def fit(self, *args: Any, **kwargs: Any) -> None: - r""":meta private:""" # noqa: D415 - x = kwargs["x"] if "x" in kwargs else args[0] - if not isinstance(x, TableTensor): - x = TableTensor.from_tensor(x) - kwargs["_schema"] = x.schema - return super().fit(*args, **kwargs) - def _forward( self, x_context: TableTensor | None, # [..., R_context, D] @@ -172,6 +155,8 @@ def _forward( related_query_tables: RelatedTables[TableTensor] | None, cache: Cache | None, generator: torch.Generator | None, + *, + categorical_mask: Tensor, **kwargs: Any, ) -> TableTensor: # [..., R_query, num_classes or 1] @@ -204,22 +189,6 @@ def _forward( f"(got {len(classes)})" ) - if cache is None or cache.is_recording: - assert x_context is not None - schema: TableSchema = kwargs["_schema"] - categorical_columns = set(schema.columns[Stype.categorical]) - categorical_mask = torch.tensor( - [ - column in categorical_columns - for column in x_context.columns[Stype.numerical] - ], - device=x.device, - dtype=torch.bool, - ) - if cache is not None: - cache["categorical_mask"] = categorical_mask - else: - categorical_mask = cast(Tensor, cache["categorical_mask"]) categorical_mask = categorical_mask.expand(*x.size()[:-2], -1) task = Task.classification if classes is not None else Task.regression diff --git a/sdm/processing/execution.py b/sdm/processing/execution.py index 25b4fe9f6..3f03f4c77 100644 --- a/sdm/processing/execution.py +++ b/sdm/processing/execution.py @@ -21,6 +21,8 @@ class MemberContext(NamedTuple): x: TableTensor y: TableTensor related_tables: RelatedTables[TableTensor] | None + #: Semantic types of the raw context columns before ``recipe.features``. + input_stypes: Mapping[str, Stype] class MemberQuery(NamedTuple): @@ -124,8 +126,11 @@ def fit_transform( if isinstance(module, sp.TableDispatch): module._route = "task" - x = _to_ensemble_table(x, num_members) - x = self.recipe.features.fit_transform_ensemble(x, generator=generator) + inputs = _to_ensemble_table(x, num_members) + x = self.recipe.features.fit_transform_ensemble( + inputs, + generator=generator, + ) if len(x) != self.num_members: raise ValueError( "Expected inputs to map to the same number of ensemble members" @@ -148,6 +153,7 @@ def fit_transform( x=x[member_id], y=y[member_id], related_tables=related_tables_i, + input_stypes=inputs[member_id].stypes, ) ) diff --git a/test/explain/test_gradient.py b/test/explain/test_gradient.py index 914c4394d..2340284ba 100644 --- a/test/explain/test_gradient.py +++ b/test/explain/test_gradient.py @@ -102,3 +102,37 @@ def test_returns_query_input_gradients(fitted: bool) -> None: ) assert result.related_tables.relationships == related_tables.relationships assert result.related_tables.task_links == related_tables.task_links + + +@pytest.mark.parametrize("fitted", [False, True]) +def test_gradients_with_estimator_batching(fitted: bool) -> None: + model = _LinearModel() + x_context = torch.zeros(1, 2) + y_context = torch.zeros(1, 1) + x_query = torch.ones(1, 2) + explainer = GradientExplainer( + output=lambda prediction: prediction.numerical + ) + + if fitted: + model.fit( + x=x_context, + y=y_context, + num_estimators=3, + estimator_batch_size=None, + ) + result = explainer.explain(model, x_query) + else: + result = explainer.explain( + model, + x_query, + x_context=x_context, + y_context=y_context, + num_estimators=3, + estimator_batch_size=None, + ) + + torch.testing.assert_close( + result.x.numerical, torch.full_like(x_query, 2.0) + ) + assert result.related_tables is None diff --git a/test/models/kumo/tabular/test_model.py b/test/models/kumo/tabular/test_model.py index 66bccbbff..48b2cf008 100644 --- a/test/models/kumo/tabular/test_model.py +++ b/test/models/kumo/tabular/test_model.py @@ -130,6 +130,60 @@ def test_categorical_features_are_marked(cls_model: KumoTabular) -> None: assert not categorical.allclose(numerical) +@pytest.mark.parametrize("task", ["classification", "regression"]) +@pytest.mark.parametrize("estimator_batch_size", [2, None]) +def test_estimator_batching( + task: Literal["classification", "regression"], + size: Literal["small", "medium", "large"], + estimator_batch_size: int | None, +) -> None: + model = _build(task, size) + x_context, x_query = _features() + target = _cls_target() if task == "classification" else _reg_target() + # Match the shuffled features and target categories in both executions. + expected = model( + x_context=x_context, + y_context=target, + x_query=x_query, + num_estimators=5, + generator=torch.Generator().manual_seed(0), + ) + + actual = model( + x_context=x_context, + y_context=target, + x_query=x_query, + num_estimators=5, + estimator_batch_size=estimator_batch_size, + generator=torch.Generator().manual_seed(0), + ) + assert actual.columns == expected.columns + assert actual.dtype == expected.dtype + assert torch.is_inference(actual) + torch.testing.assert_close( + actual=actual.numerical, + expected=expected.numerical, + atol=1e-4, + rtol=1e-4, + ) + + model.fit( + x=x_context, + y=target, + num_estimators=5, + estimator_batch_size=estimator_batch_size, + generator=torch.Generator().manual_seed(0), + ) + actual = model.predict(x_query) + assert actual.columns == expected.columns + torch.testing.assert_close( + actual=actual.numerical, + expected=expected.numerical, + atol=1e-4, + rtol=1e-4, + ) + + def test_fit_predict( cls_model: KumoTabular, reg_model: KumoTabular, diff --git a/test/models/tabfm/test_model.py b/test/models/tabfm/test_model.py index 2c08f74be..850a686ea 100644 --- a/test/models/tabfm/test_model.py +++ b/test/models/tabfm/test_model.py @@ -14,9 +14,11 @@ @withCUDA @pytest.mark.parametrize("dtype", [torch.int64, torch.float32]) +@pytest.mark.parametrize("estimator_batch_size", [1, 2, None]) def test_forward( device: torch.device, dtype: torch.dtype, + estimator_batch_size: int | None, monkeypatch: pytest.MonkeyPatch, ) -> None: monkeypatch.setattr( @@ -36,6 +38,10 @@ def test_forward( pretrained=False, device=device, ) + # Make predictions sensitive to cached attention and feature permutations. + for parameter in model.parameters(): + if not parameter.any(): + torch.nn.init.normal_(parameter, std=0.02) if device.type == "cpu": assert repr(model) == "TabFM()" else: @@ -55,7 +61,13 @@ def test_forward( y_context = torch.tensor([0, 1, 0, 1, 0], device=device).unsqueeze(-1) generator = torch.Generator(device=device).manual_seed(1) - out = model(x_context, y_context, x_query, generator=generator) + out = model( + x_context=x_context, + y_context=y_context, + x_query=x_query, + num_estimators=9, + generator=generator, + ) assert out.dtype == x_context.dtype assert out.device == device assert torch.is_inference(out) @@ -65,7 +77,24 @@ def test_forward( assert out.size() == (3, 2) generator = torch.Generator(device=device).manual_seed(1) - model.fit(x_context, y_context, generator=generator) + batched = model( + x_context=x_context, + y_context=y_context, + x_query=x_query, + num_estimators=9, + estimator_batch_size=estimator_batch_size, + generator=generator, + ) + assert batched.allclose(out, atol=1e-4, rtol=1e-4) + + generator = torch.Generator(device=device).manual_seed(1) + model.fit( + x=x_context, + y=y_context, + num_estimators=9, + estimator_batch_size=estimator_batch_size, + generator=generator, + ) assert model._cache is not None assert model._cache.size() > 0 assert model.predict(x_query).allclose(out, atol=1e-4, rtol=1e-4) diff --git a/test/models/tabiclv2/test_model.py b/test/models/tabiclv2/test_model.py index 8dacb0a61..43f98ddfe 100644 --- a/test/models/tabiclv2/test_model.py +++ b/test/models/tabiclv2/test_model.py @@ -109,14 +109,23 @@ def test_autocast_preserves_gradients_in_forward( assert any(p.grad is not None for p in model.parameters()) -def test_train_mode_preserves_gradients_with_ensembling() -> None: +@pytest.mark.parametrize("estimator_batch_size", [1, 2, None]) +def test_train_mode_preserves_gradients_with_ensembling( + estimator_batch_size: int | None, +) -> None: model = TabICLv2(pretrained=False) model.train() x_context = torch.eye(2) y_context = torch.tensor([[0.0], [1.0]]) x_query = torch.ones(1, 2) - out = model(x_context, y_context, x_query, num_estimators=2) + out = model( + x_context=x_context, + y_context=y_context, + x_query=x_query, + num_estimators=3, + estimator_batch_size=estimator_batch_size, + ) assert not torch.is_inference(out.numerical) out.numerical.sum().backward() @@ -124,7 +133,10 @@ def test_train_mode_preserves_gradients_with_ensembling() -> None: @pytest.mark.parametrize("batch_shape", [(), (2,)]) -def test_num_estimators(batch_shape: tuple[int, ...]) -> None: +@pytest.mark.parametrize("estimator_batch_size", [1, 2, None]) +def test_num_estimators( + batch_shape: tuple[int, ...], estimator_batch_size: int | None +) -> None: model = TabICLv2(pretrained=False) R_context, R_query, C = 5, 3, 6 @@ -133,7 +145,13 @@ def test_num_estimators(batch_shape: tuple[int, ...]) -> None: x_query = torch.randn(*batch_shape, R_query, C) y_context = torch.randn(*batch_shape, R_context, 1) - out = model(x_context, y_context, x_query, num_estimators=2) + out = model( + x_context=x_context, + y_context=y_context, + x_query=x_query, + num_estimators=3, + estimator_batch_size=estimator_batch_size, + ) assert out.size() == (*batch_shape, R_query, 999) model.fit(x_context, y_context, num_estimators=3) @@ -237,8 +255,10 @@ def forward( @withCUDA +@pytest.mark.parametrize("estimator_batch_size", [1, 2, None]) def test_tabiclv2_many_classes_forward_and_cache( device: torch.device, + estimator_batch_size: int | None, ) -> None: torch.manual_seed(1) model = TabICLv2(pretrained=False, device=device) @@ -252,7 +272,7 @@ def test_tabiclv2_many_classes_forward_and_cache( ).unsqueeze(-1) torch.manual_seed(1) - out = model(x_context, y_context, x_query) + out = model(x_context, y_context, x_query, num_estimators=3) assert out.size() == (test_size, num_classes) probabilities = out.numerical @@ -262,7 +282,12 @@ def test_tabiclv2_many_classes_forward_and_cache( ) torch.manual_seed(1) - model.fit(x_context, y_context) + model.fit( + x=x_context, + y=y_context, + num_estimators=3, + estimator_batch_size=estimator_batch_size, + ) assert model.predict(x_query).allclose(out) diff --git a/test/models/test_base.py b/test/models/test_base.py index 53d3bd581..b5ad6205e 100644 --- a/test/models/test_base.py +++ b/test/models/test_base.py @@ -8,7 +8,14 @@ import torch import sdm.processing as sp -from sdm import ColumnarTensor, RelatedTables, Stype, TableTensor +from sdm import ( + CategoricalTensor, + ColumnarTensor, + EnsembleTable, + RelatedTables, + Stype, + TableTensor, +) from sdm.cache import Cache from sdm.models import ICLModel from sdm.models.callback import Callback @@ -62,6 +69,52 @@ def default_recipe(cls) -> sp.Recipe: return sp.Recipe() +class _ClassFrequencyModel(ICLModel): + """Predict the context class frequencies, one column per class.""" + + supported_feature_stypes = frozenset({Stype.numerical}) + supported_target_stypes = frozenset({Stype.categorical}) + supports_multi_target = False + supports_related_tables = False + + def __init__(self) -> None: + super().__init__(task=None) + self.eval() + + def _forward( + self, + x_context: TableTensor | None, + y_context: TableTensor | None, + x_query: TableTensor | None, + related_context_tables: RelatedTables | None, + related_query_tables: RelatedTables | None, + cache: Cache | None, + generator: torch.Generator | None, + **kwargs: Any, + ) -> TableTensor: + if cache is None or cache.is_recording: + assert y_context is not None + classes = y_context.categorical.categories[0] + code = y_context.categorical.code.squeeze(-1).long() # [..., R] + counts = torch.nn.functional.one_hot(code, len(classes)).float() + frequency = counts.mean(dim=-2, keepdim=True) # [..., 1, K] + if cache is not None: + cache["frequency"] = frequency + else: + classes = cast(torch.Tensor, cache["classes"]) + frequency = cast(torch.Tensor, cache["frequency"]) + if x_query is None: + return TableTensor(numerical=frequency) + return TableTensor( + columns={Stype.numerical: [str(c) for c in classes.tolist()]}, + numerical=frequency.expand(*x_query.size()[:-1], -1), + ) + + @classmethod + def default_recipe(cls) -> sp.Recipe: + return sp.Recipe() + + class _UnsupportedRecordingModel(_RecordingModel): supported_feature_stypes = frozenset({Stype.numerical}) supported_target_stypes = frozenset({Stype.numerical, Stype.categorical}) @@ -304,18 +357,34 @@ def test_callback() -> None: ] -def test_train_mode_enables_grad() -> None: +@pytest.mark.parametrize("num_estimators", [1, 3]) +@pytest.mark.parametrize("estimator_batch_size", [1, 2, None]) +def test_train_mode_enables_grad( + num_estimators: int, estimator_batch_size: int | None +) -> None: model = _RecordingModel() x_context = torch.tensor([[0.0], [2.0]]) y_context = torch.tensor([[0.0], [1.0]]) x_query = torch.tensor([[3.0]]) model.eval() - out = model(x_context, y_context, x_query) + out = model( + x_context=x_context, + y_context=y_context, + x_query=x_query, + num_estimators=num_estimators, + estimator_batch_size=estimator_batch_size, + ) assert torch.is_inference(out) model.train() - out = model(x_context, y_context, x_query) + out = model( + x_context=x_context, + y_context=y_context, + x_query=x_query, + num_estimators=num_estimators, + estimator_batch_size=estimator_batch_size, + ) assert not torch.is_inference(out) @@ -533,3 +602,281 @@ def test_ensemble_output_reduce() -> None: ) assert out.size() == (2, 3) + + +@pytest.mark.parametrize("estimator_batch_size", [1, 2, None]) +def test_estimator_callbacks(estimator_batch_size: int | None) -> None: + model = _RecordingModel() + x = torch.arange(30.0).view(5, 3, 2) + y = torch.zeros(5, 3, 1) + events: list[str] = [] + callbacks = (MyCallback("affine", 2.0, 3.0, events),) + hooks = ( + "context_preprocessing_end", + "query_preprocessing_end", + "model_forward_end", + ) + + out = model( + x_context=x, + y_context=y, + x_query=x, + estimator_batch_size=estimator_batch_size, + callbacks=callbacks, + ) + torch.testing.assert_close(out.numerical, 2.0 * x + 3.0) + for hook in hooks: + assert events.count(f"affine_{hook}") == 5 + + events.clear() + model.fit( + x=x, + y=y, + estimator_batch_size=estimator_batch_size, + callbacks=callbacks, + ) + out = model.predict(x, callbacks=callbacks) + torch.testing.assert_close(out.numerical, 2.0 * x + 3.0) + for hook in hooks: + assert events.count(f"affine_{hook}") == 5 + + +@pytest.mark.parametrize( + ("estimator_batch_size", "num_calls", "query_size"), + [(1, 4, (2, 2)), (2, 2, (2, 2, 2)), (None, 1, (4, 2, 2))], +) +def test_estimator_batching_groups_consecutive_members( + estimator_batch_size: int | None, + num_calls: int, + query_size: tuple[int, ...], +) -> None: + model = _RecordingModel() + x = torch.randn(4, 3, 2) + y = torch.zeros(4, 3, 1) + x_query = torch.randn(4, 2, 2) + + out = model(x, y, x_query, estimator_batch_size=estimator_batch_size) + + torch.testing.assert_close(out.numerical, x_query) + assert len(model.calls) == num_calls + assert model.calls[0].x_query is not None + assert model.calls[0].x_query.size() == query_size + + model.calls.clear() + model.fit(x, y, estimator_batch_size=estimator_batch_size) + out = model.predict(x_query) + + torch.testing.assert_close(out.numerical, x_query) + assert len(model.calls) == 2 * num_calls + assert model.calls[-1].x_query is not None + assert model.calls[-1].x_query.size() == query_size + + +@pytest.mark.parametrize("estimator_batch_size", [2, None]) +def test_estimator_batching_relabels_shuffled_classes( + estimator_batch_size: int | None, +) -> None: + model = _ClassFrequencyModel() + x = torch.randn(4, 2) + y = TableTensor( + categorical=CategoricalTensor( + code=torch.tensor([[0], [1], [0], [0]]), + categories=(torch.tensor([10, 20]),), + ), + ) + # Shift the class order of every other estimator. + recipe = sp.Recipe( + target=sp.StypeDispatch( + categorical=sp.ShuffleCategories(method="shift") + ), + ) + + def _check(out: TableTensor) -> None: + columns = out.columns[Stype.numerical] + assert sorted(columns) == ["10", "20"] + frequency = torch.tensor( + [0.75 if column == "10" else 0.25 for column in columns] + ) + torch.testing.assert_close(out.numerical, frequency.expand(4, 4, -1)) + + _check( + model( + x_context=x, + y_context=y, + x_query=x, + recipe=recipe, + num_estimators=4, + estimator_batch_size=estimator_batch_size, + generator=torch.Generator().manual_seed(0), + ) + ) + + model.fit( + x=x, + y=y, + recipe=recipe, + num_estimators=4, + estimator_batch_size=estimator_batch_size, + generator=torch.Generator().manual_seed(0), + ) + _check(model.predict(x)) + + +def test_estimator_batching_rejects_related_tables() -> None: + model = _RecordingModel() + x_context = _table([0.0, 2.0], [1, 2], value_column="feature") + x_query = _table([3.0], [3], value_column="feature") + y_context = TableTensor.from_tensor(torch.tensor([[0.0], [1.0]])) + related_context = _related_tables(query=False) + related_query = _related_tables(query=True) + + with pytest.raises(ValueError, match="Related tables require"): + model( + x_context=x_context, + y_context=y_context, + x_query=x_query, + related_context_tables=related_context, + related_query_tables=related_query, + num_estimators=2, + estimator_batch_size=None, + ) + with pytest.raises(ValueError, match="Related tables require"): + model.fit( + x=x_context, + y=y_context, + related_tables=related_context, + num_estimators=2, + estimator_batch_size=None, + ) + + +def test_estimator_batching_incompatible_shapes() -> None: + model = _RecordingModel() + x = EnsembleTable.from_tables( + tables=[ + TableTensor(numerical=torch.ones(3, 2)), + TableTensor(numerical=torch.ones(4, 2)), + ], + member_table_ids=(0, 1), + ) + y = EnsembleTable.from_tables( + tables=[ + TableTensor(numerical=torch.ones(3, 1)), + TableTensor(numerical=torch.ones(4, 1)), + ], + member_table_ids=(0, 1), + ) + with pytest.raises(RuntimeError, match="stack expects"): + model.fit(x, y, estimator_batch_size=None) + query = TableTensor(numerical=torch.randn(2, 2, 2)) + with pytest.raises(RuntimeError, match="stack expects"): + model(x, y, query, estimator_batch_size=None) + model.fit(x, y) + torch.testing.assert_close(model.predict(query).numerical, query.numerical) + + +def test_estimator_batching_incompatible_categories() -> None: + model = _RecordingModel() + y = EnsembleTable.from_tables( + tables=[ + TableTensor( + categorical=CategoricalTensor( + code=torch.zeros(3, 1, dtype=torch.long), + categories=(torch.arange(count),), + ), + ) + for count in (2, 3) + ], + member_table_ids=(0, 1), + ) + with pytest.raises(ValueError, match="Estimators in one batch"): + model.fit( + x=torch.ones(3, 2), + y=y, + num_estimators=2, + estimator_batch_size=None, + ) + with pytest.raises(ValueError, match="Estimators in one batch"): + model( + x_context=torch.ones(3, 2), + y_context=y, + x_query=torch.ones(1, 2), + num_estimators=2, + estimator_batch_size=None, + ) + + +@pytest.mark.parametrize("columns", [("a", "c"), ("a",)]) +def test_estimator_batching_incompatible_columns( + columns: tuple[str, ...], +) -> None: + model = _RecordingModel() + tables = [ + TableTensor( + columns={Stype.numerical: ("a", "b")}, + numerical=torch.ones(3, 2), + ), + TableTensor( + columns={Stype.numerical: columns}, + numerical=torch.ones(3, len(columns)), + ), + ] + x = EnsembleTable.from_tables(tables=tables, member_table_ids=(0, 1)) + y = torch.zeros(2, 3, 1) + x_query = EnsembleTable.from_tables( + tables=[table[:1] for table in tables], + member_table_ids=(0, 1), + ) + with pytest.raises(ValueError, match="Estimators in one batch"): + model(x, y, x_query, estimator_batch_size=None) + with pytest.raises(ValueError, match="Estimators in one batch"): + model.fit(x, y, estimator_batch_size=None) + + +def test_estimator_batching_incompatible_class_values() -> None: + model = _ClassFrequencyModel() + x = torch.randn(3, 2) + y = EnsembleTable.from_tables( + tables=[ + TableTensor( + categorical=CategoricalTensor( + code=torch.tensor([[0], [1], [0]]), + categories=(torch.tensor(categories),), + ), + ) + for categories in ([10, 20], [10, 30]) + ], + member_table_ids=(0, 1), + ) + with pytest.raises(ValueError, match="Estimators in one batch"): + model(x, y, x, num_estimators=2, estimator_batch_size=None) + with pytest.raises(ValueError, match="Estimators in one batch"): + model.fit(x, y, num_estimators=2, estimator_batch_size=None) + + +@pytest.mark.parametrize("estimator_batch_size", [1, 2, None]) +def test_forward_batching_validates_each_query_schema( + estimator_batch_size: int | None, +) -> None: + model = _RecordingModel() + context = TableTensor( + columns={Stype.numerical: ("a", "b")}, + numerical=torch.ones(3, 2), + ) + x = EnsembleTable.from_tables( + tables=[ + context, + TableTensor( + columns={Stype.numerical: ("b", "a")}, + numerical=torch.ones(3, 2), + ), + ], + member_table_ids=(0, 1), + ) + y = torch.zeros(2, 3, 1) + query = EnsembleTable.from_tables( + tables=[context[:1]], + member_table_ids=(0, 0), + ) + with pytest.raises(ValueError, match="share the same schema"): + model(x, y, query, estimator_batch_size=estimator_batch_size) diff --git a/test/processing/test_execution.py b/test/processing/test_execution.py index fd753dd30..ef8b55b8a 100644 --- a/test/processing/test_execution.py +++ b/test/processing/test_execution.py @@ -5,7 +5,14 @@ import torch import sdm.processing as sp -from sdm import EnsembleTable, Recipe, RelatedTables, TableTensor +from sdm import ( + CategoricalTensor, + EnsembleTable, + Recipe, + RelatedTables, + Stype, + TableTensor, +) from sdm.processing.execution import RecipeExecution @@ -354,3 +361,22 @@ def test_sequence_rejects_queries_that_cannot_form_fitted_batch() -> None: task_links=[], ), ) + + +def test_member_context_exposes_input_stypes() -> None: + x = TableTensor( + columns={Stype.numerical: ("n",), Stype.categorical: ("c",)}, + numerical=torch.randn(3, 1), + categorical=CategoricalTensor.from_tensor(torch.zeros(3, 1).long()), + ) + y = torch.zeros(3, 1) + recipe = Recipe(features=sp.ToNumerical()) + + (context,) = RecipeExecution(recipe).fit_transform( + x=x, + y=y, + related_tables=None, + ) + + assert context.input_stypes == x.stypes + assert context.x.columns[Stype.numerical] == ("n", "c")