diff --git a/sdm/models/base.py b/sdm/models/base.py index 83d9bc3f5..9903ab4eb 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 @@ -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,6 +134,9 @@ def forward( 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: ``(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. @@ -184,6 +189,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, @@ -212,6 +219,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, @@ -240,6 +249,9 @@ def fit( 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: ``(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. @@ -273,11 +285,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 +400,13 @@ def predict( self._cache["recipe_execution"], ) num_batches = cast(int, self._cache["num_batches"]) + 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,6 +441,7 @@ def predict( assert cache is not None x_schemas = cast(tuple[TableSchema, ...], cache["x_schemas"]) + classes = cast(Tensor | None, cache["classes"]) batch_queries = [ self._prepare_query( query=query, @@ -443,20 +467,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=batch_queries, + 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,6 +620,8 @@ 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, @@ -583,6 +634,9 @@ def _forward_members( 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 +646,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,6 +676,8 @@ def _forward_members( queries=queries, class_values=class_values, estimator_batch_size=estimator_batch_size, + max_cost=estimator_max_cost, + cost=estimator_cost, ): outs += self._forward_batch( contexts=contexts[batch], @@ -849,10 +907,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 +922,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 +946,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..db7746418 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 @@ -890,3 +891,107 @@ 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_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 diff --git a/test/models/test_ecoc.py b/test/models/test_ecoc.py index e812c6cd9..675590b6a 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"), + [(10, 1), (11, 8)], +) +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 + )