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
12 changes: 3 additions & 9 deletions benchmark/tabular/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -222,22 +222,16 @@ def _predict_proba(
self.autocast_dtype,
enabled=x_query.is_cuda,
):
outputs = self.model._forward_members(
outputs, dtypes = self.model._forward_members(
contexts=self._contexts,
queries=queries,
estimator_batch_size=self._get_model_params()[
"estimator_batch_size"
],
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, dtypes)

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
61 changes: 27 additions & 34 deletions sdm/models/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -195,7 +195,7 @@ def forward(
related_tables=related_query_tables,
)

outs = self._forward_members(
outs, dtypes = self._forward_members(
contexts=contexts,
queries=queries,
estimator_batch_size=estimator_batch_size,
Expand All @@ -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,
dtypes=dtypes,
)

def fit(
self,
Expand Down Expand Up @@ -442,6 +437,7 @@ def predict(
compute_stream.wait_stream(transfer_stream)

outs: list[TableTensor] = []
dtypes: list[torch.dtype] = []
start = 0
for i in range(len(caches)):
cache, next_cache = next_cache, None
Expand All @@ -465,6 +461,7 @@ def predict(
strict=True,
)
]
dtypes.extend(member.x.dtype for member in members)
start += len(x_schemas)

if i + 1 < len(caches):
Expand Down Expand Up @@ -529,19 +526,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,
dtypes=dtypes,
)

def clear(self) -> None:
r"""Clear cached context state created by :meth:`fit`."""
Expand Down Expand Up @@ -644,13 +636,13 @@ def _forward_members(
callbacks: Sequence[Callback] | None = None,
generator: torch.Generator | None = None,
**kwargs: Any,
) -> list[TableTensor]:
) -> tuple[list[TableTensor], list[torch.dtype]]:
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``.
Returns member outputs before target inversion and ``recipe.output``,
and their callback-prepared query dtypes.
"""
callbacks = () if callbacks is None else callbacks
requires_grad = self.training
Expand All @@ -661,18 +653,19 @@ def _forward_members(
# Complete each sequential member before preparing the next member.
if estimator_batch_size == 1 and len(contexts) > 1:
outs: list[TableTensor] = []
dtypes: list[torch.dtype] = []
for context, query in zip(contexts, queries, strict=True):
outs.extend(
self._forward_members(
contexts=(context,),
queries=(query,),
estimator_batch_size=1,
callbacks=callbacks,
generator=generator,
**kwargs,
)
member_outs, member_dtypes = self._forward_members(
contexts=(context,),
queries=(query,),
estimator_batch_size=1,
callbacks=callbacks,
generator=generator,
**kwargs,
)
return outs
outs.extend(member_outs)
dtypes.extend(member_dtypes)
return outs, dtypes

members = [
self._prepare_context(context, callbacks) for context in contexts
Expand Down Expand Up @@ -716,7 +709,7 @@ def _forward_members(
generator=generator,
**kwargs,
)
return outs
return outs, [query.x.dtype for query in queries]

def _prepare_context(
self,
Expand Down Expand Up @@ -796,7 +789,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
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
Loading