From a4c614bb3b5be1dd84be54ae209b5ec4c3f79b4a Mon Sep 17 00:00:00 2001 From: Jingang Qu Date: Mon, 28 Sep 2026 22:40:27 +0200 Subject: [PATCH 1/2] Chunk Kumo row embeddings - Without gradients on CUDA, embed the context rows once and the query rows in balanced passes that replay the recorded context state, when the query cells would exceed the chunk memory limit. - Align passes with the chunks of the row attention, so every row runs in a chunk of the same size as in a single pass. Passes then match a single pass up to rare rounding differences in small passes. - Free the label embedding and each layer's full key/value early in the ICL block. Signed-off-by: Jingang Qu --- sdm/models/kumo/tabular/icl.py | 11 +- sdm/models/kumo/tabular/row_embedding.py | 105 +++++++++++++++++- test/models/kumo/tabular/test_model.py | 48 ++++++++ .../models/kumo/tabular/test_row_embedding.py | 98 +++++++++++++++- 4 files changed, 255 insertions(+), 7 deletions(-) diff --git a/sdm/models/kumo/tabular/icl.py b/sdm/models/kumo/tabular/icl.py index be431353d..067d5f619 100644 --- a/sdm/models/kumo/tabular/icl.py +++ b/sdm/models/kumo/tabular/icl.py @@ -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}" @@ -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: diff --git a/sdm/models/kumo/tabular/row_embedding.py b/sdm/models/kumo/tabular/row_embedding.py index f69ff3e0e..ccbb0afff 100644 --- a/sdm/models/kumo/tabular/row_embedding.py +++ b/sdm/models/kumo/tabular/row_embedding.py @@ -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 @@ -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, @@ -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, @@ -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) diff --git a/test/models/kumo/tabular/test_model.py b/test/models/kumo/tabular/test_model.py index 0f153a8a8..c7840478d 100644 --- a/test/models/kumo/tabular/test_model.py +++ b/test/models/kumo/tabular/test_model.py @@ -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( @@ -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, diff --git a/test/models/kumo/tabular/test_row_embedding.py b/test/models/kumo/tabular/test_row_embedding.py index 29fa0380e..76d69fa63 100644 --- a/test/models/kumo/tabular/test_row_embedding.py +++ b/test/models/kumo/tabular/test_row_embedding.py @@ -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 @@ -46,3 +47,98 @@ 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", + ) + # 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 From dfd428758f80229a884ebc0c38391f40b1d6c811 Mon Sep 17 00:00:00 2001 From: Martin Jurkovic Date: Tue, 29 Sep 2026 10:50:12 +0200 Subject: [PATCH 2/2] Randomize residual branches in Kumo row embedding tests --- test/models/kumo/tabular/test_row_embedding.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/test/models/kumo/tabular/test_row_embedding.py b/test/models/kumo/tabular/test_row_embedding.py index 76d69fa63..31c514271 100644 --- a/test/models/kumo/tabular/test_row_embedding.py +++ b/test/models/kumo/tabular/test_row_embedding.py @@ -66,6 +66,11 @@ def test_row_embedding_passes( 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")