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..31c514271 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,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