Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
113 changes: 63 additions & 50 deletions benchmark/tabular/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
import abc
import copy
import math
from collections.abc import Callable
from dataclasses import dataclass
from typing import Any, ClassVar, Literal, cast

Expand Down Expand Up @@ -44,6 +45,7 @@ class SDMModel(AbstractTorchModel, abc.ABC):
default_num_estimators: ClassVar[int]
autocast_dtype: ClassVar[torch.dtype]
low_cardinality: ClassVar[Literal["off", "infer"]] = "off"
_estimator_batch_budget: ClassVar[int] = 2**20

@staticmethod
@abc.abstractmethod
Expand All @@ -63,6 +65,27 @@ def _set_default_params(self) -> None:
self._set_default_param_value("kv_cache", False)
self._set_default_param_value("estimator_batch_size", "auto")

def _estimator_cost(
self,
x: sdm.TableTensor,
num_classes: int,
) -> int:
del num_classes
rows = math.prod(x.size()[:-1])
return rows * (x.size(-1) + 32)

def _estimator_batching(
self,
) -> tuple[
int | None,
Callable[[sdm.TableTensor, int], int] | None,
int | None,
]:
estimator_batch_size = self._get_model_params()["estimator_batch_size"]
if estimator_batch_size == "auto":
return None, self._estimator_cost, self._estimator_batch_budget
return cast(int | None, estimator_batch_size), None, None

