Skip to content
Merged
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
19 changes: 19 additions & 0 deletions sdm/_memory.py
Original file line number Diff line number Diff line change
@@ -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"))
)
10 changes: 3 additions & 7 deletions sdm/models/tabfm/cell_embedding.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__(
Expand Down Expand Up @@ -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)

Expand Down
62 changes: 35 additions & 27 deletions sdm/nn/attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,14 +4,14 @@
"""Attention modules for structured tensor models."""

import math
import os
from typing import Any, Literal, overload

import torch
import torch.nn.functional as F
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

Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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,
Expand Down
Loading