From c4b4f8bb772fd84d335777e7127ff1f6b1dae85a Mon Sep 17 00:00:00 2001 From: Akihiro Nitta Date: Tue, 22 Sep 2026 21:00:57 +0000 Subject: [PATCH 01/11] update --- sdm/models/_batch.py | 163 +++++++++++++++++++++++ sdm/models/base.py | 115 +++++++++++++++-- sdm/models/kumo/tabular/model.py | 15 +-- sdm/models/tabfm/model.py | 15 +-- test/explain/test_gradient.py | 19 ++- test/models/tabiclv2/test_model.py | 2 - test/models/test_base.py | 82 +++++++++++- test/models/test_estimator_batching.py | 171 +++++++++++++++++++++++++ 8 files changed, 547 insertions(+), 35 deletions(-) create mode 100644 sdm/models/_batch.py create mode 100644 test/models/test_estimator_batching.py diff --git a/sdm/models/_batch.py b/sdm/models/_batch.py new file mode 100644 index 000000000..b6aba327b --- /dev/null +++ b/sdm/models/_batch.py @@ -0,0 +1,163 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from collections.abc import Sequence +from typing import cast + +import torch +from torch import Tensor + +from sdm import RelatedTables, Stype, StypeLike, TableTensor +from sdm.processing.execution import MemberContext, MemberQuery +from sdm.tensor.table import TableSchema + + +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 len(contexts) == 1 or 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 _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/base.py b/sdm/models/base.py index 43293d194..caf573f62 100644 --- a/sdm/models/base.py +++ b/sdm/models/base.py @@ -21,6 +21,12 @@ from sdm._inference import inference_mode from sdm._warnings import warn_once from sdm.cache import Cache +from sdm.models._batch import ( + _output_columns, + _stack_contexts, + _stack_queries, + _unstack_output, +) from sdm.models.callback import Callback from sdm.processing.execution import ( MemberContext, @@ -102,6 +108,7 @@ def forward( *, recipe: Recipe | None = None, num_estimators: int | None = None, + estimator_batch_size: int | None = None, callbacks: Sequence[Callback] | None = None, generator: torch.Generator | None = None, **kwargs: Any, @@ -124,6 +131,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: Callbacks applied in sequence to this model call. generator: Pseudorandom number generator used for sampling during pre-processing and model execution. @@ -172,8 +184,28 @@ def forward( related_tables=related_query_tables, ) + batch_size = ( + len(contexts) + if estimator_batch_size is None + else estimator_batch_size + ) outs: list[TableTensor] = [] - for context, query in zip(contexts, queries): + for start in range(0, len(contexts), batch_size): + context_members = contexts[start : start + batch_size] + query_members = queries[start : start + batch_size] + if len(context_members) > 1: + for context, query in zip(context_members, query_members): + 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("no_grad" if requires_grad else "inference"): + context = _stack_contexts(context_members) + query = _stack_queries(query_members) for callback in callbacks: context = MemberContext( *callback.on_context_preprocessing_end(self, *context) @@ -206,6 +238,11 @@ def forward( related_query_tables=query.related_tables, cache=None, generator=generator, + _x_schemas=( + tuple(member.x.schema for member in context_members) + if len(context_members) > 1 + else (context.x.schema,) + ), **kwargs, ) @@ -213,7 +250,14 @@ def forward( out = callback.on_model_forward_end(self, out) out = cast(TableTensor, out.to(query.x.dtype)) - outs.append(out) + with inference_mode("grad" if requires_grad else "inference"): + outs.extend( + _unstack_output( + out=out, + num_members=len(context_members), + columns=_output_columns(context_members), + ) + ) # Regression: invert target before stacking estimator outputs. if contexts[0].y.numerical.size(-1) > 0: @@ -237,6 +281,7 @@ def fit( *, recipe: Recipe | None = None, num_estimators: int | None = None, + estimator_batch_size: int | None = None, callbacks: Sequence[Callback] | None = None, generator: torch.Generator | None = None, **kwargs: Any, @@ -258,6 +303,11 @@ 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. Prediction + reuses the same batches. Larger batches use more device memory. callbacks: Callbacks applied in sequence to this model call. generator: Pseudorandom number generator used for sampling during pre-processing and model execution. @@ -282,11 +332,21 @@ def fit( generator=generator, ) + batch_size = ( + len(contexts) + if estimator_batch_size is None + else estimator_batch_size + ) + starts = range(0, len(contexts), batch_size) cache = Cache( recipe_execution=recipe_execution, kwargs=kwargs, + estimator_batch_size=batch_size, ) - for i, context in enumerate(contexts): + for i, start in enumerate(starts): + members = contexts[start : start + batch_size] + with inference_mode("no_grad"): + context = _stack_contexts(members) for callback in callbacks: context = MemberContext( *callback.on_context_preprocessing_end(self, *context) @@ -298,6 +358,12 @@ def fit( ) estimator_cache = Cache( x_schema=context.x.schema, + x_schemas=( + tuple(member.x.schema for member in members) + if len(members) > 1 + else (context.x.schema,) + ), + 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 +387,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( @@ -393,10 +459,9 @@ def predict( RecipeExecution, self._cache["recipe_execution"], ) - caches = [ - cast(Cache, self._cache[i]) - for i in range(recipe_execution.num_members) - ] + batch_size = cast(int, self._cache["estimator_batch_size"]) + starts = range(0, recipe_execution.num_members, batch_size) + caches = [cast(Cache, self._cache[i]) for i in range(len(starts))] next_cache = caches[0] compute_stream: torch.cuda.Stream | None = None @@ -424,9 +489,29 @@ 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 + batch_size] + if len(members) > 1: + for schema, member in zip( + cast(tuple[TableSchema, ...], cache["x_schemas"]), + members, + strict=True, + ): + self._validate_query( + x_context=schema, + x_query=member.x, + related_context_tables=cast( + RelatedTablesSchema | None, + cache["related_tables_schema"], + ), + related_query_tables=member.related_tables, + ) + with inference_mode( + "no_grad" if requires_grad else "inference" + ): + query = _stack_queries(members) for callback in callbacks: query = MemberQuery( @@ -470,7 +555,17 @@ def predict( tensor.record_stream(compute_stream) out = cast(TableTensor, out.to(query.x.dtype)) - outs.append(out) + with inference_mode("grad" if requires_grad else "inference"): + outs.extend( + _unstack_output( + out=out, + num_members=len(members), + columns=cast( + tuple[tuple[str, ...], ...] | None, + cache["output_columns"], + ), + ) + ) if x.is_cuda and next_cache is not None: assert compute_stream is not None diff --git a/sdm/models/kumo/tabular/model.py b/sdm/models/kumo/tabular/model.py index c5a151287..6bc897a38 100644 --- a/sdm/models/kumo/tabular/model.py +++ b/sdm/models/kumo/tabular/model.py @@ -14,6 +14,7 @@ from sdm import Recipe, RelatedTables, Stype, TableTensor, Task, TaskLike from sdm.cache import Cache from sdm.models import ECOC, ICLModel +from sdm.models._batch import _categorical_mask from sdm.models._huggingface import download_checkpoint from sdm.models.kumo.tabular.icl import ICLBlock from sdm.models.kumo.tabular.recipe import default_recipe @@ -218,14 +219,12 @@ 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"]), ) if cache is not None: cache["categorical_mask"] = categorical_mask diff --git a/sdm/models/tabfm/model.py b/sdm/models/tabfm/model.py index 8b53228bf..44b42b73e 100644 --- a/sdm/models/tabfm/model.py +++ b/sdm/models/tabfm/model.py @@ -28,6 +28,7 @@ from sdm import Recipe, RelatedTables, Stype, TableTensor, Task, TaskLike from sdm.cache import Cache +from sdm.models._batch import _categorical_mask from sdm.models._huggingface import download_checkpoint from sdm.models.base import ICLModel from sdm.models.tabfm.ckpt import remap_ckpt @@ -207,14 +208,12 @@ 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"]), ) if cache is not None: cache["categorical_mask"] = categorical_mask diff --git a/test/explain/test_gradient.py b/test/explain/test_gradient.py index 914c4394d..3ef44a915 100644 --- a/test/explain/test_gradient.py +++ b/test/explain/test_gradient.py @@ -48,8 +48,13 @@ 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)], +) +def test_returns_query_input_gradients( + fitted: bool, num_estimators: int +) -> None: model = _LinearModel() x_context = torch.zeros(1, 2) y_context = torch.zeros(1, 1) @@ -76,7 +81,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,8 +96,11 @@ 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, ) + if num_estimators > 1: + x_query = x_query.expand(num_estimators, *x_query.size()) torch.testing.assert_close( result.x.numerical, torch.full_like(x_query, 2.0) ) diff --git a/test/models/tabiclv2/test_model.py b/test/models/tabiclv2/test_model.py index 8dacb0a61..161e0a255 100644 --- a/test/models/tabiclv2/test_model.py +++ b/test/models/tabiclv2/test_model.py @@ -138,8 +138,6 @@ def test_num_estimators(batch_shape: tuple[int, ...]) -> None: 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 diff --git a/test/models/test_base.py b/test/models/test_base.py index 53d3bd581..bfa8157ae 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,19 @@ def test_callback() -> None: ] -def test_train_mode_enables_grad() -> None: +@pytest.mark.parametrize("num_estimators", [1, 3]) +def test_train_mode_enables_grad(num_estimators: int) -> 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, y_context, x_query, num_estimators=num_estimators) assert torch.is_inference(out) model.train() - out = model(x_context, y_context, x_query) + out = model(x_context, y_context, x_query, num_estimators=num_estimators) assert not torch.is_inference(out) @@ -361,6 +369,7 @@ def test_related_table_preprocessing_forward_and_cache() -> None: related_query, recipe=_recipe(), num_estimators=2, + estimator_batch_size=1, ), ) @@ -409,6 +418,7 @@ def test_related_table_preprocessing_forward_and_cache() -> None: related_context, recipe=_recipe(), num_estimators=2, + estimator_batch_size=1, ) assert model._cache is not None @@ -533,3 +543,67 @@ def test_ensemble_output_reduce() -> None: ) assert out.size() == (2, 3) + + +@pytest.mark.parametrize("batch_size", [1, 2, None]) +def test_estimator_callbacks(batch_size: int | None) -> None: + model = _RecordingModel() + x = torch.arange(30.0).view(5, 3, 2) + y = torch.zeros(5, 3, 1) + model.fit(x, y, estimator_batch_size=batch_size) + out = model.predict( + x, + callbacks=(MyCallback("affine", 2.0, 3.0, []),), + ) + torch.testing.assert_close(out.numerical, 2.0 * x + 3.0) + out = model( + x_context=x, + y_context=y, + x_query=x, + estimator_batch_size=batch_size, + callbacks=(MyCallback("affine", 2.0, 3.0, []),), + ) + 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) + model.fit(x, y, estimator_batch_size=1) + 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) diff --git a/test/models/test_estimator_batching.py b/test/models/test_estimator_batching.py new file mode 100644 index 000000000..8c917bda8 --- /dev/null +++ b/test/models/test_estimator_batching.py @@ -0,0 +1,171 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import functools +from typing import Literal + +import pytest +import torch + +import sdm.processing as sp +from sdm import CategoricalTensor, Recipe, TableTensor +from sdm.models import ICLModel, KumoTabular, TabFM, TabICLv2 +from sdm.models.kumo.tabular import model as kumo_module +from sdm.models.tabfm import model as tabfm_module +from sdm.models.tabiclv2 import model as tabicl_module +from sdm.testing import withCUDA + + +def _build( + name: str, + task: Literal["classification", "regression"], + monkeypatch: pytest.MonkeyPatch, +) -> ICLModel: + if name == "tabiclv2": + monkeypatch.setattr( + tabicl_module, + "_TabICLv2", + functools.partial( + tabicl_module._TabICLv2, + channels=16, + num_embedding_layers=2, + num_embedding_heads=2, + num_inducing_points=4, + num_readout_tokens=2, + num_icl_layers=2, + num_icl_heads=2, + ), + ) + model = TabICLv2(task=task, pretrained=False) + elif name == "tabfm": + monkeypatch.setattr( + tabfm_module, + "_TabFM", + functools.partial( + tabfm_module._TabFM, + channels=16, + num_embedding_layers=2, + num_embedding_col_heads=2, + num_embedding_row_heads=2, + num_inducing_points=4, + num_readout_tokens=2, + num_icl_layers=2, + num_icl_heads=2, + ), + ) + model = TabFM(task=task, pretrained=False) + else: + assert name == "kumo" + monkeypatch.setitem( + kumo_module.MODEL_KWARGS, + "small", + { + "cell_channels": 16, + "num_embedding_layers": 2, + "num_embedding_heads": 2, + "num_inducing_points": 4, + "num_readout_tokens": 2, + "icl_channels": 32, + "num_icl_layers": 2, + "num_icl_heads": 2, + }, + ) + model = KumoTabular(task=task, pretrained=False) + + # Zero-initialized residuals would hide errors in feature permutations. + for parameter in model.parameters(): + if not parameter.any(): + torch.nn.init.normal_(parameter, std=0.02) + return model + + +@withCUDA +@pytest.mark.parametrize("name", ["tabiclv2", "tabfm", "kumo"]) +@pytest.mark.parametrize("task", ["classification", "regression"]) +@pytest.mark.parametrize("batch_shape", [(), (2,)]) +def test_estimator_batching( + device: torch.device, + name: str, + task: Literal["classification", "regression"], + batch_shape: tuple[int, ...], + monkeypatch: pytest.MonkeyPatch, +) -> None: + model = _build(name, task, monkeypatch).to(device) + x = TableTensor( + numerical=torch.randn(*batch_shape, 12, 3, device=device), + categorical=CategoricalTensor.from_tensor( + (torch.arange(12, device=device) % 2) + .view(12, 1) + .expand(*batch_shape, 12, 1) + ), + ) + context, query = x.split(8, dim=-2) + if task == "classification": + y = (10 + 10 * (torch.arange(8, device=device) % 3)).view(8, 1) + y = y.expand(*batch_shape, 8, 1) + else: + y = torch.randn(*batch_shape, 8, 1, device=device) + + default = model.default_recipe() + # Keep every estimator's output so averaging cannot hide misalignment. + recipe = Recipe(features=default.features, target=default.target) + if batch_shape: + # AlignCategories in the default recipes currently only supports one + # leading batch dimension. Exercise nested model batches separately. + recipe = Recipe( + features=[sp.ToNumerical(), sp.Standardize(), sp.ShuffleColumns()], + target=sp.StypeDispatch( + categorical=sp.ShuffleCategories(), + numerical=sp.Standardize(), + ), + ) + + def predict(batch_size: int | None) -> tuple[TableTensor, TableTensor]: + model.fit( + x=context, + y=y, + recipe=recipe, + num_estimators=5, + estimator_batch_size=batch_size, + # Compare execution with identical stochastic preprocessing. + generator=torch.Generator(device=device).manual_seed(123), + ) + first = model.predict(query) + second = model.predict(query[..., :2, :]) + return first, second + + expected, expected_short = predict(1) + assert expected.size()[: 1 + len(batch_shape)] == (5, *batch_shape) + for batch_size in (2, None, 10): + actual, actual_short = predict(batch_size) + assert actual.schema == expected.schema + assert actual.dtype == expected.dtype + assert actual.device == device + assert torch.is_inference(actual) + torch.testing.assert_close( + actual.numerical, expected.numerical, atol=1e-4, rtol=1e-4 + ) + torch.testing.assert_close( + actual_short.numerical, + expected_short.numerical, + atol=1e-4, + rtol=1e-4, + ) + + for batch_size in (1, 2, None, 10): + actual = model( + x_context=context, + y_context=y, + x_query=query, + recipe=recipe, + num_estimators=5, + estimator_batch_size=batch_size, + generator=torch.Generator(device=device).manual_seed(123), + ) + assert actual.schema == expected.schema + assert actual.dtype == expected.dtype + assert actual.device == device + assert torch.is_inference(actual) + torch.testing.assert_close( + actual.numerical, expected.numerical, atol=1e-4, rtol=1e-4 + ) From 16875512403393db1f237db15cb5b3860a0d3426 Mon Sep 17 00:00:00 2001 From: Akihiro Nitta Date: Tue, 22 Sep 2026 21:03:25 +0000 Subject: [PATCH 02/11] update --- sdm/models/_batch.py | 163 ------------------------------- sdm/models/base.py | 158 ++++++++++++++++++++++++++++-- sdm/models/kumo/tabular/model.py | 2 +- sdm/models/tabfm/model.py | 3 +- 4 files changed, 154 insertions(+), 172 deletions(-) delete mode 100644 sdm/models/_batch.py diff --git a/sdm/models/_batch.py b/sdm/models/_batch.py deleted file mode 100644 index b6aba327b..000000000 --- a/sdm/models/_batch.py +++ /dev/null @@ -1,163 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 - -from collections.abc import Sequence -from typing import cast - -import torch -from torch import Tensor - -from sdm import RelatedTables, Stype, StypeLike, TableTensor -from sdm.processing.execution import MemberContext, MemberQuery -from sdm.tensor.table import TableSchema - - -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 len(contexts) == 1 or 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 _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/base.py b/sdm/models/base.py index caf573f62..afdd85234 100644 --- a/sdm/models/base.py +++ b/sdm/models/base.py @@ -14,6 +14,7 @@ Recipe, RelatedTables, Stype, + StypeLike, TableTensor, Task, TaskLike, @@ -21,12 +22,6 @@ from sdm._inference import inference_mode from sdm._warnings import warn_once from sdm.cache import Cache -from sdm.models._batch import ( - _output_columns, - _stack_contexts, - _stack_queries, - _unstack_output, -) from sdm.models.callback import Callback from sdm.processing.execution import ( MemberContext, @@ -727,3 +722,154 @@ 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 len(contexts) == 1 or 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 _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 6bc897a38..6b990a9a2 100644 --- a/sdm/models/kumo/tabular/model.py +++ b/sdm/models/kumo/tabular/model.py @@ -14,8 +14,8 @@ from sdm import Recipe, RelatedTables, Stype, TableTensor, Task, TaskLike from sdm.cache import Cache from sdm.models import ECOC, ICLModel -from sdm.models._batch import _categorical_mask 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 diff --git a/sdm/models/tabfm/model.py b/sdm/models/tabfm/model.py index 44b42b73e..581592f0f 100644 --- a/sdm/models/tabfm/model.py +++ b/sdm/models/tabfm/model.py @@ -28,9 +28,8 @@ from sdm import Recipe, RelatedTables, Stype, TableTensor, Task, TaskLike from sdm.cache import Cache -from sdm.models._batch import _categorical_mask 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 From d25345460a49e0a75166c386a132042c8a8ee57e Mon Sep 17 00:00:00 2001 From: Akihiro Nitta Date: Wed, 23 Sep 2026 04:21:42 +0000 Subject: [PATCH 03/11] actually dont expose it --- sdm/models/base.py | 293 ++++++++++++++++++------- sdm/models/kumo/tabular/model.py | 4 +- sdm/models/tabfm/model.py | 4 +- test/benchmark/test_tabular_model.py | 39 ++++ test/explain/test_gradient.py | 18 +- test/models/test_base.py | 23 +- test/models/test_estimator_batching.py | 161 ++++++++++++-- 7 files changed, 415 insertions(+), 127 deletions(-) create mode 100644 test/benchmark/test_tabular_model.py diff --git a/sdm/models/base.py b/sdm/models/base.py index afdd85234..a84fca68d 100644 --- a/sdm/models/base.py +++ b/sdm/models/base.py @@ -21,7 +21,7 @@ ) 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, @@ -103,7 +103,6 @@ def forward( *, recipe: Recipe | None = None, num_estimators: int | None = None, - estimator_batch_size: int | None = None, callbacks: Sequence[Callback] | None = None, generator: torch.Generator | None = None, **kwargs: Any, @@ -126,11 +125,6 @@ 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: Callbacks applied in sequence to this model call. generator: Pseudorandom number generator used for sampling during pre-processing and model execution. @@ -179,28 +173,8 @@ def forward( related_tables=related_query_tables, ) - batch_size = ( - len(contexts) - if estimator_batch_size is None - else estimator_batch_size - ) outs: list[TableTensor] = [] - for start in range(0, len(contexts), batch_size): - context_members = contexts[start : start + batch_size] - query_members = queries[start : start + batch_size] - if len(context_members) > 1: - for context, query in zip(context_members, query_members): - 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("no_grad" if requires_grad else "inference"): - context = _stack_contexts(context_members) - query = _stack_queries(query_members) + for context, query in zip(contexts, queries): for callback in callbacks: context = MemberContext( *callback.on_context_preprocessing_end(self, *context) @@ -233,11 +207,6 @@ def forward( related_query_tables=query.related_tables, cache=None, generator=generator, - _x_schemas=( - tuple(member.x.schema for member in context_members) - if len(context_members) > 1 - else (context.x.schema,) - ), **kwargs, ) @@ -245,14 +214,7 @@ def forward( out = callback.on_model_forward_end(self, out) out = cast(TableTensor, out.to(query.x.dtype)) - with inference_mode("grad" if requires_grad else "inference"): - outs.extend( - _unstack_output( - out=out, - num_members=len(context_members), - columns=_output_columns(context_members), - ) - ) + outs.append(out) # Regression: invert target before stacking estimator outputs. if contexts[0].y.numerical.size(-1) > 0: @@ -276,7 +238,7 @@ def fit( *, recipe: Recipe | None = None, num_estimators: int | None = None, - estimator_batch_size: int | None = None, + estimator_batch_size: int | None = 1, callbacks: Sequence[Callback] | None = None, generator: torch.Generator | None = None, **kwargs: Any, @@ -301,8 +263,8 @@ def fit( 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. Prediction - reuses the same batches. Larger batches use more device memory. + shapes, target columns, and feature categories. Larger batches + use more device memory. callbacks: Callbacks applied in sequence to this model call. generator: Pseudorandom number generator used for sampling during pre-processing and model execution. @@ -327,19 +289,16 @@ def fit( generator=generator, ) - batch_size = ( - len(contexts) - if estimator_batch_size is None - else estimator_batch_size - ) - starts = range(0, len(contexts), batch_size) + 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=batch_size, + estimator_batch_size=estimator_batch_size, ) for i, start in enumerate(starts): - members = contexts[start : start + batch_size] + members = contexts[start : start + estimator_batch_size] with inference_mode("no_grad"): context = _stack_contexts(members) for callback in callbacks: @@ -353,12 +312,8 @@ def fit( ) estimator_cache = Cache( x_schema=context.x.schema, - x_schemas=( - tuple(member.x.schema for member in members) - if len(members) > 1 - else (context.x.schema,) - ), - output_columns=_output_columns(members), + x_schemas=_batch_schemas(members, context.x.schema), + output_columns=_output_columns(members, context.y), y_schema=context.y.schema, related_tables_schema=context.related_tables.schema if context.related_tables is not None @@ -404,6 +359,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. @@ -416,6 +372,10 @@ 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. callbacks: Callbacks applied in sequence to this model call. Returns: @@ -454,9 +414,11 @@ def predict( RecipeExecution, self._cache["recipe_execution"], ) - batch_size = cast(int, self._cache["estimator_batch_size"]) - starts = range(0, recipe_execution.num_members, batch_size) - caches = [cast(Cache, self._cache[i]) for i in range(len(starts))] + if estimator_batch_size is None: + estimator_batch_size = recipe_execution.num_members + 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 @@ -487,22 +449,7 @@ def predict( for i, start in enumerate(starts): cache, next_cache = next_cache, None assert cache is not None - members = queries[start : start + batch_size] - if len(members) > 1: - for schema, member in zip( - cast(tuple[TableSchema, ...], cache["x_schemas"]), - members, - strict=True, - ): - self._validate_query( - x_context=schema, - x_query=member.x, - related_context_tables=cast( - RelatedTablesSchema | None, - cache["related_tables_schema"], - ), - related_query_tables=member.related_tables, - ) + members = queries[start : start + estimator_batch_size] with inference_mode( "no_grad" if requires_grad else "inference" ): @@ -521,6 +468,13 @@ def predict( ), related_query_tables=query.related_tables, ) + if len(members) > 1 and cache["x_schemas"] != _batch_schemas( + members, query.x.schema + ): + raise ValueError( + "Expected context and query features to share " + "the same schema" + ) if i + 1 < len(caches): next_cache = caches[i + 1] @@ -541,6 +495,32 @@ def predict( **cast(dict[str, Any], self._cache["kwargs"]), ) + output_columns = cast( + tuple[tuple[str, ...], ...] | None, + cache["output_columns"], + ) + if callbacks and output_columns is not None: + outputs = _unstack_output( + out=out, + num_members=len(members), + columns=output_columns, + ) + out = outputs[0] + if len(outputs) > 1: + names = out.columns[Stype.numerical] + # Align raw tensors to preserve gradients. + values: list[Tensor] = [] + for output in outputs: + columns = output.columns[Stype.numerical] + indices = [ + columns.index(name) for name in names + ] + values.append(output.numerical[..., indices]) + out = TableTensor( + columns={Stype.numerical: names}, + numerical=torch.stack(values), + ) + output_columns = None for callback in callbacks: out = callback.on_model_forward_end(self, out) @@ -555,10 +535,7 @@ def predict( _unstack_output( out=out, num_members=len(members), - columns=cast( - tuple[tuple[str, ...], ...] | None, - cache["output_columns"], - ), + columns=output_columns, ) ) @@ -620,6 +597,7 @@ def _forward( generator: torch.Generator | None, **kwargs: Any, ) -> TableTensor: # [..., R_query, *] + # Regroupable cache tensors retain all leading input batch dimensions. pass @classmethod @@ -817,18 +795,167 @@ def _stack_queries(queries: Sequence[MemberQuery]) -> MemberQuery: ) +def _remap_columns( + columns: Sequence[tuple[str, ...]], + before: tuple[str, ...], + after: tuple[str, ...], +) -> tuple[tuple[str, ...], ...]: + if before == after: + return tuple(columns) + positions = {column: i for i, column in enumerate(before)} + return tuple( + tuple( + names[positions[column]] if column in positions else column + for column in after + ) + for names in columns + ) + + +def _batch_schemas( + members: Sequence[MemberContext] | Sequence[MemberQuery], + schema: TableSchema, +) -> tuple[TableSchema, ...]: + schemas = tuple(member.x.schema for member in members) + if schemas[0] == schema: + return schemas + # Apply callback column selections/reordering to each estimator's layout. + columns = _remap_columns( + columns=[s.columns[Stype.numerical] for s in schemas], + before=schemas[0].columns[Stype.numerical], + after=schema.columns[Stype.numerical], + ) + return tuple( + TableSchema(columns={**schema.columns, Stype.numerical: names}) + for names in columns + ) + + def _output_columns( contexts: Sequence[MemberContext], + y: TableTensor, ) -> tuple[tuple[str, ...], ...] | None: - if len(contexts) == 1 or contexts[0].y.categorical.size(-1) == 0: + if y.categorical.size(-1) == 0: return None - return tuple( + names = tuple(str(value) for value in y.categorical.categories[0].tolist()) + if contexts[0].y.categorical.size(-1) == 0: + return (names,) * len(contexts) + columns = tuple( tuple( str(value) for value in context.y.categorical.categories[0].tolist() ) for context in contexts ) + if len(names) == len(columns[0]) and set(names) != set(columns[0]): + renamed = dict(zip(columns[0], names, strict=True)) + return tuple( + tuple(renamed[name] for name in group) for group in columns + ) + return _remap_columns(columns=columns, before=columns[0], after=names) + + +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 + + metadata_keys = { + "x_schema", + "x_schemas", + "y_schema", + "related_tables_schema", + "classes", + "output_columns", + } + 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 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 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( diff --git a/sdm/models/kumo/tabular/model.py b/sdm/models/kumo/tabular/model.py index 6b990a9a2..87f17d4a2 100644 --- a/sdm/models/kumo/tabular/model.py +++ b/sdm/models/kumo/tabular/model.py @@ -222,15 +222,15 @@ def _forward( categorical_mask = _categorical_mask( x=x_context, schema=schema, - schemas=kwargs.get("_x_schemas", (x_context.schema,)) + 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 581592f0f..783ffb942 100644 --- a/sdm/models/tabfm/model.py +++ b/sdm/models/tabfm/model.py @@ -210,15 +210,15 @@ def _forward( categorical_mask = _categorical_mask( x=x_context, schema=schema, - schemas=kwargs.get("_x_schemas", (x_context.schema,)) + 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/benchmark/test_tabular_model.py b/test/benchmark/test_tabular_model.py new file mode 100644 index 000000000..b726dede0 --- /dev/null +++ b/test/benchmark/test_tabular_model.py @@ -0,0 +1,39 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import pytest + +pytest.importorskip("autogluon.tabular") +pytest.importorskip("tabarena") + +from benchmark.tabular.model import _kumo_prediction_batch_size + + +@pytest.mark.parametrize( + ("context_rows", "query_rows", "columns", "memory_gib", "expected"), + [ + (1_000, 500, 30, 24, 8), + (1_000, 500, 30, 4, 1), + (10_000, 500, 30, 240, 1), + (1_000, 10_000, 30, 240, 1), + (1_000, 500, 100, 240, 1), + (1, 1, 500, 8, 1), + ], +) +def test_kumo_prediction_batch_size( + context_rows: int, + query_rows: int, + columns: int, + memory_gib: int, + expected: int, +) -> None: + assert ( + _kumo_prediction_batch_size( + context_rows=context_rows, + query_rows=query_rows, + columns=columns, + num_estimators=8, + available_memory=memory_gib * 2**30, + ) + == expected + ) diff --git a/test/explain/test_gradient.py b/test/explain/test_gradient.py index 3ef44a915..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 @@ -52,12 +52,20 @@ def default_recipe(cls) -> Recipe: ("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 + 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={ @@ -99,8 +107,6 @@ def test_returns_query_input_gradients( num_estimators=num_estimators, ) - if num_estimators > 1: - x_query = x_query.expand(num_estimators, *x_query.size()) torch.testing.assert_close( result.x.numerical, torch.full_like(x_query, 2.0) ) diff --git a/test/models/test_base.py b/test/models/test_base.py index bfa8157ae..5d6ae09f6 100644 --- a/test/models/test_base.py +++ b/test/models/test_base.py @@ -369,7 +369,6 @@ def test_related_table_preprocessing_forward_and_cache() -> None: related_query, recipe=_recipe(), num_estimators=2, - estimator_batch_size=1, ), ) @@ -418,7 +417,6 @@ def test_related_table_preprocessing_forward_and_cache() -> None: related_context, recipe=_recipe(), num_estimators=2, - estimator_batch_size=1, ) assert model._cache is not None @@ -545,12 +543,12 @@ def test_ensemble_output_reduce() -> None: assert out.size() == (2, 3) -@pytest.mark.parametrize("batch_size", [1, 2, None]) -def test_estimator_callbacks(batch_size: int | None) -> None: +@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) - model.fit(x, y, estimator_batch_size=batch_size) + model.fit(x, y, estimator_batch_size=estimator_batch_size) out = model.predict( x, callbacks=(MyCallback("affine", 2.0, 3.0, []),), @@ -560,7 +558,6 @@ def test_estimator_callbacks(batch_size: int | None) -> None: x_context=x, y_context=y, x_query=x, - estimator_batch_size=batch_size, callbacks=(MyCallback("affine", 2.0, 3.0, []),), ) torch.testing.assert_close(out.numerical, 2.0 * x + 3.0) @@ -583,11 +580,12 @@ def test_estimator_batching_incompatible_shapes() -> None: member_table_ids=(0, 1), ) with pytest.raises(RuntimeError, match="stack expects"): - model.fit(x, y) - model.fit(x, y, estimator_batch_size=1) + model.fit(x, y, 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 + model.predict(TableTensor(numerical=query)).numerical, + query, ) @@ -606,4 +604,9 @@ def test_estimator_batching_incompatible_categories() -> None: member_table_ids=(0, 1), ) with pytest.raises(ValueError, match="matching class counts"): - model.fit(torch.ones(3, 2), y, num_estimators=2) + model.fit( + torch.ones(3, 2), y, 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) diff --git a/test/models/test_estimator_batching.py b/test/models/test_estimator_batching.py index 8c917bda8..57d061bf8 100644 --- a/test/models/test_estimator_batching.py +++ b/test/models/test_estimator_batching.py @@ -8,8 +8,9 @@ import torch import sdm.processing as sp -from sdm import CategoricalTensor, Recipe, TableTensor +from sdm import CategoricalTensor, Recipe, RelatedTables, Stype, TableTensor from sdm.models import ICLModel, KumoTabular, TabFM, TabICLv2 +from sdm.models.callback import Callback from sdm.models.kumo.tabular import model as kumo_module from sdm.models.tabfm import model as tabfm_module from sdm.models.tabiclv2 import model as tabicl_module @@ -70,7 +71,7 @@ def _build( "num_icl_heads": 2, }, ) - model = KumoTabular(task=task, pretrained=False) + model = KumoTabular(task=task, size="small", pretrained=False) # Zero-initialized residuals would hide errors in feature permutations. for parameter in model.parameters(): @@ -120,24 +121,28 @@ def test_estimator_batching( ), ) - def predict(batch_size: int | None) -> tuple[TableTensor, TableTensor]: + def predict( + estimator_batch_size: int | None, + ) -> tuple[TableTensor, TableTensor]: model.fit( x=context, y=y, recipe=recipe, - num_estimators=5, - estimator_batch_size=batch_size, + num_estimators=9, + estimator_batch_size=estimator_batch_size, # Compare execution with identical stochastic preprocessing. generator=torch.Generator(device=device).manual_seed(123), ) - first = model.predict(query) - second = model.predict(query[..., :2, :]) + first = model.predict(query, estimator_batch_size=estimator_batch_size) + second = model.predict( + query[..., :2, :], estimator_batch_size=estimator_batch_size + ) return first, second expected, expected_short = predict(1) - assert expected.size()[: 1 + len(batch_shape)] == (5, *batch_shape) - for batch_size in (2, None, 10): - actual, actual_short = predict(batch_size) + assert expected.size()[: 1 + len(batch_shape)] == (9, *batch_shape) + for estimator_batch_size in (2, None, 10): + actual, actual_short = predict(estimator_batch_size) assert actual.schema == expected.schema assert actual.dtype == expected.dtype assert actual.device == device @@ -152,20 +157,128 @@ def predict(batch_size: int | None) -> tuple[TableTensor, TableTensor]: rtol=1e-4, ) - for batch_size in (1, 2, None, 10): - actual = model( - x_context=context, - y_context=y, - x_query=query, + # Reuse the fit across prediction groupings, including partial batches. + for fit_estimator_batch_size in (1, 2, None): + predict(fit_estimator_batch_size) + for predict_estimator_batch_size in (None, 8, 2, 1): + actual = model.predict( + query, estimator_batch_size=predict_estimator_batch_size + ) + assert actual.schema == expected.schema + torch.testing.assert_close( + actual.numerical, expected.numerical, atol=1e-4, rtol=1e-4 + ) + + actual = model( + x_context=context, + y_context=y, + x_query=query, + recipe=recipe, + num_estimators=9, + generator=torch.Generator(device=device).manual_seed(123), + ) + assert actual.schema == expected.schema + assert actual.dtype == expected.dtype + assert actual.device == device + assert torch.is_inference(actual) + torch.testing.assert_close( + actual.numerical, expected.numerical, atol=1e-4, rtol=1e-4 + ) + + +@withCUDA +@pytest.mark.parametrize("name", ["tabfm", "kumo"]) +def test_estimator_batching_callback_columns( + device: torch.device, + name: str, + monkeypatch: pytest.MonkeyPatch, +) -> None: + class SelectColumns(Callback): + requires_grad = True + + def on_context_preprocessing_end( + self, + model: torch.nn.Module, + x: TableTensor, + y: TableTensor, + related_tables: RelatedTables[TableTensor] | None, + ) -> tuple[ + TableTensor, TableTensor, RelatedTables[TableTensor] | None + ]: + x = x.select_columns(x.columns[Stype.numerical][::2]) + y = y.replace_blocks( + categorical=CategoricalTensor( + code=y.categorical.code, + categories=tuple( + c + 100 for c in y.categorical.categories + ), + ), + ) + return x, y, related_tables + + def on_query_preprocessing_end( + self, + model: torch.nn.Module, + x: TableTensor, + related_tables: RelatedTables[TableTensor] | None, + ) -> tuple[TableTensor, RelatedTables[TableTensor] | None]: + x = x.select_columns(x.columns[Stype.numerical][::2]) + self.numerical = x.numerical.detach().requires_grad_(True) + return x.replace_blocks(numerical=self.numerical), related_tables + + def on_model_forward_end( + self, model: torch.nn.Module, out: TableTensor + ) -> TableTensor: + out = out.select_columns("120") + (gradient,) = torch.autograd.grad( + out.numerical.sum(), self.numerical + ) + assert gradient.isfinite().all() + return out + + model = _build(name, "classification", monkeypatch).to(device) + x = TableTensor( + numerical=torch.randn(9, 2, device=device), + categorical=CategoricalTensor.from_tensor( + (torch.arange(9, device=device) % 2).view(-1, 1) + ), + ) + y = TableTensor( + categorical=CategoricalTensor.from_tensor( + (10 + 10 * (torch.arange(9, device=device) % 3)).view(-1, 1) + ), + ) + recipe = Recipe( + features=[sp.ToNumerical(), sp.ShuffleColumns()], + target=sp.ShuffleCategories(), + ) + callbacks = (SelectColumns(),) + expected = model( + x_context=x, + y_context=y, + x_query=x, + recipe=recipe, + num_estimators=3, + callbacks=callbacks, + generator=torch.Generator(device=device).manual_seed(123), + ) + for fit_estimator_batch_size in (1, 2, None): + model.fit( + x=x, + y=y, recipe=recipe, - num_estimators=5, - estimator_batch_size=batch_size, + num_estimators=3, + estimator_batch_size=fit_estimator_batch_size, + callbacks=callbacks, generator=torch.Generator(device=device).manual_seed(123), ) - assert actual.schema == expected.schema - assert actual.dtype == expected.dtype - assert actual.device == device - assert torch.is_inference(actual) - torch.testing.assert_close( - actual.numerical, expected.numerical, atol=1e-4, rtol=1e-4 - ) + for predict_estimator_batch_size in (1, 2, None): + actual = model.predict( + x, + estimator_batch_size=predict_estimator_batch_size, + callbacks=callbacks, + ) + assert actual.schema == expected.schema + torch.testing.assert_close( + actual.numerical, expected.numerical, atol=1e-4, rtol=1e-4 + ) From 7ef2d5a42d2d1b616b9c43e30fedb6599a475d03 Mon Sep 17 00:00:00 2001 From: Akihiro Nitta Date: Wed, 23 Sep 2026 04:23:03 +0000 Subject: [PATCH 04/11] auto --- benchmark/tabular/model.py | 64 +++++++++++++++++++++++++++++++++++++- 1 file changed, 63 insertions(+), 1 deletion(-) diff --git a/benchmark/tabular/model.py b/benchmark/tabular/model.py index 5a5caead5..a8e57a19e 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,10 @@ 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.size(-2), + x_context.numerical.size(-1) + 2 * x_context.categorical.size(-1), + ) recipe = self._create_recipe() if params["max_columns"] is not None: @@ -146,6 +151,7 @@ def _fit( y=y_context, recipe=recipe, num_estimators=num_estimators, + estimator_batch_size=1, generator=generator, ) return @@ -202,7 +208,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._prediction_batch_size(x_query), + ) else: with ( torch.inference_mode(), @@ -251,6 +260,35 @@ def _predict_proba( probabilities = out.numerical[..., indices].float().cpu().numpy() return self._convert_proba_to_unified_form(probabilities) + def _prediction_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 + or not isinstance(self.model, sdm.models.KumoTabular) + ): + return 1 + + free, total = torch.cuda.mem_get_info(x.device) + allocated = torch.cuda.memory_allocated(x.device) + available = min( + free + torch.cuda.memory_reserved(x.device) - allocated, + total * torch.cuda.get_per_process_memory_fraction(x.device) + - allocated, + ) + rows, columns = self._context_shape + return _kumo_prediction_batch_size( + context_rows=rows, + query_rows=x.size(-2), + columns=columns, + num_estimators=self._num_estimators, + available_memory=available, + ) + def get_device(self) -> str: return str(next(self.model.parameters()).device) @@ -415,3 +453,27 @@ def beyondarena_method_name(self) -> str: model_cls=SDMTabFMModel, ), } + + +def _kumo_prediction_batch_size( + context_rows: int, + query_rows: int, + columns: int, + num_estimators: int, + available_memory: float, +) -> int: + rows = max(context_rows, query_rows) + if rows > 2_000 or rows * columns > 50_000: + return 1 + # Kumo large: 24 ICL layers, 6 column layers, 16-bit keys and values. + cache_bytes = ( + num_estimators + * 2 + * 2 + * (24 * context_rows * 128 + 6 * columns * 256 * 256) + ) + # Conservative working-memory allowance, including four readout tokens. + peak_bytes = ( + cache_bytes + num_estimators * query_rows * (columns + 4) * 8_192 + ) + return num_estimators if peak_bytes <= available_memory * 0.25 else 1 From e6b0653ab8e4c14d16a1a58a2d4c5f13d9dfba70 Mon Sep 17 00:00:00 2001 From: Akihiro Nitta Date: Wed, 23 Sep 2026 05:03:23 +0000 Subject: [PATCH 05/11] update --- benchmark/tabular/model.py | 63 ++++++---------------------- sdm/models/base.py | 31 +++++++++----- test/benchmark/test_tabular_model.py | 39 ----------------- 3 files changed, 34 insertions(+), 99 deletions(-) delete mode 100644 test/benchmark/test_tabular_model.py diff --git a/benchmark/tabular/model.py b/benchmark/tabular/model.py index a8e57a19e..dfccc0743 100644 --- a/benchmark/tabular/model.py +++ b/benchmark/tabular/model.py @@ -21,6 +21,8 @@ import sdm import sdm.processing as sp +from sdm.cache import Cache +from sdm.models.base import _can_batch_cache from sdm.processing.execution import RecipeExecution Task = Literal["classification", "regression"] @@ -128,10 +130,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.size(-2), - x_context.numerical.size(-1) + 2 * x_context.categorical.size(-1), - ) + self._context_shape = x_context.shape[-2:] recipe = self._create_recipe() if params["max_columns"] is not None: @@ -210,7 +209,7 @@ def _predict_proba( ): out = self.model.predict( x_query, - estimator_batch_size=self._prediction_batch_size(x_query), + estimator_batch_size=self._estimator_batch_size(x_query), ) else: with ( @@ -260,34 +259,22 @@ def _predict_proba( probabilities = out.numerical[..., indices].float().cpu().numpy() return self._convert_proba_to_unified_form(probabilities) - def _prediction_batch_size(self, x: torch.Tensor) -> int | None: + 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 - or not isinstance(self.model, sdm.models.KumoTabular) - ): + if not x.is_cuda or self._expand_query: return 1 - free, total = torch.cuda.mem_get_info(x.device) - allocated = torch.cuda.memory_allocated(x.device) - available = min( - free + torch.cuda.memory_reserved(x.device) - allocated, - total * torch.cuda.get_per_process_memory_fraction(x.device) - - allocated, - ) - rows, columns = self._context_shape - return _kumo_prediction_batch_size( - context_rows=rows, - query_rows=x.size(-2), - columns=columns, - num_estimators=self._num_estimators, - available_memory=available, - ) + num_rows, num_cols = self._context_shape + num_rows = max(num_rows, x.size(-2)) + if num_rows > 2_000 or num_rows * num_cols >= 50_000: + return 1 + if not _can_batch_cache(cast(Cache, self.model._cache)): + return 1 + return self._num_estimators def get_device(self) -> str: return str(next(self.model.parameters()).device) @@ -453,27 +440,3 @@ def beyondarena_method_name(self) -> str: model_cls=SDMTabFMModel, ), } - - -def _kumo_prediction_batch_size( - context_rows: int, - query_rows: int, - columns: int, - num_estimators: int, - available_memory: float, -) -> int: - rows = max(context_rows, query_rows) - if rows > 2_000 or rows * columns > 50_000: - return 1 - # Kumo large: 24 ICL layers, 6 column layers, 16-bit keys and values. - cache_bytes = ( - num_estimators - * 2 - * 2 - * (24 * context_rows * 128 + 6 * columns * 256 * 256) - ) - # Conservative working-memory allowance, including four readout tokens. - peak_bytes = ( - cache_bytes + num_estimators * query_rows * (columns + 4) * 8_192 - ) - return num_estimators if peak_bytes <= available_memory * 0.25 else 1 diff --git a/sdm/models/base.py b/sdm/models/base.py index a84fca68d..77794ff1a 100644 --- a/sdm/models/base.py +++ b/sdm/models/base.py @@ -31,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. @@ -855,6 +864,16 @@ def _output_columns( return _remap_columns(columns=columns, before=columns[0], after=names) +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"]) @@ -867,14 +886,6 @@ def _prediction_caches(cache: Cache, estimator_batch_size: int) -> list[Cache]: ): return caches - metadata_keys = { - "x_schema", - "x_schemas", - "y_schema", - "related_tables_schema", - "classes", - "output_columns", - } members: list[Cache] = [] for fitted in caches: schemas = cast(tuple[TableSchema, ...], fitted["x_schemas"]) @@ -890,7 +901,7 @@ def _prediction_caches(cache: Cache, estimator_batch_size: int) -> list[Cache]: ) if len(schemas) > 1: for key, value in fitted.items(): - if key not in metadata_keys: + if key not in _CACHE_METADATA_KEYS: if isinstance(value, Tensor): member[key] = value[index] elif isinstance(value, KVCacheEntry): @@ -940,7 +951,7 @@ def _prediction_caches(cache: Cache, estimator_batch_size: int) -> list[Cache]: ) ) for key, value in first.items(): - if key not in metadata_keys: + 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)) diff --git a/test/benchmark/test_tabular_model.py b/test/benchmark/test_tabular_model.py deleted file mode 100644 index b726dede0..000000000 --- a/test/benchmark/test_tabular_model.py +++ /dev/null @@ -1,39 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 - -import pytest - -pytest.importorskip("autogluon.tabular") -pytest.importorskip("tabarena") - -from benchmark.tabular.model import _kumo_prediction_batch_size - - -@pytest.mark.parametrize( - ("context_rows", "query_rows", "columns", "memory_gib", "expected"), - [ - (1_000, 500, 30, 24, 8), - (1_000, 500, 30, 4, 1), - (10_000, 500, 30, 240, 1), - (1_000, 10_000, 30, 240, 1), - (1_000, 500, 100, 240, 1), - (1, 1, 500, 8, 1), - ], -) -def test_kumo_prediction_batch_size( - context_rows: int, - query_rows: int, - columns: int, - memory_gib: int, - expected: int, -) -> None: - assert ( - _kumo_prediction_batch_size( - context_rows=context_rows, - query_rows=query_rows, - columns=columns, - num_estimators=8, - available_memory=memory_gib * 2**30, - ) - == expected - ) From 1829d5349ece43b7d50d5efcce627b9b4ad7fbe5 Mon Sep 17 00:00:00 2001 From: Akihiro Nitta Date: Wed, 23 Sep 2026 05:12:17 +0000 Subject: [PATCH 06/11] update --- test/models/tabfm/test_model.py | 31 ++- test/models/test_estimator_batching.py | 284 ------------------------- 2 files changed, 28 insertions(+), 287 deletions(-) delete mode 100644 test/models/test_estimator_batching.py diff --git a/test/models/tabfm/test_model.py b/test/models/tabfm/test_model.py index 2c08f74be..d92048342 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,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,8 +81,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/test_estimator_batching.py b/test/models/test_estimator_batching.py deleted file mode 100644 index 57d061bf8..000000000 --- a/test/models/test_estimator_batching.py +++ /dev/null @@ -1,284 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 - -import functools -from typing import Literal - -import pytest -import torch - -import sdm.processing as sp -from sdm import CategoricalTensor, Recipe, RelatedTables, Stype, TableTensor -from sdm.models import ICLModel, KumoTabular, TabFM, TabICLv2 -from sdm.models.callback import Callback -from sdm.models.kumo.tabular import model as kumo_module -from sdm.models.tabfm import model as tabfm_module -from sdm.models.tabiclv2 import model as tabicl_module -from sdm.testing import withCUDA - - -def _build( - name: str, - task: Literal["classification", "regression"], - monkeypatch: pytest.MonkeyPatch, -) -> ICLModel: - if name == "tabiclv2": - monkeypatch.setattr( - tabicl_module, - "_TabICLv2", - functools.partial( - tabicl_module._TabICLv2, - channels=16, - num_embedding_layers=2, - num_embedding_heads=2, - num_inducing_points=4, - num_readout_tokens=2, - num_icl_layers=2, - num_icl_heads=2, - ), - ) - model = TabICLv2(task=task, pretrained=False) - elif name == "tabfm": - monkeypatch.setattr( - tabfm_module, - "_TabFM", - functools.partial( - tabfm_module._TabFM, - channels=16, - num_embedding_layers=2, - num_embedding_col_heads=2, - num_embedding_row_heads=2, - num_inducing_points=4, - num_readout_tokens=2, - num_icl_layers=2, - num_icl_heads=2, - ), - ) - model = TabFM(task=task, pretrained=False) - else: - assert name == "kumo" - monkeypatch.setitem( - kumo_module.MODEL_KWARGS, - "small", - { - "cell_channels": 16, - "num_embedding_layers": 2, - "num_embedding_heads": 2, - "num_inducing_points": 4, - "num_readout_tokens": 2, - "icl_channels": 32, - "num_icl_layers": 2, - "num_icl_heads": 2, - }, - ) - model = KumoTabular(task=task, size="small", pretrained=False) - - # Zero-initialized residuals would hide errors in feature permutations. - for parameter in model.parameters(): - if not parameter.any(): - torch.nn.init.normal_(parameter, std=0.02) - return model - - -@withCUDA -@pytest.mark.parametrize("name", ["tabiclv2", "tabfm", "kumo"]) -@pytest.mark.parametrize("task", ["classification", "regression"]) -@pytest.mark.parametrize("batch_shape", [(), (2,)]) -def test_estimator_batching( - device: torch.device, - name: str, - task: Literal["classification", "regression"], - batch_shape: tuple[int, ...], - monkeypatch: pytest.MonkeyPatch, -) -> None: - model = _build(name, task, monkeypatch).to(device) - x = TableTensor( - numerical=torch.randn(*batch_shape, 12, 3, device=device), - categorical=CategoricalTensor.from_tensor( - (torch.arange(12, device=device) % 2) - .view(12, 1) - .expand(*batch_shape, 12, 1) - ), - ) - context, query = x.split(8, dim=-2) - if task == "classification": - y = (10 + 10 * (torch.arange(8, device=device) % 3)).view(8, 1) - y = y.expand(*batch_shape, 8, 1) - else: - y = torch.randn(*batch_shape, 8, 1, device=device) - - default = model.default_recipe() - # Keep every estimator's output so averaging cannot hide misalignment. - recipe = Recipe(features=default.features, target=default.target) - if batch_shape: - # AlignCategories in the default recipes currently only supports one - # leading batch dimension. Exercise nested model batches separately. - recipe = Recipe( - features=[sp.ToNumerical(), sp.Standardize(), sp.ShuffleColumns()], - target=sp.StypeDispatch( - categorical=sp.ShuffleCategories(), - numerical=sp.Standardize(), - ), - ) - - def predict( - estimator_batch_size: int | None, - ) -> tuple[TableTensor, TableTensor]: - model.fit( - x=context, - y=y, - recipe=recipe, - num_estimators=9, - estimator_batch_size=estimator_batch_size, - # Compare execution with identical stochastic preprocessing. - generator=torch.Generator(device=device).manual_seed(123), - ) - first = model.predict(query, estimator_batch_size=estimator_batch_size) - second = model.predict( - query[..., :2, :], estimator_batch_size=estimator_batch_size - ) - return first, second - - expected, expected_short = predict(1) - assert expected.size()[: 1 + len(batch_shape)] == (9, *batch_shape) - for estimator_batch_size in (2, None, 10): - actual, actual_short = predict(estimator_batch_size) - assert actual.schema == expected.schema - assert actual.dtype == expected.dtype - assert actual.device == device - assert torch.is_inference(actual) - torch.testing.assert_close( - actual.numerical, expected.numerical, atol=1e-4, rtol=1e-4 - ) - torch.testing.assert_close( - actual_short.numerical, - expected_short.numerical, - atol=1e-4, - rtol=1e-4, - ) - - # Reuse the fit across prediction groupings, including partial batches. - for fit_estimator_batch_size in (1, 2, None): - predict(fit_estimator_batch_size) - for predict_estimator_batch_size in (None, 8, 2, 1): - actual = model.predict( - query, estimator_batch_size=predict_estimator_batch_size - ) - assert actual.schema == expected.schema - torch.testing.assert_close( - actual.numerical, expected.numerical, atol=1e-4, rtol=1e-4 - ) - - actual = model( - x_context=context, - y_context=y, - x_query=query, - recipe=recipe, - num_estimators=9, - generator=torch.Generator(device=device).manual_seed(123), - ) - assert actual.schema == expected.schema - assert actual.dtype == expected.dtype - assert actual.device == device - assert torch.is_inference(actual) - torch.testing.assert_close( - actual.numerical, expected.numerical, atol=1e-4, rtol=1e-4 - ) - - -@withCUDA -@pytest.mark.parametrize("name", ["tabfm", "kumo"]) -def test_estimator_batching_callback_columns( - device: torch.device, - name: str, - monkeypatch: pytest.MonkeyPatch, -) -> None: - class SelectColumns(Callback): - requires_grad = True - - def on_context_preprocessing_end( - self, - model: torch.nn.Module, - x: TableTensor, - y: TableTensor, - related_tables: RelatedTables[TableTensor] | None, - ) -> tuple[ - TableTensor, TableTensor, RelatedTables[TableTensor] | None - ]: - x = x.select_columns(x.columns[Stype.numerical][::2]) - y = y.replace_blocks( - categorical=CategoricalTensor( - code=y.categorical.code, - categories=tuple( - c + 100 for c in y.categorical.categories - ), - ), - ) - return x, y, related_tables - - def on_query_preprocessing_end( - self, - model: torch.nn.Module, - x: TableTensor, - related_tables: RelatedTables[TableTensor] | None, - ) -> tuple[TableTensor, RelatedTables[TableTensor] | None]: - x = x.select_columns(x.columns[Stype.numerical][::2]) - self.numerical = x.numerical.detach().requires_grad_(True) - return x.replace_blocks(numerical=self.numerical), related_tables - - def on_model_forward_end( - self, model: torch.nn.Module, out: TableTensor - ) -> TableTensor: - out = out.select_columns("120") - (gradient,) = torch.autograd.grad( - out.numerical.sum(), self.numerical - ) - assert gradient.isfinite().all() - return out - - model = _build(name, "classification", monkeypatch).to(device) - x = TableTensor( - numerical=torch.randn(9, 2, device=device), - categorical=CategoricalTensor.from_tensor( - (torch.arange(9, device=device) % 2).view(-1, 1) - ), - ) - y = TableTensor( - categorical=CategoricalTensor.from_tensor( - (10 + 10 * (torch.arange(9, device=device) % 3)).view(-1, 1) - ), - ) - recipe = Recipe( - features=[sp.ToNumerical(), sp.ShuffleColumns()], - target=sp.ShuffleCategories(), - ) - callbacks = (SelectColumns(),) - expected = model( - x_context=x, - y_context=y, - x_query=x, - recipe=recipe, - num_estimators=3, - callbacks=callbacks, - generator=torch.Generator(device=device).manual_seed(123), - ) - for fit_estimator_batch_size in (1, 2, None): - model.fit( - x=x, - y=y, - recipe=recipe, - num_estimators=3, - estimator_batch_size=fit_estimator_batch_size, - callbacks=callbacks, - generator=torch.Generator(device=device).manual_seed(123), - ) - for predict_estimator_batch_size in (1, 2, None): - actual = model.predict( - x, - estimator_batch_size=predict_estimator_batch_size, - callbacks=callbacks, - ) - assert actual.schema == expected.schema - torch.testing.assert_close( - actual.numerical, expected.numerical, atol=1e-4, rtol=1e-4 - ) From 414a2e03d08f307b2e2ac5769b94279d180bef98 Mon Sep 17 00:00:00 2001 From: Akihiro Nitta Date: Wed, 23 Sep 2026 05:35:50 +0000 Subject: [PATCH 07/11] update --- sdm/models/base.py | 118 ++++++++++----------------------------- test/models/test_base.py | 23 ++++++-- 2 files changed, 47 insertions(+), 94 deletions(-) diff --git a/sdm/models/base.py b/sdm/models/base.py index 77794ff1a..846054168 100644 --- a/sdm/models/base.py +++ b/sdm/models/base.py @@ -273,13 +273,15 @@ def fit( 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. + 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() @@ -314,6 +316,8 @@ def fit( context = MemberContext( *callback.on_context_preprocessing_end(self, *context) ) + if len(members) == 1: + members = (context,) self._validate_context( x=context.x, y=context.y, @@ -321,8 +325,8 @@ def fit( ) estimator_cache = Cache( x_schema=context.x.schema, - x_schemas=_batch_schemas(members, context.x.schema), - output_columns=_output_columns(members, context.y), + 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 @@ -385,6 +389,7 @@ def predict( 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. + Callbacks require ``1``. callbacks: Callbacks applied in sequence to this model call. Returns: @@ -399,6 +404,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: @@ -477,8 +484,8 @@ def predict( ), related_query_tables=query.related_tables, ) - if len(members) > 1 and cache["x_schemas"] != _batch_schemas( - members, query.x.schema + 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 " @@ -504,50 +511,29 @@ def predict( **cast(dict[str, Any], self._cache["kwargs"]), ) - output_columns = cast( - tuple[tuple[str, ...], ...] | None, - cache["output_columns"], + outputs = _unstack_output( + out=out, + num_members=len(members), + columns=cast( + tuple[tuple[str, ...], ...] | None, + cache["output_columns"], + ), ) - if callbacks and output_columns is not None: - outputs = _unstack_output( - out=out, - num_members=len(members), - columns=output_columns, - ) - out = outputs[0] - if len(outputs) > 1: - names = out.columns[Stype.numerical] - # Align raw tensors to preserve gradients. - values: list[Tensor] = [] - for output in outputs: - columns = output.columns[Stype.numerical] - indices = [ - columns.index(name) for name in names - ] - values.append(output.numerical[..., indices]) - out = TableTensor( - columns={Stype.numerical: names}, - numerical=torch.stack(values), - ) - output_columns = None 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)) - with inference_mode("grad" if requires_grad else "inference"): - outs.extend( - _unstack_output( - out=out, - num_members=len(members), - columns=output_columns, - ) - ) - if x.is_cuda and next_cache is not None: assert compute_stream is not None assert transfer_stream is not None @@ -804,64 +790,18 @@ def _stack_queries(queries: Sequence[MemberQuery]) -> MemberQuery: ) -def _remap_columns( - columns: Sequence[tuple[str, ...]], - before: tuple[str, ...], - after: tuple[str, ...], -) -> tuple[tuple[str, ...], ...]: - if before == after: - return tuple(columns) - positions = {column: i for i, column in enumerate(before)} - return tuple( - tuple( - names[positions[column]] if column in positions else column - for column in after - ) - for names in columns - ) - - -def _batch_schemas( - members: Sequence[MemberContext] | Sequence[MemberQuery], - schema: TableSchema, -) -> tuple[TableSchema, ...]: - schemas = tuple(member.x.schema for member in members) - if schemas[0] == schema: - return schemas - # Apply callback column selections/reordering to each estimator's layout. - columns = _remap_columns( - columns=[s.columns[Stype.numerical] for s in schemas], - before=schemas[0].columns[Stype.numerical], - after=schema.columns[Stype.numerical], - ) - return tuple( - TableSchema(columns={**schema.columns, Stype.numerical: names}) - for names in columns - ) - - def _output_columns( contexts: Sequence[MemberContext], - y: TableTensor, ) -> tuple[tuple[str, ...], ...] | None: - if y.categorical.size(-1) == 0: - return None - names = tuple(str(value) for value in y.categorical.categories[0].tolist()) if contexts[0].y.categorical.size(-1) == 0: - return (names,) * len(contexts) - columns = tuple( + return None + return tuple( tuple( str(value) for value in context.y.categorical.categories[0].tolist() ) for context in contexts ) - if len(names) == len(columns[0]) and set(names) != set(columns[0]): - renamed = dict(zip(columns[0], names, strict=True)) - return tuple( - tuple(renamed[name] for name in group) for group in columns - ) - return _remap_columns(columns=columns, before=columns[0], after=names) def _can_batch_cache(cache: Cache) -> bool: diff --git a/test/models/test_base.py b/test/models/test_base.py index 5d6ae09f6..cd58beee2 100644 --- a/test/models/test_base.py +++ b/test/models/test_base.py @@ -548,17 +548,30 @@ 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.fit( + x=x, + y=y, + estimator_batch_size=estimator_batch_size, + callbacks=callbacks, + ) model.fit(x, y, estimator_batch_size=estimator_batch_size) - out = model.predict( - x, - callbacks=(MyCallback("affine", 2.0, 3.0, []),), - ) + 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=(MyCallback("affine", 2.0, 3.0, []),), + callbacks=callbacks, ) torch.testing.assert_close(out.numerical, 2.0 * x + 3.0) From c6e0bab7888894dfc83bbf0e1909dc19b8352d85 Mon Sep 17 00:00:00 2001 From: Akihiro Nitta Date: Wed, 23 Sep 2026 06:45:39 +0000 Subject: [PATCH 08/11] update --- benchmark/tabular/model.py | 6 +----- sdm/models/base.py | 8 ++++++++ test/models/tabiclv2/test_model.py | 11 ++++++++--- 3 files changed, 17 insertions(+), 8 deletions(-) diff --git a/benchmark/tabular/model.py b/benchmark/tabular/model.py index dfccc0743..730271075 100644 --- a/benchmark/tabular/model.py +++ b/benchmark/tabular/model.py @@ -7,7 +7,7 @@ import copy import math from dataclasses import dataclass -from typing import Any, ClassVar, Literal, cast +from typing import Any, ClassVar, Literal import numpy as np import pandas as pd @@ -21,8 +21,6 @@ import sdm import sdm.processing as sp -from sdm.cache import Cache -from sdm.models.base import _can_batch_cache from sdm.processing.execution import RecipeExecution Task = Literal["classification", "regression"] @@ -272,8 +270,6 @@ def _estimator_batch_size(self, x: torch.Tensor) -> int | None: num_rows = max(num_rows, x.size(-2)) if num_rows > 2_000 or num_rows * num_cols >= 50_000: return 1 - if not _can_batch_cache(cast(Cache, self.model._cache)): - return 1 return self._num_estimators def get_device(self) -> str: diff --git a/sdm/models/base.py b/sdm/models/base.py index 846054168..1334f2655 100644 --- a/sdm/models/base.py +++ b/sdm/models/base.py @@ -389,6 +389,8 @@ def predict( 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. @@ -432,6 +434,12 @@ def predict( ) 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) diff --git a/test/models/tabiclv2/test_model.py b/test/models/tabiclv2/test_model.py index 161e0a255..c8e645faf 100644 --- a/test/models/tabiclv2/test_model.py +++ b/test/models/tabiclv2/test_model.py @@ -235,8 +235,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) @@ -250,7 +252,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 @@ -260,8 +262,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 From 12e200fd655676954f17336a70cdb484bbf37574 Mon Sep 17 00:00:00 2001 From: Akihiro Nitta Date: Fri, 25 Sep 2026 05:04:05 +0000 Subject: [PATCH 09/11] Fix benchmark typing import --- benchmark/tabular/model.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/benchmark/tabular/model.py b/benchmark/tabular/model.py index 730271075..0a02ffa86 100644 --- a/benchmark/tabular/model.py +++ b/benchmark/tabular/model.py @@ -7,7 +7,7 @@ import copy import math from dataclasses import dataclass -from typing import Any, ClassVar, Literal +from typing import Any, ClassVar, Literal, cast import numpy as np import pandas as pd From 0ad306080bfac2610c40f2bfe5b13845e4ccd4f0 Mon Sep 17 00:00:00 2001 From: Akihiro Nitta Date: Fri, 25 Sep 2026 05:50:37 +0000 Subject: [PATCH 10/11] Batch uncached estimator inference --- benchmark/tabular/model.py | 39 ++++--- sdm/models/base.py | 154 ++++++++++++++++++------- sdm/models/kumo/tabular/model.py | 2 +- sdm/models/tabfm/model.py | 2 +- test/models/kumo/tabular/test_model.py | 38 ++++++ test/models/tabfm/test_model.py | 1 + test/models/tabiclv2/test_model.py | 26 ++++- test/models/test_base.py | 67 ++++++++++- 8 files changed, 259 insertions(+), 70 deletions(-) diff --git a/benchmark/tabular/model.py b/benchmark/tabular/model.py index 0a02ffa86..5cacfd1f3 100644 --- a/benchmark/tabular/model.py +++ b/benchmark/tabular/model.py @@ -221,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( @@ -267,7 +263,12 @@ def _estimator_batch_size(self, x: torch.Tensor) -> int | None: return 1 num_rows, num_cols = self._context_shape - num_rows = max(num_rows, x.size(-2)) + # Uncached inference processes context and query rows together. + num_rows = ( + max(num_rows, x.size(-2)) + if params["kv_cache"] + else num_rows + x.size(-2) + ) if num_rows > 2_000 or num_rows * num_cols >= 50_000: return 1 return self._num_estimators diff --git a/sdm/models/base.py b/sdm/models/base.py index 1334f2655..9b2e119f1 100644 --- a/sdm/models/base.py +++ b/sdm/models/base.py @@ -112,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, @@ -134,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. @@ -144,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) @@ -182,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: @@ -610,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, diff --git a/sdm/models/kumo/tabular/model.py b/sdm/models/kumo/tabular/model.py index 87f17d4a2..8dda7d8bf 100644 --- a/sdm/models/kumo/tabular/model.py +++ b/sdm/models/kumo/tabular/model.py @@ -222,7 +222,7 @@ def _forward( categorical_mask = _categorical_mask( x=x_context, schema=schema, - schemas=(x_context.schema,) + schemas=kwargs.get("_x_schemas", (x_context.schema,)) if cache is None else cast(tuple[TableSchema, ...], cache["x_schemas"]), ) diff --git a/sdm/models/tabfm/model.py b/sdm/models/tabfm/model.py index 783ffb942..5477d4380 100644 --- a/sdm/models/tabfm/model.py +++ b/sdm/models/tabfm/model.py @@ -210,7 +210,7 @@ def _forward( categorical_mask = _categorical_mask( x=x_context, schema=schema, - schemas=(x_context.schema,) + schemas=kwargs.get("_x_schemas", (x_context.schema,)) if cache is None else cast(tuple[TableSchema, ...], cache["x_schemas"]), ) diff --git a/test/models/kumo/tabular/test_model.py b/test/models/kumo/tabular/test_model.py index 4f013e244..ab558ea45 100644 --- a/test/models/kumo/tabular/test_model.py +++ b/test/models/kumo/tabular/test_model.py @@ -128,6 +128,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 d92048342..d4930ae98 100644 --- a/test/models/tabfm/test_model.py +++ b/test/models/tabfm/test_model.py @@ -70,6 +70,7 @@ def test_forward( 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 diff --git a/test/models/tabiclv2/test_model.py b/test/models/tabiclv2/test_model.py index c8e645faf..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,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) diff --git a/test/models/test_base.py b/test/models/test_base.py index cd58beee2..0e55fe93b 100644 --- a/test/models/test_base.py +++ b/test/models/test_base.py @@ -312,18 +312,33 @@ def test_callback() -> None: @pytest.mark.parametrize("num_estimators", [1, 3]) -def test_train_mode_enables_grad(num_estimators: int) -> None: +@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, num_estimators=num_estimators) + 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, num_estimators=num_estimators) + 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) @@ -550,6 +565,14 @@ def test_estimator_callbacks(estimator_batch_size: int | None) -> None: 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, @@ -594,6 +617,8 @@ def test_estimator_batching_incompatible_shapes() -> None: ) 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( @@ -620,6 +645,42 @@ def test_estimator_batching_incompatible_categories() -> None: 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) From 66b3b19db6e5c307f3725a55fdf815ef8bafc2c7 Mon Sep 17 00:00:00 2001 From: Akihiro Nitta Date: Fri, 25 Sep 2026 00:46:18 -0700 Subject: [PATCH 11/11] Tune uncached estimator batching thresholds --- benchmark/tabular/model.py | 14 ++++++++------ 1 file changed, 8 insertions(+), 6 deletions(-) diff --git a/benchmark/tabular/model.py b/benchmark/tabular/model.py index 5cacfd1f3..607a0e852 100644 --- a/benchmark/tabular/model.py +++ b/benchmark/tabular/model.py @@ -263,12 +263,14 @@ def _estimator_batch_size(self, x: torch.Tensor) -> int | None: return 1 num_rows, num_cols = self._context_shape - # Uncached inference processes context and query rows together. - num_rows = ( - max(num_rows, x.size(-2)) - if params["kv_cache"] - else num_rows + x.size(-2) - ) + 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