diff --git a/benchmark/tabular/model.py b/benchmark/tabular/model.py index 5a5caead5..607a0e852 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,6 +148,7 @@ def _fit( y=y_context, recipe=recipe, num_estimators=num_estimators, + estimator_batch_size=1, generator=generator, ) return @@ -202,7 +205,10 @@ def _predict_proba( self.autocast_dtype, enabled=x_query.is_cuda, ): - out = self.model.predict(x_query) + out = self.model.predict( + x_query, + estimator_batch_size=self._estimator_batch_size(x_query), + ) else: with ( torch.inference_mode(), @@ -215,24 +221,20 @@ def _predict_proba( generator = torch.Generator(self._device).set_state( self._rng_state ) - outputs = [] - for context, query in zip(self._contexts, queries): - 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, - generator=generator, - _schema=self._schema, - ) - outputs.append(out.to(query.x.dtype)) + with torch.amp.autocast( + self._device.type, + self.autocast_dtype, + enabled=x_query.is_cuda, + ): + outputs = self.model._forward_estimators( + contexts=self._contexts, + queries=queries, + estimator_batch_size=self._estimator_batch_size( + x_query + ), + generator=generator, + _schema=self._schema, + ) if self.problem_type == REGRESSION: outputs = list( @@ -251,6 +253,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/models/base.py b/sdm/models/base.py index 43293d194..9b2e119f1 100644 --- a/sdm/models/base.py +++ b/sdm/models/base.py @@ -14,13 +14,14 @@ Recipe, RelatedTables, Stype, + StypeLike, TableTensor, Task, TaskLike, ) from sdm._inference import inference_mode from sdm._warnings import warn_once -from sdm.cache import Cache +from sdm.cache import Cache, KVCacheEntry from sdm.models.callback import Callback from sdm.processing.execution import ( MemberContext, @@ -30,6 +31,15 @@ from sdm.relational.task import RelatedTablesSchema from sdm.tensor.table import TableSchema +_CACHE_METADATA_KEYS = { + "x_schema", + "x_schemas", + "y_schema", + "related_tables_schema", + "classes", + "output_columns", +} + class ICLModel(torch.nn.Module, abc.ABC): r"""Base model for in-context foundation models on structured data. @@ -102,6 +112,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 +135,11 @@ 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 per model + execution. ``None`` runs all estimators together; ``1`` runs + them sequentially. Batched estimators must have compatible + shapes, target columns, and feature categories. Larger batches + use more device memory. Callbacks require ``1``. callbacks: Callbacks applied in sequence to this model call. generator: Pseudorandom number generator used for sampling during pre-processing and model execution. @@ -134,6 +150,8 @@ def forward( stacked estimator outputs with shape ``[E, ..., R_query, *]``. """ callbacks = () if callbacks is None else callbacks + if callbacks and estimator_batch_size != 1: + raise ValueError("Callbacks require 'estimator_batch_size=1'") requires_grad = self.training requires_grad |= any(callback.requires_grad for callback in callbacks) @@ -172,48 +190,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_estimators( + 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 +221,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,12 +243,19 @@ 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 per model + execution. ``None`` runs all estimators together; ``1`` runs + them sequentially. Batched estimators must have compatible + shapes, target columns, and feature categories. Larger batches + use more device memory. Callbacks require ``1``. callbacks: Callbacks applied in sequence to this model call. generator: Pseudorandom number generator used for sampling during pre-processing and model execution. kwargs: Additional keyword arguments passed to the model. """ callbacks = () if callbacks is None else callbacks + if callbacks and estimator_batch_size != 1: + raise ValueError("Callbacks require 'estimator_batch_size=1'") self.clear() @@ -282,15 +274,24 @@ def fit( generator=generator, ) + if estimator_batch_size is None: + estimator_batch_size = len(contexts) + starts = 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 i, start in enumerate(starts): + members = contexts[start : start + estimator_batch_size] + with inference_mode("no_grad"): + context = _stack_contexts(members) for callback in callbacks: context = MemberContext( *callback.on_context_preprocessing_end(self, *context) ) + if len(members) == 1: + members = (context,) self._validate_context( x=context.x, y=context.y, @@ -298,6 +299,8 @@ def fit( ) estimator_cache = Cache( x_schema=context.x.schema, + x_schemas=tuple(member.x.schema for member in members), + output_columns=_output_columns(members), y_schema=context.y.schema, related_tables_schema=context.related_tables.schema if context.related_tables is not None @@ -321,7 +324,7 @@ def fit( **kwargs, ) - if x.is_cuda and len(contexts) > 1: + if x.is_cuda and len(starts) > 1: try: # Copy to pinned CPU memory: estimator_cache = estimator_cache._apply_tensor( lambda tensor: torch.ops.aten._to_copy.default( @@ -343,6 +346,7 @@ def predict( x: Tensor | TableTensor | EnsembleTable, # [..., R, D] related_tables: RelatedTables | None = None, *, + estimator_batch_size: int | None = 1, callbacks: Sequence[Callback] | None = None, ) -> TableTensor: # Recipe-defined output shape. r"""Predict unseen query examples. @@ -355,6 +359,13 @@ def predict( x: The feature tensor of query examples with shape ``[..., R, D]`` with ``R`` rows and ``D`` columns. related_tables: Related context for query examples. + estimator_batch_size: Maximum number of estimators per model + execution. ``None`` runs all estimators together; ``1`` runs + them sequentially. This can differ from the batch size used + for :meth:`fit` when the model uses compatible tensor caches. + Sequentially fitted caches that cannot be batched stay + sequential. + Callbacks require ``1``. callbacks: Callbacks applied in sequence to this model call. Returns: @@ -369,6 +380,8 @@ def predict( ) callbacks = () if callbacks is None else callbacks + if callbacks and estimator_batch_size != 1: + raise ValueError("Callbacks require 'estimator_batch_size=1'") requires_grad = any(callback.requires_grad for callback in callbacks) if self._cache is None: @@ -393,10 +406,17 @@ def predict( RecipeExecution, self._cache["recipe_execution"], ) - caches = [ - cast(Cache, self._cache[i]) - for i in range(recipe_execution.num_members) - ] + if estimator_batch_size is None: + estimator_batch_size = recipe_execution.num_members + if ( + estimator_batch_size > 1 + and self._cache["estimator_batch_size"] == 1 + and not _can_batch_cache(self._cache) + ): + estimator_batch_size = 1 + starts = range(0, recipe_execution.num_members, estimator_batch_size) + with inference_mode("no_grad" if requires_grad else "inference"): + caches = _prediction_caches(self._cache, estimator_batch_size) next_cache = caches[0] compute_stream: torch.cuda.Stream | None = None @@ -424,9 +444,14 @@ def predict( compute_stream.wait_stream(transfer_stream) outs: list[TableTensor] = [] - for i, query in enumerate(queries): + for i, start in enumerate(starts): cache, next_cache = next_cache, None assert cache is not None + members = queries[start : start + estimator_batch_size] + with inference_mode( + "no_grad" if requires_grad else "inference" + ): + query = _stack_queries(members) for callback in callbacks: query = MemberQuery( @@ -441,6 +466,13 @@ def predict( ), related_query_tables=query.related_tables, ) + if len(members) > 1 and cache["x_schemas"] != tuple( + member.x.schema for member in members + ): + raise ValueError( + "Expected context and query features to share " + "the same schema" + ) if i + 1 < len(caches): next_cache = caches[i + 1] @@ -461,17 +493,29 @@ def predict( **cast(dict[str, Any], self._cache["kwargs"]), ) + outputs = _unstack_output( + out=out, + num_members=len(members), + columns=cast( + tuple[tuple[str, ...], ...] | None, + cache["output_columns"], + ), + ) for callback in callbacks: - out = callback.on_model_forward_end(self, out) + outputs = [ + callback.on_model_forward_end(self, out) + for out in outputs + ] + outs.extend( + cast(TableTensor, out.to(query.x.dtype)) + for out in outputs + ) 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 @@ -530,6 +574,7 @@ def _forward( generator: torch.Generator | None, **kwargs: Any, ) -> TableTensor: # [..., R_query, *] + # Regroupable cache tensors retain all leading input batch dimensions. pass @classmethod @@ -539,6 +584,102 @@ def default_recipe(cls) -> Recipe: # Helpers ################################################################# + def _forward_estimators( + self, + contexts: Sequence[MemberContext], + queries: Sequence[MemberQuery], + *, + estimator_batch_size: int | None = 1, + callbacks: Sequence[Callback] = (), + generator: torch.Generator | None = None, + **kwargs: Any, + ) -> list[TableTensor]: + if estimator_batch_size is None: + estimator_batch_size = len(contexts) + requires_grad = self.training or any( + callback.requires_grad for callback in callbacks + ) + outs: list[TableTensor] = [] + for start in range(0, len(contexts), estimator_batch_size): + members = contexts[start : start + estimator_batch_size] + query_members = queries[start : start + estimator_batch_size] + with inference_mode("no_grad" if requires_grad else "inference"): + context = _stack_contexts(members) + query = _stack_queries(query_members) + if context.x.device != query.x.device: + # Transfer only the current batch of offloaded contexts. + context = MemberContext( + x=cast(TableTensor, context.x.to(query.x.device)), + y=cast(TableTensor, context.y.to(query.x.device)), + related_tables=context.related_tables.to( + query.x.device + ) + if context.related_tables is not None + else None, + ) + + 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, + ) + if len(members) > 1 and any( + context.x.schema != query.x.schema + for context, query in zip(members, query_members) + ): + raise ValueError( + "Expected context and query features to share " + "the same schema" + ) + + 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, + _x_schemas=tuple(member.x.schema for member in members) + if len(members) > 1 + else (context.x.schema,), + **kwargs, + ) + outputs = _unstack_output( + out=out, + num_members=len(members), + columns=_output_columns(members) + if len(members) > 1 + else None, + ) + for callback in callbacks: + outputs = [ + callback.on_model_forward_end(self, out) + for out in outputs + ] + outs.extend( + cast(TableTensor, out.to(query.x.dtype)) for out in outputs + ) + return outs + def _validate_context( self, x: TableTensor, @@ -632,3 +773,259 @@ def _validate_query( "Expected related context and query tables to share the " "same schema" ) + + +def _stack_tables( + tables: Sequence[TableTensor], + *, + target: bool = False, +) -> TableTensor: + if len(tables) == 1: + return tables[0] + + ref = tables[0] + for table in tables[1:]: + # Numerical feature positions can differ after column shuffling. + # Other metadata must agree because the batch shares one schema. + if table.active_stypes != ref.active_stypes or any( + columns != ref.columns[stype] + for stype, columns in table.columns.items() + if target or stype != Stype.numerical + ): + raise ValueError( + "Estimator batches require compatible column layouts; " + "use 'estimator_batch_size=1' for incompatible estimators" + ) + for categories, ref_categories in zip( + table.categorical.categories, + ref.categorical.categories, + strict=True, + ): + compatible = ( + len(categories) == len(ref_categories) + if target + else categories is ref_categories + or categories.equal(ref_categories) + ) + if not compatible: + raise ValueError( + "Estimator batches require matching class counts and " + "compatible feature categories; use " + "'estimator_batch_size=1' for incompatible estimators" + ) + + # TableTensor.stack aligns names, which would undo feature permutations. + # Stack blocks by position; target class labels are restored on output. + return TableTensor( + columns=cast(dict[StypeLike, tuple[str, ...]], ref.columns), + size=(len(tables), *ref.size()[:-1]), + device=ref.device, + **{ + stype: torch.stack( + [table.blocks[stype] for table in tables], dim=0 + ) + for stype, _ in ref.items() + }, + ) + + +def _stack_related( + tables: Sequence[RelatedTables[TableTensor] | None], +) -> RelatedTables[TableTensor] | None: + ref = tables[0] + if len(tables) == 1 or ref is None: + return ref + if any(table is None or table.schema != ref.schema for table in tables): + raise ValueError( + "Estimator batches require compatible related table schemas; " + "use 'estimator_batch_size=1' for incompatible estimators" + ) + tables = cast(Sequence[RelatedTables[TableTensor]], tables) + return ref.replace_tables( + { + name: _stack_tables([table.tables[name] for table in tables]) + for name in ref.tables + } + ) + + +def _stack_contexts(contexts: Sequence[MemberContext]) -> MemberContext: + return MemberContext( + x=_stack_tables([context.x for context in contexts]), + y=_stack_tables([context.y for context in contexts], target=True), + related_tables=_stack_related( + [context.related_tables for context in contexts] + ), + ) + + +def _stack_queries(queries: Sequence[MemberQuery]) -> MemberQuery: + return MemberQuery( + x=_stack_tables([query.x for query in queries]), + related_tables=_stack_related( + [query.related_tables for query in queries] + ), + ) + + +def _output_columns( + contexts: Sequence[MemberContext], +) -> tuple[tuple[str, ...], ...] | None: + if contexts[0].y.categorical.size(-1) == 0: + return None + return tuple( + tuple( + str(value) + for value in context.y.categorical.categories[0].tolist() + ) + for context in contexts + ) + + +def _can_batch_cache(cache: Cache) -> bool: + return all( + isinstance(value, Tensor | KVCacheEntry) + for estimator_cache in cache.values() + if isinstance(estimator_cache, Cache) + for key, value in estimator_cache.items() + if key not in _CACHE_METADATA_KEYS + ) + + +def _prediction_caches(cache: Cache, estimator_batch_size: int) -> list[Cache]: + num_members = cast(RecipeExecution, cache["recipe_execution"]).num_members + fitted_estimator_batch_size = cast(int, cache["estimator_batch_size"]) + caches = [ + cast(Cache, cache[i]) + for i in range(len(range(0, num_members, fitted_estimator_batch_size))) + ] + if min(estimator_batch_size, num_members) == min( + fitted_estimator_batch_size, num_members + ): + return caches + + members: list[Cache] = [] + for fitted in caches: + schemas = cast(tuple[TableSchema, ...], fitted["x_schemas"]) + columns = cast( + tuple[tuple[str, ...], ...] | None, fitted["output_columns"] + ) + for index, schema in enumerate(schemas): + member = Cache(fitted) + member["x_schema"] = schema + member["x_schemas"] = (schema,) + member["output_columns"] = ( + None if columns is None else (columns[index],) + ) + if len(schemas) > 1: + for key, value in fitted.items(): + if key not in _CACHE_METADATA_KEYS: + if isinstance(value, Tensor): + member[key] = value[index] + elif isinstance(value, KVCacheEntry): + member[key] = KVCacheEntry( + key=value.key[index], value=value.value[index] + ) + else: + raise ValueError( + "Changing estimator batch size requires " + "tensor caches" + ) + members.append(member) + + batches: list[Cache] = [] + for start in range(0, num_members, estimator_batch_size): + group = members[start : start + estimator_batch_size] + first = group[0] + if len(group) == 1: + batches.append(first.freeze()) + continue + classes = cast(Tensor | None, first["classes"]) + for member in group[1:]: + other_classes = cast(Tensor | None, member["classes"]) + if ( + member.keys() != first.keys() + or member["y_schema"] != first["y_schema"] + or member["related_tables_schema"] + != first["related_tables_schema"] + or (classes is None) != (other_classes is None) + or ( + classes is not None + and len(classes) != len(cast(Tensor, other_classes)) + ) + ): + raise ValueError( + "Estimator batches require compatible caches; " + "use 'estimator_batch_size=1'" + ) + batched = Cache(first) + batched["x_schemas"] = tuple(member["x_schema"] for member in group) + batched["output_columns"] = ( + None + if classes is None + else tuple( + cast(tuple[tuple[str, ...], ...], member["output_columns"])[0] + for member in group + ) + ) + for key, value in first.items(): + if key not in _CACHE_METADATA_KEYS: + values = [member[key] for member in group] + if isinstance(value, Tensor): + batched[key] = torch.stack(cast(list[Tensor], values)) + elif isinstance(value, KVCacheEntry): + entries = cast(list[KVCacheEntry], values) + batched[key] = KVCacheEntry( + key=torch.stack([entry.key for entry in entries]), + value=torch.stack([entry.value for entry in entries]), + ) + else: + raise ValueError( + "Changing estimator batch size requires tensor caches" + ) + batches.append(batched.freeze()) + return batches + + +def _unstack_output( + out: TableTensor, + num_members: int, + columns: Sequence[tuple[str, ...]] | None, +) -> list[TableTensor]: + outputs = ( + [out] + if num_members == 1 + else list(cast(tuple[TableTensor, ...], out.unbind(0))) + ) + if columns is None: + return outputs + return [ + TableTensor( + columns={Stype.numerical: names}, + numerical=output.numerical, + ) + for output, names in zip(outputs, columns, strict=True) + ] + + +def _categorical_mask( + x: TableTensor, + schema: TableSchema, + schemas: Sequence[TableSchema], +) -> Tensor: + categorical_columns = set(schema.columns[Stype.categorical]) + mask = torch.tensor( + [ + [ + column in categorical_columns + for column in s.columns[Stype.numerical] + ] + for s in schemas + ], + device=x.device, + dtype=torch.bool, + ) + if len(schemas) == 1: + return mask[0] + # [E, C] -> [E, 1, ..., C], preserving existing input batch dimensions. + return mask.view(len(schemas), *((1,) * (x.dim() - 3)), -1) diff --git a/sdm/models/kumo/tabular/model.py b/sdm/models/kumo/tabular/model.py index d60bf15a4..7c315ce72 100644 --- a/sdm/models/kumo/tabular/model.py +++ b/sdm/models/kumo/tabular/model.py @@ -15,6 +15,7 @@ from sdm.cache import Cache from sdm.models import ECOC, ICLModel from sdm.models._huggingface import download_checkpoint +from sdm.models.base import _categorical_mask 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 @@ -220,20 +221,18 @@ def _forward( 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, + categorical_mask = _categorical_mask( + x=x_context, + schema=schema, + schemas=kwargs.get("_x_schemas", (x_context.schema,)) + if cache is None + else cast(tuple[TableSchema, ...], cache["x_schemas"]), ) + categorical_mask = categorical_mask.expand(*x.size()[:-2], -1) 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: out = self.models[Task.regression]( diff --git a/sdm/models/tabfm/model.py b/sdm/models/tabfm/model.py index 8b53228bf..5477d4380 100644 --- a/sdm/models/tabfm/model.py +++ b/sdm/models/tabfm/model.py @@ -29,7 +29,7 @@ from sdm import Recipe, RelatedTables, Stype, TableTensor, Task, TaskLike from sdm.cache import Cache from sdm.models._huggingface import download_checkpoint -from sdm.models.base import ICLModel +from sdm.models.base import ICLModel, _categorical_mask from sdm.models.tabfm.ckpt import remap_ckpt from sdm.models.tabfm.icl import ICLBlock from sdm.models.tabfm.recipe import default_recipe @@ -207,20 +207,18 @@ def _forward( 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, + categorical_mask = _categorical_mask( + x=x_context, + schema=schema, + schemas=kwargs.get("_x_schemas", (x_context.schema,)) + if cache is None + else cast(tuple[TableSchema, ...], cache["x_schemas"]), ) + categorical_mask = categorical_mask.expand(*x.size()[:-2], -1) 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 out = self.models[task](x, y, categorical_mask, cache=cache) diff --git a/test/explain/test_gradient.py b/test/explain/test_gradient.py index 914c4394d..dca38532c 100644 --- a/test/explain/test_gradient.py +++ b/test/explain/test_gradient.py @@ -6,7 +6,7 @@ import pytest import torch -from sdm import Recipe, RelatedTables, Stype, TableTensor +from sdm import CategoricalTensor, Recipe, RelatedTables, Stype, TableTensor from sdm.cache import Cache from sdm.explain import GradientExplainer from sdm.models import ICLModel @@ -14,7 +14,7 @@ class _LinearModel(ICLModel): supported_feature_stypes = frozenset({Stype.numerical}) - supported_target_stypes = frozenset({Stype.numerical}) + supported_target_stypes = frozenset({Stype.numerical, Stype.categorical}) supports_multi_target = False supports_related_tables = True @@ -48,11 +48,24 @@ def default_recipe(cls) -> Recipe: return Recipe() -@pytest.mark.parametrize("fitted", [False, True]) -def test_returns_query_input_gradients(fitted: bool) -> None: +@pytest.mark.parametrize( + ("fitted", "num_estimators"), + [(False, 1), (False, 3), (True, 1), (True, 3)], +) +@pytest.mark.parametrize("classification", [False, True]) +def test_returns_query_input_gradients( + fitted: bool, num_estimators: int, classification: bool +) -> None: model = _LinearModel() x_context = torch.zeros(1, 2) - y_context = torch.zeros(1, 1) + y_context = TableTensor(numerical=torch.zeros(1, 1)) + if classification: + y_context = TableTensor( + categorical=CategoricalTensor( + code=torch.zeros(1, 1, dtype=torch.long), + categories=(torch.arange(2),), + ), + ) x_query = torch.ones(1, 2) related_tables = RelatedTables( tables={ @@ -76,7 +89,12 @@ def test_returns_query_input_gradients(fitted: bool) -> None: ) if fitted: - model.fit(x_context, y_context, related_tables) + model.fit( + x=x_context, + y=y_context, + related_tables=related_tables, + num_estimators=num_estimators, + ) result = explainer.explain(model, x_query, related_tables) else: result = explainer.explain( @@ -86,6 +104,7 @@ def test_returns_query_input_gradients(fitted: bool) -> None: x_context=x_context, y_context=y_context, related_context_tables=related_tables, + num_estimators=num_estimators, ) torch.testing.assert_close( diff --git a/test/models/kumo/tabular/test_model.py b/test/models/kumo/tabular/test_model.py index 66bccbbff..6e1301f2f 100644 --- a/test/models/kumo/tabular/test_model.py +++ b/test/models/kumo/tabular/test_model.py @@ -130,6 +130,44 @@ 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_forward_estimator_batching( + task: Literal["classification", "regression"], + size: Literal["small", "medium"], + 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, + ) + + 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..d4930ae98 100644 --- a/test/models/tabfm/test_model.py +++ b/test/models/tabfm/test_model.py @@ -14,9 +14,15 @@ @withCUDA @pytest.mark.parametrize("dtype", [torch.int64, torch.float32]) +@pytest.mark.parametrize( + ("fit_estimator_batch_size", "predict_estimator_batch_size"), + [(1, 1), (1, 8), (8, 1), (2, None), (None, 2)], +) def test_forward( device: torch.device, dtype: torch.dtype, + fit_estimator_batch_size: int | None, + predict_estimator_batch_size: int | None, monkeypatch: pytest.MonkeyPatch, ) -> None: monkeypatch.setattr( @@ -36,6 +42,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 +65,14 @@ 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, + estimator_batch_size=predict_estimator_batch_size, + generator=generator, + ) assert out.dtype == x_context.dtype assert out.device == device assert torch.is_inference(out) @@ -65,8 +82,17 @@ def test_forward( assert out.size() == (3, 2) generator = torch.Generator(device=device).manual_seed(1) - model.fit(x_context, y_context, generator=generator) + model.fit( + x=x_context, + y=y_context, + num_estimators=9, + estimator_batch_size=fit_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) + assert model.predict( + x=x_query, + estimator_batch_size=predict_estimator_batch_size, + ).allclose(out, atol=1e-4, rtol=1e-4) model.clear() diff --git a/test/models/tabiclv2/test_model.py b/test/models/tabiclv2/test_model.py index 8dacb0a61..0facb6c47 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,13 +145,17 @@ 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) assert model._cache is not None - assert 0 in model._cache - assert 1 in model._cache assert model._cache.size() > 0 assert model._cache.is_cpu @@ -237,8 +253,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 +270,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,8 +280,11 @@ def test_tabiclv2_many_classes_forward_and_cache( ) torch.manual_seed(1) - model.fit(x_context, y_context) - assert model.predict(x_query).allclose(out) + model.fit(x_context, y_context, num_estimators=3) + assert model.predict( + x=x_query, + estimator_batch_size=estimator_batch_size, + ).allclose(out) @onlyCUDA diff --git a/test/models/test_base.py b/test/models/test_base.py index 53d3bd581..0e55fe93b 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 @@ -304,18 +311,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 +556,131 @@ 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) + callbacks = (MyCallback("affine", 2.0, 3.0, []),) + if estimator_batch_size != 1: + with pytest.raises(ValueError, match="Callbacks require"): + model( + x_context=x, + y_context=y, + x_query=x, + estimator_batch_size=estimator_batch_size, + callbacks=callbacks, + ) + with pytest.raises(ValueError, match="Callbacks require"): + model.fit( + x=x, + y=y, + estimator_batch_size=estimator_batch_size, + callbacks=callbacks, + ) + model.fit(x, y, estimator_batch_size=estimator_batch_size) + if estimator_batch_size != 1: + with pytest.raises(ValueError, match="Callbacks require"): + model.predict( + x=x, + 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) + out = model( + x_context=x, + y_context=y, + x_query=x, + callbacks=callbacks, + ) + torch.testing.assert_close(out.numerical, 2.0 * x + 3.0) + + +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) + with pytest.raises(RuntimeError, match="stack expects"): + model(x, y, torch.ones(2, 2, 2), estimator_batch_size=None) + model.fit(x, y) + query = torch.randn(2, 2, 2) + torch.testing.assert_close( + model.predict(TableTensor(numerical=query)).numerical, + query, + ) + + +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="matching class counts"): + model.fit( + torch.ones(3, 2), y, num_estimators=2, estimator_batch_size=None + ) + with pytest.raises(ValueError, match="matching class counts"): + model( + x_context=torch.ones(3, 2), + y_context=y, + x_query=torch.ones(1, 2), + num_estimators=2, + estimator_batch_size=None, + ) + model.fit(torch.ones(3, 2), y, num_estimators=2) + with pytest.raises(ValueError, match="compatible caches"): + model.predict(torch.ones(1, 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)