From c8fc4f71244d0828cd0b8c0509214048f1a68169 Mon Sep 17 00:00:00 2001 From: RBendias Date: Mon, 28 Sep 2026 14:41:49 +0000 Subject: [PATCH 1/3] Add automatic estimator batching Signed-off-by: RBendias --- sdm/models/base.py | 181 ++++++++++++++++++++++++++++++++------- sdm/models/ecoc.py | 16 ++-- test/models/test_base.py | 148 ++++++++++++++++++++++++++++++++ test/models/test_ecoc.py | 28 ++++++ 4 files changed, 338 insertions(+), 35 deletions(-) diff --git a/sdm/models/base.py b/sdm/models/base.py index 83d9bc3f5..812749831 100644 --- a/sdm/models/base.py +++ b/sdm/models/base.py @@ -3,7 +3,7 @@ import abc import copy -from collections.abc import Hashable, Iterable, Mapping, Sequence +from collections.abc import Callable, Hashable, Iterable, Mapping, Sequence from typing import Any, ClassVar, cast import torch @@ -128,11 +128,15 @@ def forward( per estimator). estimator_batch_size: Maximum number of consecutive estimators run through the model in one call. ``1`` (default) runs estimators - one by one; ``None`` batches as many as possible. Estimators - whose preprocessed tables differ in shape or target class + one by one, which minimizes device memory; ``None`` batches as + many as possible. Estimators whose + preprocessed tables differ in shape or target class set, or that come with related tables, run in separate - calls. Device memory grows with the batch size. - callbacks: Callbacks applied in sequence to this model call. + calls. Batched and sequential predictions are equal up to + floating-point rounding. + callbacks: Callbacks applied in sequence to this model call. The + preprocessing hooks of all estimators run before their model + forward hooks. generator: Pseudorandom number generator used for sampling during pre-processing and model execution. kwargs: Additional keyword arguments passed to the model. @@ -142,6 +146,14 @@ def forward( stacked estimator outputs with shape ``[E, ..., R_query, *]``. """ callbacks = () if callbacks is None else callbacks + estimator_cost = cast( + Callable[[TableTensor, int], int] | None, + kwargs.pop("_estimator_cost", None), + ) + estimator_max_cost = cast( + int | None, + kwargs.pop("_estimator_max_cost", None), + ) requires_grad = self.training requires_grad |= any(callback.requires_grad for callback in callbacks) @@ -184,6 +196,8 @@ def forward( contexts=contexts, queries=queries, estimator_batch_size=estimator_batch_size, + estimator_cost=estimator_cost, + estimator_max_cost=estimator_max_cost, callbacks=callbacks, generator=generator, **kwargs, @@ -235,17 +249,27 @@ def fit( per estimator). estimator_batch_size: Maximum number of consecutive estimators run through the model in one call. ``1`` (default) runs estimators - one by one; ``None`` batches as many as possible. Estimators - whose preprocessed tables differ in shape or target class + one by one, which minimizes device memory; ``None`` batches as + many as possible. Estimators whose + preprocessed tables differ in shape or target class set, or that come with related tables, run in separate - calls. Device memory grows with the batch size. Estimators - fitted together are predicted together. + calls. Estimators fitted together are predicted together. + Batched and sequential predictions are equal up to + floating-point rounding. 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 + estimator_cost = cast( + Callable[[TableTensor, int], int] | None, + kwargs.pop("_estimator_cost", None), + ) + estimator_max_cost = cast( + int | None, + kwargs.pop("_estimator_max_cost", None), + ) self.clear() @@ -273,11 +297,15 @@ def fit( queries=None, class_values=class_values, estimator_batch_size=estimator_batch_size, + max_cost=estimator_max_cost, + cost=estimator_cost, ) cache = Cache( recipe_execution=recipe_execution, kwargs=kwargs, num_batches=len(batches), + estimator_cost=estimator_cost, + estimator_max_cost=estimator_max_cost, ) for i, batch in enumerate(batches): with inference_mode("no_grad"): @@ -384,6 +412,14 @@ def predict( self._cache["recipe_execution"], ) num_batches = cast(int, self._cache["num_batches"]) + # Callbacks see every estimator output once, so queries stay whole. + estimator_max_cost = ( + None if callbacks else self._cache["estimator_max_cost"] + ) + estimator_cost = cast( + Callable[[TableTensor, int], int] | None, + self._cache["estimator_cost"], + ) caches = [cast(Cache, self._cache[i]) for i in range(num_batches)] next_cache = caches[0] @@ -418,7 +454,8 @@ def predict( assert cache is not None x_schemas = cast(tuple[TableSchema, ...], cache["x_schemas"]) - batch_queries = [ + classes = cast(Tensor | None, cache["classes"]) + members = [ self._prepare_query( query=query, x_schema=x_schema, @@ -443,20 +480,45 @@ def predict( with torch.cuda.stream(transfer_stream): next_cache = next_cache.to(x.device, non_blocking=True) - outs += self._forward_batch( - contexts=None, - queries=batch_queries, - cache=cache, - categorical_mask=cast(Tensor, cache["categorical_mask"]), - class_values=cast( - Sequence[tuple[Any, ...] | None] | None, - cache["class_values"], - ), - callbacks=callbacks, - requires_grad=requires_grad, - generator=None, - **cast(dict[str, Any], self._cache["kwargs"]), - ) + chunks = [ + self._forward_batch( + contexts=None, + queries=chunk, + cache=cache, + categorical_mask=cast( + Tensor, cache["categorical_mask"] + ), + class_values=cast( + Sequence[tuple[Any, ...] | None] | None, + cache["class_values"], + ), + callbacks=callbacks, + requires_grad=requires_grad, + generator=None, + **cast(dict[str, Any], self._cache["kwargs"]), + ) + for chunk in _query_chunks( + queries=members, + max_cost=cast(int | None, estimator_max_cost), + cost=estimator_cost, + num_classes=( + 0 if classes is None else classes.numel() + ), + ) + ] + + if len(chunks) == 1: + batch_outputs = chunks[0] + else: + batch_outputs = [ + cast(TableTensor, torch.cat(parts, dim=-2)) + for parts in zip(*chunks, strict=True) + ] + + # Drop source chunk references before the next estimator + # batch. A single chunk stays alive through `batch_outputs`. + del chunks + outs.extend(batch_outputs) if x.is_cuda: assert compute_stream is not None @@ -571,18 +633,25 @@ def _forward_members( queries: Sequence[MemberQuery], *, estimator_batch_size: int | None = 1, + estimator_cost: Callable[[TableTensor, int], int] | None = None, + estimator_max_cost: int | None = None, callbacks: Sequence[Callback] | None = None, generator: torch.Generator | None = None, **kwargs: Any, ) -> list[TableTensor]: - r"""Run recipe-transformed members that live on the model device. + r"""Run recipe-transformed members on the device of their queries. + Context features and targets may live on another device; each batch + is moved when it runs. Returns one output per member before target inversion and ``recipe.output``. """ callbacks = () if callbacks is None else callbacks requires_grad = self.training requires_grad |= any(callback.requires_grad for callback in callbacks) + if requires_grad and estimator_max_cost is not None: + estimator_batch_size = 1 + estimator_max_cost = None if estimator_batch_size == 1 and len(contexts) > 1: outs: list[TableTensor] = [] @@ -592,6 +661,8 @@ def _forward_members( contexts=(context,), queries=(query,), estimator_batch_size=1, + estimator_cost=estimator_cost, + estimator_max_cost=estimator_max_cost, callbacks=callbacks, generator=generator, **kwargs, @@ -620,9 +691,18 @@ def _forward_members( queries=queries, class_values=class_values, estimator_batch_size=estimator_batch_size, + max_cost=estimator_max_cost, + cost=estimator_cost, ): + device = queries[batch.start].x.device outs += self._forward_batch( - contexts=contexts[batch], + contexts=[ + context._replace( + x=cast(TableTensor, context.x.to(device)), + y=cast(TableTensor, context.y.to(device)), + ) + for context in contexts[batch] + ], queries=queries[batch], cache=None, categorical_mask=None, @@ -849,10 +929,12 @@ def _batch_slices( queries: Sequence[MemberQuery] | None, class_values: Sequence[tuple[Any, ...] | None], estimator_batch_size: int | None, + max_cost: int | None, + cost: Callable[[TableTensor, int], int] | None, ) -> list[slice]: # Consecutive estimators that can go through one `_forward` together, - # split when shapes/dtypes or the target class set change, or the - # batch is full. + # split when shapes/dtypes or the target class set change, the batch is + # full, or the optional cost budget would be exceeded. related = any(context.related_tables is not None for context in contexts) if queries is not None: related = related or any( @@ -862,7 +944,7 @@ def _batch_slices( return [slice(i, i + 1) for i in range(len(contexts))] batches: list[slice] = [] - start = 0 + start = total = 0 key: Hashable = None # Stack when table shapes/dtypes match and the target class set matches. for i, context in enumerate(contexts): @@ -886,20 +968,59 @@ def _batch_slices( ), None if classes is None else frozenset(classes), ) + member_cost = 0 + if max_cost is not None: + assert cost is not None + num_classes = ( + context.y.categorical.categories[0].numel() + if context.y.categorical.size(-1) > 0 + else 0 + ) + member_cost = cost(context.x, num_classes) + if query is not None: + member_cost += cost(query.x, num_classes) if i > start and ( member_key != key or ( estimator_batch_size is not None and i - start == estimator_batch_size ) + or (max_cost is not None and total + member_cost > max_cost) ): batches.append(slice(start, i)) - start = i + start, total = i, 0 key = member_key + total += member_cost batches.append(slice(start, len(contexts))) return batches +def _query_chunks( + queries: Sequence[MemberQuery], + max_cost: int | None, + cost: Callable[[TableTensor, int], int] | None, + num_classes: int, +) -> list[list[MemberQuery]]: + # Row chunks of a batch's queries within the cost budget, but never smaller + # than the queries of one estimator on their own. + if max_cost is None or len(queries) == 1: + return [list(queries)] + assert cost is not None + total = sum(cost(query.x, num_classes) for query in queries) + if total <= max_cost: + return [list(queries)] + rows = queries[0].x.size(-2) + budget = max(max_cost, total // len(queries)) + splits = [ + query.x.split(max(1, budget * rows // total), dim=-2) + for query in queries + ] + return [ + [MemberQuery(x=x, related_tables=None) for x in xs] + for xs in zip(*splits, strict=True) + ] + + def _categorical_mask(members: Sequence[MemberContext]) -> Tensor: x = members[0].x mask = torch.tensor( diff --git a/sdm/models/ecoc.py b/sdm/models/ecoc.py index 901b61835..daf6df26e 100644 --- a/sdm/models/ecoc.py +++ b/sdm/models/ecoc.py @@ -138,6 +138,16 @@ def forward( active = index != self.max_classes - 1 return scores.masked_fill(~active, 0).sum(dim=0) / active.sum(dim=0) + def num_tasks(self, num_classes: int) -> int: + """Return the number of tasks the model runs for ``num_classes``.""" + if num_classes <= self.max_classes: + return 1 + return max( + # Give every class its own output in at least one task. + math.ceil(num_classes / (self.max_classes - 1)), + 4 * math.ceil(math.log(num_classes, self.max_classes)), + ) + def _draw_codebook( self, num_classes: int, @@ -145,11 +155,7 @@ def _draw_codebook( generator: torch.Generator | None, ) -> Tensor: rest_idx = self.max_classes - 1 - num_codes = max( - # Give every class its own output in at least one task. - math.ceil(num_classes / rest_idx), - 4 * math.ceil(math.log(num_classes, self.max_classes)), - ) + num_codes = self.num_tasks(num_classes) # Bound the quadratic distance search for large targets. num_draws = 50 if num_classes <= 200 else 1 codebook = torch.full( diff --git a/test/models/test_base.py b/test/models/test_base.py index 96c796730..e08121930 100644 --- a/test/models/test_base.py +++ b/test/models/test_base.py @@ -1,6 +1,7 @@ # SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 +import math from dataclasses import dataclass from typing import Any, ClassVar, cast @@ -20,6 +21,7 @@ from sdm.models import ICLModel from sdm.models.callback import Callback from sdm.processing import InvertibleMixin, Processor +from sdm.processing.execution import RecipeExecution @dataclass @@ -890,3 +892,149 @@ def table(i: int) -> TableTensor: ) assert len(model.calls) == 2 + + +def _estimator_cost(x: TableTensor, num_classes: int) -> int: + del num_classes + return math.prod(x.size()[:-1]) * (x.size(-1) + 1) + + +def test_estimator_batching_keeps_cost_budget() -> None: + # Every estimator costs 3 context and 2 query rows of 2 columns plus 1. + model = _RecordingModel() + x = torch.randn(4, 3, 2) + y = torch.zeros(4, 3, 1) + x_query = torch.randn(4, 2, 2) + + out = model( + x, + y, + x_query, + estimator_batch_size=None, + _estimator_cost=_estimator_cost, + _estimator_max_cost=30, + ) + + torch.testing.assert_close(out.numerical, x_query) + assert [ + cast(TableTensor, call.x_query).size() for call in model.calls + ] == [(2, 2, 2), (2, 2, 2)] + + +def test_estimator_batching_runs_over_budget_member_alone() -> None: + model = _RecordingModel() + x_query = torch.randn(2, 2, 2) + + out = model( + torch.randn(2, 3, 2), + torch.zeros(2, 3, 1), + x_query, + estimator_batch_size=None, + _estimator_cost=_estimator_cost, + _estimator_max_cost=14, + ) + + torch.testing.assert_close(out.numerical, x_query) + assert len(model.calls) == 2 + + +def test_estimator_batching_counts_rows_without_columns() -> None: + model = _RecordingModel() + + model( + torch.randn(2, 3, 0), + torch.zeros(2, 3, 1), + torch.randn(2, 2, 0), + estimator_batch_size=None, + _estimator_cost=_estimator_cost, + _estimator_max_cost=9, + ) + + assert len(model.calls) == 2 + + +def test_estimator_cost_batching_is_sequential_with_gradients() -> None: + model = _RecordingModel() + model.train() + + model( + torch.randn(3, 3, 2), + torch.zeros(3, 3, 1), + torch.randn(3, 2, 2), + estimator_batch_size=None, + _estimator_cost=_estimator_cost, + _estimator_max_cost=2**20, + ) + + assert len(model.calls) == 3 + + +def test_predict_splits_batched_query_rows_within_budget() -> None: + # Two estimators with 3 context rows of 2 columns plus 1 each fit one + # batch, but their 5 query rows exceed the budget together. + model = _RecordingModel() + model.fit( + torch.randn(2, 3, 2), + torch.zeros(2, 3, 1), + estimator_batch_size=None, + _estimator_cost=_estimator_cost, + _estimator_max_cost=18, + ) + model.calls.clear() + x_query = torch.randn(2, 5, 2) + + out = model.predict(x_query) + + torch.testing.assert_close(out.numerical, x_query) + assert [ + cast(TableTensor, call.x_query).size() for call in model.calls + ] == [(2, 3, 2), (2, 2, 2)] + + +def test_predict_keeps_queries_whole_with_callbacks() -> None: + # Without callbacks, this budget splits the query rows into two calls. + model = _RecordingModel() + model.fit( + torch.randn(2, 3, 2), + torch.zeros(2, 3, 1), + estimator_batch_size=None, + _estimator_cost=_estimator_cost, + _estimator_max_cost=18, + ) + model.calls.clear() + events: list[str] = [] + + model.predict( + torch.randn(2, 5, 2), + callbacks=(MyCallback("affine", 2.0, 3.0, events),), + ) + + assert len(model.calls) == 1 + assert events.count("affine_model_forward_end") == 2 + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_forward_members_moves_contexts_to_query_device() -> None: + model = _RecordingModel() + recipe = RecipeExecution(model.default_recipe()) + contexts = recipe.fit_transform( + x=torch.randn(4, 3, 2), + y=torch.zeros(4, 3, 1), + related_tables=None, + num_members=None, + ) + x_query = torch.randn(4, 2, 2) + queries = [ + query._replace(x=cast(TableTensor, query.x.cuda())) + for query in recipe.transform(x=x_query, related_tables=None) + ] + + outs = model._forward_members(contexts=contexts, queries=queries) + + assert all( + cast(TableTensor, call.x_context).is_cuda for call in model.calls + ) + torch.testing.assert_close( + torch.stack([out.numerical.cpu() for out in outs]), + x_query, + ) diff --git a/test/models/test_ecoc.py b/test/models/test_ecoc.py index e812c6cd9..b55f202c3 100644 --- a/test/models/test_ecoc.py +++ b/test/models/test_ecoc.py @@ -131,3 +131,31 @@ def test_ecoc_members(device: torch.device) -> None: cache.freeze() replayed = forward(x=x[..., num_context:, :], y=y[..., :0], cache=cache) torch.testing.assert_close(replayed, scores) + + +@pytest.mark.parametrize( + ("num_classes", "expected_num_tasks"), + [(3, 1), (10, 1), (11, 8), (100, 12), (201, 23)], +) +def test_ecoc_num_tasks( + num_classes: int, + expected_num_tasks: int, +) -> None: + ecoc = ECOC(max_classes=10) + x = torch.randn(num_classes + 2, num_classes, dtype=torch.float64) + y = torch.arange(num_classes) + cache = Cache() + + ecoc( + model=MyModel(10), + x=x[:num_classes], + y=y, + num_classes=num_classes, + cache=cache, + ) + + assert ecoc.num_tasks(num_classes) == expected_num_tasks + if num_classes > 10: + assert ( + cast(Tensor, cache["ecoc_codebook"]).size(-2) == expected_num_tasks + ) From f8af851027ff2dc2e2840ca4b4acbc8a9b996ef4 Mon Sep 17 00:00:00 2001 From: RBendias Date: Tue, 29 Sep 2026 12:56:07 +0000 Subject: [PATCH 2/3] update --- sdm/models/base.py | 47 +++++++++++++++++++--------------------- test/models/test_base.py | 24 ++++++++++---------- 2 files changed, 34 insertions(+), 37 deletions(-) diff --git a/sdm/models/base.py b/sdm/models/base.py index 812749831..bdf33bfba 100644 --- a/sdm/models/base.py +++ b/sdm/models/base.py @@ -104,6 +104,8 @@ def forward( recipe: Recipe | None = None, num_estimators: int | None = None, estimator_batch_size: int | None = 1, + estimator_cost: Callable[[TableTensor, int], int] | None = None, + estimator_max_cost: int | None = None, callbacks: Sequence[Callback] | None = None, generator: torch.Generator | None = None, **kwargs: Any, @@ -132,11 +134,14 @@ def forward( many as possible. Estimators whose preprocessed tables differ in shape or target class set, or that come with related tables, run in separate - calls. Batched and sequential predictions are equal up to - floating-point rounding. - callbacks: Callbacks applied in sequence to this model call. The - preprocessing hooks of all estimators run before their model - forward hooks. + calls. Device memory grows with the batch size. + estimator_cost: Cost of a preprocessed table given its number of + target classes. Used with ``estimator_max_cost``. + estimator_max_cost: Maximum total ``estimator_cost`` of the + tables in one model call. Consecutive estimators that fit + the budget run together; a single estimator that exceeds it + still runs. Has no effect during gradient-based training. + 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. @@ -146,14 +151,6 @@ def forward( stacked estimator outputs with shape ``[E, ..., R_query, *]``. """ callbacks = () if callbacks is None else callbacks - estimator_cost = cast( - Callable[[TableTensor, int], int] | None, - kwargs.pop("_estimator_cost", None), - ) - estimator_max_cost = cast( - int | None, - kwargs.pop("_estimator_max_cost", None), - ) requires_grad = self.training requires_grad |= any(callback.requires_grad for callback in callbacks) @@ -226,6 +223,8 @@ def fit( recipe: Recipe | None = None, num_estimators: int | None = None, estimator_batch_size: int | None = 1, + estimator_cost: Callable[[TableTensor, int], int] | None = None, + estimator_max_cost: int | None = None, callbacks: Sequence[Callback] | None = None, generator: torch.Generator | None = None, **kwargs: Any, @@ -253,24 +252,22 @@ def fit( many as possible. Estimators whose preprocessed tables differ in shape or target class set, or that come with related tables, run in separate - calls. Estimators fitted together are predicted together. - Batched and sequential predictions are equal up to - floating-point rounding. + calls. Device memory grows with the batch size. Estimators + fitted together are predicted together. + estimator_cost: Cost of a preprocessed table given its number of + target classes. Used with ``estimator_max_cost``. + estimator_max_cost: Maximum total ``estimator_cost`` of the + tables in one model call. Consecutive estimators that fit + the budget are fitted together and later predicted together; + a single estimator that exceeds it still runs. On + :meth:`predict`, query rows of a fitted batch may split to + stay within the budget. 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 - estimator_cost = cast( - Callable[[TableTensor, int], int] | None, - kwargs.pop("_estimator_cost", None), - ) - estimator_max_cost = cast( - int | None, - kwargs.pop("_estimator_max_cost", None), - ) - self.clear() recipe_execution = RecipeExecution( diff --git a/test/models/test_base.py b/test/models/test_base.py index e08121930..9a6903ff1 100644 --- a/test/models/test_base.py +++ b/test/models/test_base.py @@ -911,8 +911,8 @@ def test_estimator_batching_keeps_cost_budget() -> None: y, x_query, estimator_batch_size=None, - _estimator_cost=_estimator_cost, - _estimator_max_cost=30, + estimator_cost=_estimator_cost, + estimator_max_cost=30, ) torch.testing.assert_close(out.numerical, x_query) @@ -930,8 +930,8 @@ def test_estimator_batching_runs_over_budget_member_alone() -> None: torch.zeros(2, 3, 1), x_query, estimator_batch_size=None, - _estimator_cost=_estimator_cost, - _estimator_max_cost=14, + estimator_cost=_estimator_cost, + estimator_max_cost=14, ) torch.testing.assert_close(out.numerical, x_query) @@ -946,8 +946,8 @@ def test_estimator_batching_counts_rows_without_columns() -> None: torch.zeros(2, 3, 1), torch.randn(2, 2, 0), estimator_batch_size=None, - _estimator_cost=_estimator_cost, - _estimator_max_cost=9, + estimator_cost=_estimator_cost, + estimator_max_cost=9, ) assert len(model.calls) == 2 @@ -962,8 +962,8 @@ def test_estimator_cost_batching_is_sequential_with_gradients() -> None: torch.zeros(3, 3, 1), torch.randn(3, 2, 2), estimator_batch_size=None, - _estimator_cost=_estimator_cost, - _estimator_max_cost=2**20, + estimator_cost=_estimator_cost, + estimator_max_cost=2**20, ) assert len(model.calls) == 3 @@ -977,8 +977,8 @@ def test_predict_splits_batched_query_rows_within_budget() -> None: torch.randn(2, 3, 2), torch.zeros(2, 3, 1), estimator_batch_size=None, - _estimator_cost=_estimator_cost, - _estimator_max_cost=18, + estimator_cost=_estimator_cost, + estimator_max_cost=18, ) model.calls.clear() x_query = torch.randn(2, 5, 2) @@ -998,8 +998,8 @@ def test_predict_keeps_queries_whole_with_callbacks() -> None: torch.randn(2, 3, 2), torch.zeros(2, 3, 1), estimator_batch_size=None, - _estimator_cost=_estimator_cost, - _estimator_max_cost=18, + estimator_cost=_estimator_cost, + estimator_max_cost=18, ) model.calls.clear() events: list[str] = [] From 9444a4ecb95765edcbc14bed295a03b7951a53a3 Mon Sep 17 00:00:00 2001 From: RBendias Date: Tue, 29 Sep 2026 13:13:02 +0000 Subject: [PATCH 3/3] update --- sdm/models/base.py | 49 ++++++++++++---------------------------- test/models/test_base.py | 43 ----------------------------------- test/models/test_ecoc.py | 2 +- 3 files changed, 16 insertions(+), 78 deletions(-) diff --git a/sdm/models/base.py b/sdm/models/base.py index bdf33bfba..9903ab4eb 100644 --- a/sdm/models/base.py +++ b/sdm/models/base.py @@ -130,17 +130,13 @@ def forward( per estimator). estimator_batch_size: Maximum number of consecutive estimators run through the model in one call. ``1`` (default) runs estimators - one by one, which minimizes device memory; ``None`` batches as - many as possible. Estimators whose - preprocessed tables differ in shape or target class + one by one; ``None`` batches as many as possible. Estimators + whose preprocessed tables differ in shape or target class set, or that come with related tables, run in separate calls. Device memory grows with the batch size. - estimator_cost: Cost of a preprocessed table given its number of - target classes. Used with ``estimator_max_cost``. - estimator_max_cost: Maximum total ``estimator_cost`` of the - tables in one model call. Consecutive estimators that fit - the budget run together; a single estimator that exceeds it - still runs. Has no effect during gradient-based training. + estimator_cost: ``(table, num_classes) -> int`` cost of one table. + estimator_max_cost: Run estimators together until this budget + would be exceeded. callbacks: Callbacks applied in sequence to this model call. generator: Pseudorandom number generator used for sampling during pre-processing and model execution. @@ -248,26 +244,21 @@ def fit( per estimator). estimator_batch_size: Maximum number of consecutive estimators run through the model in one call. ``1`` (default) runs estimators - one by one, which minimizes device memory; ``None`` batches as - many as possible. Estimators whose - preprocessed tables differ in shape or target class + one by one; ``None`` batches as many as possible. Estimators + whose preprocessed tables differ in shape or target class set, or that come with related tables, run in separate calls. Device memory grows with the batch size. Estimators fitted together are predicted together. - estimator_cost: Cost of a preprocessed table given its number of - target classes. Used with ``estimator_max_cost``. - estimator_max_cost: Maximum total ``estimator_cost`` of the - tables in one model call. Consecutive estimators that fit - the budget are fitted together and later predicted together; - a single estimator that exceeds it still runs. On - :meth:`predict`, query rows of a fitted batch may split to - stay within the budget. + estimator_cost: ``(table, num_classes) -> int`` cost of one table. + estimator_max_cost: Run estimators together until this budget + would be exceeded. 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 + self.clear() recipe_execution = RecipeExecution( @@ -409,7 +400,6 @@ def predict( self._cache["recipe_execution"], ) num_batches = cast(int, self._cache["num_batches"]) - # Callbacks see every estimator output once, so queries stay whole. estimator_max_cost = ( None if callbacks else self._cache["estimator_max_cost"] ) @@ -452,7 +442,7 @@ def predict( x_schemas = cast(tuple[TableSchema, ...], cache["x_schemas"]) classes = cast(Tensor | None, cache["classes"]) - members = [ + batch_queries = [ self._prepare_query( query=query, x_schema=x_schema, @@ -495,7 +485,7 @@ def predict( **cast(dict[str, Any], self._cache["kwargs"]), ) for chunk in _query_chunks( - queries=members, + queries=batch_queries, max_cost=cast(int | None, estimator_max_cost), cost=estimator_cost, num_classes=( @@ -636,10 +626,8 @@ def _forward_members( generator: torch.Generator | None = None, **kwargs: Any, ) -> list[TableTensor]: - r"""Run recipe-transformed members on the device of their queries. + r"""Run recipe-transformed members that live on the model device. - Context features and targets may live on another device; each batch - is moved when it runs. Returns one output per member before target inversion and ``recipe.output``. """ @@ -691,15 +679,8 @@ def _forward_members( max_cost=estimator_max_cost, cost=estimator_cost, ): - device = queries[batch.start].x.device outs += self._forward_batch( - contexts=[ - context._replace( - x=cast(TableTensor, context.x.to(device)), - y=cast(TableTensor, context.y.to(device)), - ) - for context in contexts[batch] - ], + contexts=contexts[batch], queries=queries[batch], cache=None, categorical_mask=None, diff --git a/test/models/test_base.py b/test/models/test_base.py index 9a6903ff1..db7746418 100644 --- a/test/models/test_base.py +++ b/test/models/test_base.py @@ -21,7 +21,6 @@ from sdm.models import ICLModel from sdm.models.callback import Callback from sdm.processing import InvertibleMixin, Processor -from sdm.processing.execution import RecipeExecution @dataclass @@ -938,21 +937,6 @@ def test_estimator_batching_runs_over_budget_member_alone() -> None: assert len(model.calls) == 2 -def test_estimator_batching_counts_rows_without_columns() -> None: - model = _RecordingModel() - - model( - torch.randn(2, 3, 0), - torch.zeros(2, 3, 1), - torch.randn(2, 2, 0), - estimator_batch_size=None, - estimator_cost=_estimator_cost, - estimator_max_cost=9, - ) - - assert len(model.calls) == 2 - - def test_estimator_cost_batching_is_sequential_with_gradients() -> None: model = _RecordingModel() model.train() @@ -1011,30 +995,3 @@ def test_predict_keeps_queries_whole_with_callbacks() -> None: assert len(model.calls) == 1 assert events.count("affine_model_forward_end") == 2 - - -@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") -def test_forward_members_moves_contexts_to_query_device() -> None: - model = _RecordingModel() - recipe = RecipeExecution(model.default_recipe()) - contexts = recipe.fit_transform( - x=torch.randn(4, 3, 2), - y=torch.zeros(4, 3, 1), - related_tables=None, - num_members=None, - ) - x_query = torch.randn(4, 2, 2) - queries = [ - query._replace(x=cast(TableTensor, query.x.cuda())) - for query in recipe.transform(x=x_query, related_tables=None) - ] - - outs = model._forward_members(contexts=contexts, queries=queries) - - assert all( - cast(TableTensor, call.x_context).is_cuda for call in model.calls - ) - torch.testing.assert_close( - torch.stack([out.numerical.cpu() for out in outs]), - x_query, - ) diff --git a/test/models/test_ecoc.py b/test/models/test_ecoc.py index b55f202c3..675590b6a 100644 --- a/test/models/test_ecoc.py +++ b/test/models/test_ecoc.py @@ -135,7 +135,7 @@ def test_ecoc_members(device: torch.device) -> None: @pytest.mark.parametrize( ("num_classes", "expected_num_tasks"), - [(3, 1), (10, 1), (11, 8), (100, 12), (201, 23)], + [(10, 1), (11, 8)], ) def test_ecoc_num_tasks( num_classes: int,