From c54538c63e10c1e49f6d25c4468051c4748999bd Mon Sep 17 00:00:00 2001 From: Jingang Qu Date: Mon, 28 Sep 2026 22:40:27 +0200 Subject: [PATCH 1/2] Share chunk-memory sizing - Move the `SDM_CHUNK_MEMORY_FRACTION` budget of attention and TabFM cell embedding chunks into one helper, `sdm._memory.chunk_memory_limit`. - Expose the automatic attention batch size limit as `TransformerBlock.auto_batch_size_limit`, so callers can plan passes that align with its chunks. Signed-off-by: Jingang Qu --- sdm/_memory.py | 19 ++++++++ sdm/models/tabfm/cell_embedding.py | 10 ++-- sdm/nn/attention.py | 62 +++++++++++++----------- test/models/tabfm/test_cell_embedding.py | 10 +++- test/nn/test_attention.py | 16 +++++- 5 files changed, 81 insertions(+), 36 deletions(-) create mode 100644 sdm/_memory.py diff --git a/sdm/_memory.py b/sdm/_memory.py new file mode 100644 index 000000000..47a735a3d --- /dev/null +++ b/sdm/_memory.py @@ -0,0 +1,19 @@ +# 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")) + ) diff --git a/sdm/models/tabfm/cell_embedding.py b/sdm/models/tabfm/cell_embedding.py index f2a8d64b7..f2c086408 100644 --- a/sdm/models/tabfm/cell_embedding.py +++ b/sdm/models/tabfm/cell_embedding.py @@ -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__( @@ -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) diff --git a/sdm/nn/attention.py b/sdm/nn/attention.py index b2020d9ff..42190df1f 100644 --- a/sdm/nn/attention.py +++ b/sdm/nn/attention.py @@ -4,7 +4,6 @@ """Attention modules for structured tensor models.""" import math -import os from typing import Any, Literal, overload import torch @@ -12,6 +11,7 @@ from torch import Tensor from torch.nn import Linear +from sdm._memory import chunk_memory_limit from sdm.cache import KVCacheEntry from sdm.nn import QueryScaling @@ -533,32 +533,22 @@ def forward( ) if batch_size_limit == "auto": - batch_size_limit = None - if query.is_cuda: - key_value_length: int | None = None - if isinstance(key_value, Tensor): - key_value_length = key_value.size(-2) - elif isinstance(key_value, KVCacheEntry): - key_value_length = key_value.key.size(-3) - - bytes_per_example = self.peak_bytes_per_example( - element_size=torch.empty( - size=(), - dtype=torch.get_autocast_dtype(query.device.type), - ).element_size() - if torch.is_autocast_enabled(query.device.type) - else query.element_size(), - query_length=query.size(-2), - key_value_length=key_value_length, - ) - - memory_limit = int( - torch.cuda.get_device_properties(query.device).total_memory - * torch.cuda.get_per_process_memory_fraction(query.device) - * float(os.getenv("SDM_CHUNK_MEMORY_FRACTION", "0.05")) - ) - batch_size_limit = memory_limit // max(bytes_per_example, 1) - batch_size_limit = max(batch_size_limit, 1) + key_value_length: int | None = None + if isinstance(key_value, Tensor): + key_value_length = key_value.size(-2) + elif isinstance(key_value, KVCacheEntry): + key_value_length = key_value.key.size(-3) + + batch_size_limit = self.auto_batch_size_limit( + device=query.device, + element_size=torch.get_autocast_dtype( + query.device.type + ).itemsize + if torch.is_autocast_enabled(query.device.type) + else query.element_size(), + query_length=query.size(-2), + key_value_length=key_value_length, + ) batch_size_limit = min(batch_size_limit or 65_535, 65_535) @@ -728,6 +718,24 @@ def peak_bytes_per_example( r""":meta private:""" # noqa: D415 return 0 + def auto_batch_size_limit( + self, + device: torch.device, + element_size: int, + query_length: int, + key_value_length: int | None = None, + ) -> int: + r""":meta private:""" # noqa: D415 + if device.type != "cuda": + return 65_535 + bytes_per_example = self.peak_bytes_per_example( + element_size=element_size, + query_length=query_length, + key_value_length=key_value_length, + ) + limit = chunk_memory_limit(device) // max(bytes_per_example, 1) + return min(max(limit, 1), 65_535) + def _batch_shape( query: Tensor, diff --git a/test/models/tabfm/test_cell_embedding.py b/test/models/tabfm/test_cell_embedding.py index 8ad2eb316..23fd8cc8b 100644 --- a/test/models/tabfm/test_cell_embedding.py +++ b/test/models/tabfm/test_cell_embedding.py @@ -9,7 +9,10 @@ @withCUDA -def test_cell_embedding(device: torch.device) -> None: +def test_cell_embedding( + device: torch.device, + monkeypatch: pytest.MonkeyPatch, +) -> None: module = CellEmbedding( channels=8, group_size=3, @@ -44,6 +47,11 @@ def test_cell_embedding(device: torch.device) -> None: assert module.num_lin.weight.grad is not None assert module.cat_lin.weight.grad is not None + monkeypatch.setenv("SDM_CHUNK_MEMORY_FRACTION", "1e-12") + with torch.no_grad(): + auto_out = module(x, categorical_mask, batch_size_limit="auto") + torch.testing.assert_close(auto_out, out1) + @withCUDA def test_cell_embedding_mixed_dtype_output( diff --git a/test/nn/test_attention.py b/test/nn/test_attention.py index 2dee342b4..72903e870 100644 --- a/test/nn/test_attention.py +++ b/test/nn/test_attention.py @@ -551,7 +551,11 @@ def test_return_key_value_positional_compatibility() -> None: @withCUDA @pytest.mark.parametrize("qassmax", [False, True]) -def test_transformer_block(device: torch.device, qassmax: bool) -> None: +def test_transformer_block( + device: torch.device, + qassmax: bool, + monkeypatch: pytest.MonkeyPatch, +) -> None: batch_size = 2 query_len = 3 key_value_len = 5 @@ -645,6 +649,16 @@ def test_transformer_block(device: torch.device, qassmax: bool) -> None: assert chunked_buffered_out is chunked_buffer torch.testing.assert_close(chunked_buffered_out, out1) + monkeypatch.setenv("SDM_CHUNK_MEMORY_FRACTION", "1e-12") + with torch.no_grad(): + auto_out = module( + query=query, + key_value=key_value, + seqused_key_value=seqused_key_value, + batch_size_limit="auto", + ) + torch.testing.assert_close(auto_out, out1) + # Test no padding leakage new_key_value = key_value.clone() new_key_value[0, 3:] = torch.randn_like(new_key_value[0, 3:]) * 1000 From 62280d8cfe999f008733418759ba6c6257043305 Mon Sep 17 00:00:00 2001 From: Valter Hudovernik Date: Tue, 29 Sep 2026 00:02:49 +0200 Subject: [PATCH 2/2] Remove redundant automatic chunking test additions --- test/models/tabfm/test_cell_embedding.py | 10 +--------- test/nn/test_attention.py | 16 +--------------- 2 files changed, 2 insertions(+), 24 deletions(-) diff --git a/test/models/tabfm/test_cell_embedding.py b/test/models/tabfm/test_cell_embedding.py index 23fd8cc8b..8ad2eb316 100644 --- a/test/models/tabfm/test_cell_embedding.py +++ b/test/models/tabfm/test_cell_embedding.py @@ -9,10 +9,7 @@ @withCUDA -def test_cell_embedding( - device: torch.device, - monkeypatch: pytest.MonkeyPatch, -) -> None: +def test_cell_embedding(device: torch.device) -> None: module = CellEmbedding( channels=8, group_size=3, @@ -47,11 +44,6 @@ def test_cell_embedding( assert module.num_lin.weight.grad is not None assert module.cat_lin.weight.grad is not None - monkeypatch.setenv("SDM_CHUNK_MEMORY_FRACTION", "1e-12") - with torch.no_grad(): - auto_out = module(x, categorical_mask, batch_size_limit="auto") - torch.testing.assert_close(auto_out, out1) - @withCUDA def test_cell_embedding_mixed_dtype_output( diff --git a/test/nn/test_attention.py b/test/nn/test_attention.py index 72903e870..2dee342b4 100644 --- a/test/nn/test_attention.py +++ b/test/nn/test_attention.py @@ -551,11 +551,7 @@ def test_return_key_value_positional_compatibility() -> None: @withCUDA @pytest.mark.parametrize("qassmax", [False, True]) -def test_transformer_block( - device: torch.device, - qassmax: bool, - monkeypatch: pytest.MonkeyPatch, -) -> None: +def test_transformer_block(device: torch.device, qassmax: bool) -> None: batch_size = 2 query_len = 3 key_value_len = 5 @@ -649,16 +645,6 @@ def test_transformer_block( assert chunked_buffered_out is chunked_buffer torch.testing.assert_close(chunked_buffered_out, out1) - monkeypatch.setenv("SDM_CHUNK_MEMORY_FRACTION", "1e-12") - with torch.no_grad(): - auto_out = module( - query=query, - key_value=key_value, - seqused_key_value=seqused_key_value, - batch_size_limit="auto", - ) - torch.testing.assert_close(auto_out, out1) - # Test no padding leakage new_key_value = key_value.clone() new_key_value[0, 3:] = torch.randn_like(new_key_value[0, 3:]) * 1000