Skip to content
Open
11 changes: 3 additions & 8 deletions benchmark/tabular/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -214,6 +214,7 @@ def _predict_proba(
x=x_query,
related_tables=None,
)
dtype = queries[0].x.dtype
generator = torch.Generator(self._device).set_state(
self._rng_state
)
Expand All @@ -230,14 +231,8 @@ def _predict_proba(
],
generator=generator,
)

if self.problem_type == REGRESSION:
outputs = list(
self._recipe_execution.inverse_transform_target(
outputs
)
)
out = self._recipe_execution.transform_output(outputs)
del queries
out = self._recipe_execution.transform_output(outputs, dtype)

if self.problem_type == REGRESSION:
return out.numerical.float().mean(dim=-1).cpu().numpy()
Expand Down
33 changes: 33 additions & 0 deletions sdm/_memory.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,33 @@
# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

import os

import torch


def chunk_memory_limit(device: torch.device) -> int:
r"""Bytes one chunk of a chunked operation may occupy on a CUDA device.

The limit is the ``SDM_CHUNK_MEMORY_FRACTION`` (default ``0.05``) share of
the device memory available to this process.
"""
return int(
torch.cuda.get_device_properties(device).total_memory
* torch.cuda.get_per_process_memory_fraction(device)
* float(os.getenv("SDM_CHUNK_MEMORY_FRACTION", "0.05"))
)


