Skip to content
Open
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
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
48 changes: 48 additions & 0 deletions test/models/kumo/tabular/test_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,8 @@
import sdm.processing as sp
from sdm import CategoricalTensor, Stype, TableTensor
from sdm.models import KumoTabular
from sdm.models.kumo.tabular.row_embedding import RowEmbedding
from sdm.testing import onlyCUDA


def _build(
Expand Down Expand Up @@ -90,6 +92,52 @@ def _recipe() -> sp.Recipe:
)


@onlyCUDA
@pytest.mark.parametrize("task", ["classification", "regression"])
def test_query_passes_match_single_pass(
task: Literal["classification", "regression"],
monkeypatch: pytest.MonkeyPatch,
) -> None:
model = _build(task, "small").cuda().eval()
x = torch.randn(1008, 4, device="cuda")
x[::3, 0] = float("nan")
x_context, x_query = TableTensor.from_tensor(x).split([8, 1000], dim=0)
# 12 classes run through ECOC.
target = (
TableTensor(
columns={Stype.categorical: ("target",)},
categorical=CategoricalTensor(
code=torch.arange(8, device="cuda").unsqueeze(-1),
categories=(torch.arange(12, device="cuda"),),
),
)
if task == "classification"
else TableTensor.from_tensor(torch.randn(8, 1, device="cuda"))
)

def forward() -> TableTensor:
return model(
x_context=x_context,
y_context=target,
x_query=x_query,
num_estimators=2,
generator=torch.Generator("cuda").manual_seed(0),
)

# A chunk memory limit of 1 MiB embeds the query rows in passes.
total_memory = torch.cuda.get_device_properties("cuda").total_memory
monkeypatch.setenv("SDM_CHUNK_MEMORY_FRACTION", str(2**20 / total_memory))
actual = forward()
monkeypatch.setattr(
RowEmbedding,
"_pass_starts",
lambda self, x, train_size, cache: [0],
)
expected = forward()

torch.testing.assert_close(actual.numerical, expected.numerical)


def test_forward(
cls_model: KumoTabular,
reg_model: KumoTabular,
Expand Down
103 changes: 102 additions & 1 deletion test/models/kumo/tabular/test_row_embedding.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,12 @@
# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

import pytest
import torch

from sdm.cache import Cache
from sdm.models.kumo.tabular.row_embedding import RowEmbedding
from sdm.testing import withCUDA
from sdm.testing import onlyCUDA, withCUDA


@withCUDA
Expand Down Expand Up @@ -46,3 +47,103 @@ def test_row_embedding(device: torch.device) -> None:
assert out.size() == (2, 2, 32)
assert out.device == device
torch.testing.assert_close(out, expected[:, 3:], atol=1e-5, rtol=1e-5)


@onlyCUDA
@pytest.mark.parametrize("cached", [False, True])
def test_row_embedding_passes(
cached: bool,
monkeypatch: pytest.MonkeyPatch,
) -> None:
encoder = RowEmbedding(
num_classes=0,
channels=16,
num_layers=2,
num_heads=2,
group_size=3,
num_frequencies=8,
num_inducing_points=4,
num_readout_tokens=2,
device="cuda",
)
# Randomize the zero-initialized residual branches, so that attention
# outputs reach the embedding.
for parameter in encoder.parameters():
if not parameter.any():
torch.nn.init.normal_(parameter, std=0.02)
# With many batch entries, e.g. estimators or ECOC tasks, a chunk of the
# row attention spans more rows than the chunk memory limit allows.
x = torch.randn(17, 64, 3, device="cuda")
x[x > 1.0] = torch.nan
y = torch.randn(17, 4, device="cuda")
categorical_mask = torch.tensor([True, False, False], device="cuda")

@torch.inference_mode()
def embed() -> torch.Tensor:
if not cached:
return encoder(x, y, categorical_mask)
cache = Cache()
encoder(x[:, :4], y, categorical_mask, cache=cache)
return encoder(
x[:, 4:],
y[:, :0],
categorical_mask,
cache=cache.freeze(),
)

expected = embed()
# Small, intermediate and single-pass budgets exercise pass boundaries.
total_memory = torch.cuda.get_device_properties("cuda").total_memory
for budget in (2**10, 2**14, 2**19):
fraction = budget / total_memory
monkeypatch.setenv("SDM_CHUNK_MEMORY_FRACTION", str(fraction))
torch.testing.assert_close(embed(), expected)


@onlyCUDA
def test_row_embedding_passes_on_wide_tables(
monkeypatch: pytest.MonkeyPatch,
) -> None:
# Attention over more than 256 row tokens splits keys and values for
# small batches of rows, so passes must keep the row chunks of one pass.
encoder = RowEmbedding(
num_classes=10,
channels=64,
num_layers=2,
num_heads=2,
group_size=3,
num_frequencies=8,
num_inducing_points=8,
num_readout_tokens=2,
device="cuda",
)
# Randomize the zero-initialized residual branches, so that attention
# outputs reach the embedding.
for parameter in encoder.parameters():
if not parameter.any():
torch.nn.init.normal_(parameter, std=0.02)
x = torch.randn(3, 1169, 300, device="cuda")
x[x > 2.0] = torch.nan
y = torch.randint(10, (3, 200), device="cuda")
categorical_mask = torch.arange(300, device="cuda") % 4 == 0

def embed() -> tuple[torch.Tensor, int]:
torch.cuda.synchronize()
baseline = torch.cuda.memory_allocated()
torch.cuda.reset_peak_memory_stats()
with torch.no_grad(), torch.autocast("cuda", dtype=torch.float16):
out = encoder(x, y, categorical_mask)
torch.cuda.synchronize()
return out, torch.cuda.max_memory_allocated() - baseline

monkeypatch.setenv("SDM_CHUNK_MEMORY_FRACTION", "0.001")
actual, peak = embed()
monkeypatch.setattr(
RowEmbedding,
"_pass_starts",
lambda self, x, train_size, cache: [0],
)
expected, expected_peak = embed()

assert torch.equal(actual, expected)
assert peak < expected_peak
Loading