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,