def _fit(
self,
X: pd.DataFrame,
Expand Down Expand Up @@ -130,7 +153,6 @@ def _fit(
y_context = y_context[perm].unflatten(0, shape)
num_estimators = None
self._expand_query = num_estimators is None
self._context_shape = x_context.shape[-2:]

recipe = self.model.default_recipe()
if params["max_columns"] is not None:
Expand All @@ -140,6 +162,9 @@ def _fit(

self._recipe_execution: RecipeExecution | None = None
if params["kv_cache"]:
estimator_batch_size, estimator_cost, estimator_max_cost = (
self._estimator_batching()
)
with torch.amp.autocast(
self._device.type,
self.autocast_dtype,
Expand All @@ -150,8 +175,10 @@ def _fit(
y=y_context,
recipe=recipe,
num_estimators=num_estimators,
estimator_batch_size=self._estimator_batch_size(x_context),
estimator_batch_size=estimator_batch_size,
generator=generator,
_estimator_cost=estimator_cost,
_estimator_max_cost=estimator_max_cost,
)
return

Expand Down Expand Up @@ -218,32 +245,29 @@ def _predict_proba(
generator = torch.Generator(self._device).set_state(
self._rng_state
)
outputs: list[sdm.TableTensor] = []
size = self._estimator_batch_size(x_query)
size = len(self._contexts) if size is None else size
for start in range(0, len(self._contexts), size):
batch = [
context._replace(
x=cast(
sdm.TableTensor, context.x.to(self._device)
),
y=cast(
sdm.TableTensor, context.y.to(self._device)
),
)
for context in self._contexts[start : start + size]
]
with torch.amp.autocast(
self._device.type,
self.autocast_dtype,
enabled=x_query.is_cuda,
):
outputs += self.model._forward_members(
contexts=batch,
queries=queries[start : start + size],
estimator_batch_size=None,
generator=generator,
)
estimator_batch_size, estimator_cost, estimator_max_cost = (
self._estimator_batching()
)
contexts = [
context._replace(
x=cast(sdm.TableTensor, context.x.to(self._device)),
y=cast(sdm.TableTensor, context.y.to(self._device)),
)
for context in self._contexts
]
with torch.amp.autocast(
self._device.type,
self.autocast_dtype,
enabled=x_query.is_cuda,
):
outputs = self.model._forward_members(
contexts=contexts,
queries=queries,
estimator_batch_size=estimator_batch_size,
estimator_cost=estimator_cost,
estimator_max_cost=estimator_max_cost,
generator=generator,
)

if self.problem_type == REGRESSION:
outputs = list(
Expand All @@ -262,28 +286,6 @@ def _predict_proba(
probabilities = out.numerical[..., indices].float().cpu().numpy()
return self._convert_proba_to_unified_form(probabilities)

def _estimator_batch_size(self, x: torch.Tensor) -> int | None:
params = self._get_model_params()
estimator_batch_size = params["estimator_batch_size"]
if estimator_batch_size != "auto":
return estimator_batch_size
# Subsampled contexts can produce different cache shapes per estimator.
if not x.is_cuda or self._expand_query:
return 1

num_rows, num_cols = self._context_shape
if not params["kv_cache"]:
# Uncached inference processes context and query rows together.
num_rows += x.size(-2)
if num_rows > 3_000 or num_rows * num_cols > 50_000:
return 1
return self._num_estimators

num_rows = max(num_rows, x.size(-2))
if num_rows > 2_000 or num_rows * num_cols >= 50_000:
return 1
return self._num_estimators

def get_device(self) -> str:
return str(next(self.model.parameters()).device)

Expand Down Expand Up @@ -342,6 +344,17 @@ class SDMKumoTabularModel(SDMModel):
key=("task", "size"),
)

def _estimator_cost(
self,
x: sdm.TableTensor,
num_classes: int,
) -> int:
cost = super()._estimator_cost(x, num_classes)
if num_classes == 0:
return cost
model = cast(sdm.models.KumoTabular, self.model)
return cost * model.ecoc.num_tasks(num_classes)

@classmethod
def warmup(
cls,
Expand Down
96 changes: 96 additions & 0 deletions test/benchmark/tabular/test_model.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,96 @@
# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

from types import SimpleNamespace

import pytest
import torch

pytest.importorskip("autogluon")
pytest.importorskip("tabarena")

from benchmark.tabular.model import (
SDMKumoTabularSmallModel,
SDMTabICLv2Model,
)
from sdm import CategoricalTensor, Stype, TableTensor
from sdm.models import ECOC
from sdm.models.base import _batch_slices, _query_chunks
from sdm.processing.execution import MemberContext, MemberQuery


def test_estimator_cost() -> None:
adapter = object.__new__(SDMTabICLv2Model)
x = TableTensor.from_tensor(torch.empty(2, 5, 7))

assert adapter._estimator_cost(x, num_classes=0) == 10 * (7 + 32)
assert adapter._estimator_cost(x, num_classes=5) == 10 * (7 + 32)


def test_auto_estimator_batching(
monkeypatch: pytest.MonkeyPatch,
) -> None:
adapter = object.__new__(SDMTabICLv2Model)
monkeypatch.setattr(
SDMTabICLv2Model,
"_get_model_params",
lambda _: {"estimator_batch_size": "auto"},
)

size, cost, max_cost = adapter._estimator_batching()

assert size is None
assert cost == adapter._estimator_cost
assert max_cost == 2**20


def test_kumo_ecoc_cost_controls_batches_and_query_chunks() -> None:
adapter = object.__new__(SDMKumoTabularSmallModel)
adapter.model = SimpleNamespace(ecoc=ECOC(max_classes=10))
x = TableTensor.from_tensor(torch.empty(500, 4))

base_cost = 500 * (4 + 32)
assert adapter._estimator_cost(x, num_classes=0) == base_cost
assert adapter._estimator_cost(x, num_classes=5) == base_cost
assert adapter._estimator_cost(x, num_classes=12) == base_cost * 8

y = TableTensor(
columns={Stype.categorical: ("target",)},
categorical=CategoricalTensor(
code=torch.zeros(500, 1, dtype=torch.long),
categories=(torch.arange(12),),
),
)
context = MemberContext(
x=x,
y=y,
related_tables=None,
input_stypes={},
)
batches = _batch_slices(
contexts=[context] * 8,
queries=None,
class_values=[tuple(range(12))] * 8,
estimator_batch_size=None,
max_cost=adapter._estimator_batch_budget,
cost=adapter._estimator_cost,
)
assert [(batch.start, batch.stop) for batch in batches] == [
(0, 7),
(7, 8),
]

query = MemberQuery(
x=TableTensor.from_tensor(torch.empty(1_000, 4)),
related_tables=None,
)
chunks = _query_chunks(
queries=[query] * 7,
max_cost=adapter._estimator_batch_budget,
cost=adapter._estimator_cost,
num_classes=12,
)
assert [[member.x.size(-2) for member in chunk] for chunk in chunks] == [
[520] * 7,
[480] * 7,
]
Loading