def split_size(num_items: int, item_bytes: int, device: torch.device) -> int:
r"""Split size of balanced chunks of items of ``item_bytes`` bytes each.

On CUDA devices, chunks fit :func:`chunk_memory_limit`. Autograd keeps
the memory of every chunk, so one chunk holds all items elsewhere or
while gradients are enabled.
"""
if device.type != "cuda" or torch.is_grad_enabled():
return max(num_items, 1)
limit = max(chunk_memory_limit(device), 1)
num_chunks = max(-(-num_items * item_bytes // limit), 1)
return max(-(-num_items // num_chunks), 1)
19 changes: 12 additions & 7 deletions sdm/ensemble.py
Original file line number Diff line number Diff line change
Expand Up @@ -212,15 +212,20 @@ def _select_members(self, member_ids: Sequence[int]) -> Self:
group = self._groups[group_id]
selected_positions = tuple(positions)
group_ids[group_id] = new_group_id
start = selected_positions[0]
step = (
selected_positions[1] - start
if len(selected_positions) > 1
else 1
)
stop = start + step * len(selected_positions)
if selected_positions == tuple(range(group.size(0))):
groups.append(group)
elif len(selected_positions) == 1:
groups.append(
cast(
TableTensor,
group.narrow(0, selected_positions[0], 1),
)
)
elif step > 0 and selected_positions == tuple(
range(start, stop, step)
):
# Evenly spaced members are selected as a view, not a copy.
groups.append(group[start:stop:step])
else:
groups.append(
cast(
Expand Down
28 changes: 9 additions & 19 deletions sdm/models/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -204,19 +204,14 @@ def forward(
**kwargs,
)

# Regression: invert target before stacking estimator outputs.
if contexts[0].y.numerical.size(-1) > 0:
with (
torch.amp.autocast(x_query.device.type, enabled=False),
inference_mode("grad" if requires_grad else "inference"),
):
outs = list(recipe_execution.inverse_transform_target(outs))

with (
torch.amp.autocast(x_query.device.type, enabled=False),
inference_mode("grad" if requires_grad else "inference"),
):
return recipe_execution.transform_output(outs)
return recipe_execution.transform_output(
outs,
dtype=queries[0].x.dtype,
)

def fit(
self,
Expand Down Expand Up @@ -524,19 +519,14 @@ def predict(
transfer_stream.synchronize()
raise

# Regression: invert target before stacking estimator outputs.
if cast(Cache, self._cache[0])["classes"] is None:
with (
torch.amp.autocast(x.device.type, enabled=False),
inference_mode("grad" if requires_grad else "inference"),
):
outs = list(recipe_execution.inverse_transform_target(outs))

with (
torch.amp.autocast(x.device.type, enabled=False),
inference_mode("grad" if requires_grad else "inference"),
):
return recipe_execution.transform_output(outs)
return recipe_execution.transform_output(
outs,
dtype=queries[0].x.dtype,
)

def clear(self) -> None:
r"""Clear cached context state created by :meth:`fit`."""
Expand Down Expand Up @@ -775,7 +765,7 @@ def _forward_batch(
for i in range(len(outs)):
for callback in callbacks:
outs[i] = callback.on_model_forward_end(self, outs[i])
return [cast(TableTensor, out.to(query.x.dtype)) for out in outs]
return outs

def _validate_context(
self,
Expand Down
11 changes: 7 additions & 4 deletions sdm/models/kumo/tabular/icl.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,6 +74,7 @@ def forward(
y_emb = self.y_lin(y.unsqueeze(-1)) # [..., R_train, D]

x[..., :R_train, :] += y_emb.to(x.dtype)
del y_emb

for i, layer in enumerate(self.layers):
cache_key = f"icl_block.layer{i}"
Expand Down Expand Up @@ -117,12 +118,14 @@ def forward(
if last_layer
else x[..., :R_train, :],
)
key_value = KVCacheEntry(
key=key[..., : self.kv_heads, :].contiguous(),
value=value[..., : self.kv_heads, :].contiguous(),
)
del key, value
x_query = layer(
query=x[..., R_train:, :],
key_value=KVCacheEntry(
key=key[..., : self.kv_heads, :].contiguous(),
value=value[..., : self.kv_heads, :].contiguous(),
),
key_value=key_value,
out=None if torch.is_grad_enabled() else x[..., R_train:, :],
)
if last_layer:
Expand Down
3 changes: 2 additions & 1 deletion sdm/models/kumo/tabular/recipe.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@ def numerical_processor() -> sp.Sequential:
sp.RobustScale(),
sp.ClipSoft(3.0),
],
sp.RankGaussian(),
method="round_robin",
),
sp.ClipSigma(threshold=4.0),
Expand Down Expand Up @@ -49,7 +50,7 @@ def numerical_processor() -> sp.Sequential:
sp.StypeDispatch(
categorical=[
sp.AlignCategories(),
sp.ShuffleCategories(method="shift"),
sp.ShuffleCategories(method="balanced_shift"),
],
numerical=[
sp.Standardize(),
Expand Down
105 changes: 103 additions & 2 deletions sdm/models/kumo/tabular/row_embedding.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,12 +3,14 @@

# ruff: noqa: D101, D102

import math
from typing import Any, cast

import torch
from torch import Tensor
from torch.nn import Embedding, Linear, ModuleList, Parameter

from sdm._memory import chunk_memory_limit
from sdm.cache import Cache, KVCacheEntry
from sdm.models.kumo.tabular.block import KumoTabularTransformerBlock
from sdm.models.kumo.tabular.cell_embedding import CellEmbedding
Expand Down Expand Up @@ -66,7 +68,7 @@ def __init__(
**factory_kwargs,
)

self.col_blocks = ModuleList(
self.col_blocks: ModuleList[InducedTransformerBlock] = ModuleList(
InducedTransformerBlock(
channels=channels,
num_inducing_points=num_inducing_points,
Expand All @@ -89,7 +91,7 @@ def __init__(
)
for _ in range(num_layers)
)
self.row_blocks = ModuleList(
self.row_blocks: ModuleList[KumoTabularTransformerBlock] = ModuleList(
KumoTabularTransformerBlock(
channels=channels,
num_heads=num_heads,
Expand All @@ -114,6 +116,105 @@ def forward(
*,
cache: Cache | None = None,
) -> Tensor: # [..., R, K * D]
starts = self._pass_starts(x, train_size=y.size(-1), cache=cache)
if len(starts) == 1:
return self._forward(x, y, categorical_mask, cache=cache)

# Query rows only read context state, which the first pass records
# for replay in later passes.
cache = Cache() if cache is None else cache
first = self._forward(
x=x[..., : starts[1], :],
y=y,
categorical_mask=categorical_mask,
cache=cache,
) # [..., starts[1], K * D]
cache.freeze()
out = first.new_empty((*first.shape[:-2], x.size(-2), first.size(-1)))
out[..., : starts[1], :] = first
del first
ends = [*starts[2:], x.size(-2)]
for start, end in zip(starts[1:], ends, strict=True):
out[..., start:end, :] = self._forward(
x=x[..., start:end, :],
y=y[..., :0],
categorical_mask=categorical_mask,
cache=cache,
)
return out

def _pass_starts(
self,
x: Tensor, # [..., R, C]
train_size: int,
cache: Cache | None,
) -> list[int]:
# First rows of the passes that embed the rows of `x`. Passes run
# without gradients on CUDA and replay context state from a cache.
if (
torch.is_grad_enabled()
or not x.is_cuda
or (cache is not None and cache.is_recording)
):
return [0]
*B, R, C = x.size()
N = math.prod(B)
K, D = self.readout_token.size(-2), self.channels
G = self.cell_embedding.group_size
M = self.col_blocks[0].inducing_points.size(-2)
s = (
torch.get_autocast_dtype(x.device.type).itemsize
if torch.is_autocast_enabled(x.device.type)
else x.element_size()
)
budget = chunk_memory_limit(x.device)
# Bytes per row: the cell buffer, plus the missingness mask, imputed
# values and their feature groups while embedding cells.
row_bytes = N * (
(K + C) * D * s + (G + 1) * (x.element_size() + 1) * C
)
# Without a cache, the context pass records the key/value projections
# of all column blocks for the query passes. Query rows that fit the
# chunk memory budget plus these projections run with the context.
state_bytes = 0
if cache is None:
state_bytes = 2 * N * C * M * D * s * len(self.col_blocks)
if (R - train_size) * row_bytes <= budget + state_bytes:
return [0]

# Row blocks run the rows of all batch entries in chunks of `chunk`
# rows. In passes starting on multiples of `grid` rows, every row runs
# in a chunk of the same size as in a single pass, since the last pass
# holds the partial last chunk of a single pass, the last
# `N * R % chunk` rows of the last batch entry. Attention over long
# rows rounds differently in small chunks, so this keeps passes equal
# to a single pass up to rare rounding differences in small passes.
chunk = self.row_blocks[0].auto_batch_size_limit(
device=x.device,
element_size=s,
query_length=K + C,
key_value_length=K + C,
)
grid = chunk // math.gcd(N, chunk)
context = -(-train_size // grid) * grid
last = R - max(N * R % chunk, 1)
if context > last:
return [0]
# Balanced query passes need no more memory than the context pass,
# the budget or one grid of rows, whichever is more.
grids = max(max(train_size, budget // row_bytes) // grid, 1)
num_passes = -(-(R - context) // (grids * grid))
step = -(-(R - context) // (num_passes * grid)) * grid
return [0, *range(context or step, last + 1, step)]

def _forward(
self,
x: Tensor, # [..., R, C]
y: Tensor, # [..., R_train]
categorical_mask: Tensor, # [..., C]
*,
cache: Cache | None = None,
) -> Tensor: # [..., R, K * D]

*B, R, C = x.size()
R_train = y.size(-1)
Expand Down
10 changes: 3 additions & 7 deletions sdm/models/tabfm/cell_embedding.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,13 +18,14 @@
# ruff: noqa: D101, D102

import math
import os
from typing import Any, Literal

import torch
from torch import Tensor
from torch.nn import Linear

from sdm._memory import chunk_memory_limit


class CellEmbedding(torch.nn.Module):
def __init__(
Expand Down Expand Up @@ -132,12 +133,7 @@ def forward(
+ bias.numel() * bias.element_size()
)

memory_limit = int(
torch.cuda.get_device_properties(x.device).total_memory
* torch.cuda.get_per_process_memory_fraction(x.device)
* float(os.getenv("SDM_CHUNK_MEMORY_FRACTION", "0.05"))
)
memory_limit -= fixed_bytes
memory_limit = chunk_memory_limit(x.device) - fixed_bytes
batch_size_limit = memory_limit // max(bytes_per_example, 1)
batch_size_limit = max(batch_size_limit, 1)

Expand Down
Loading
Loading