Skip to content
Open
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
39 changes: 39 additions & 0 deletions tests/test_autobatching.py
Original file line number Diff line number Diff line change
Expand Up @@ -554,6 +554,45 @@ def mock_measure(*_args: Any, **_kwargs: Any) -> float:
assert max_size >= 1


@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available")
@pytest.mark.parametrize("probe_ooms", [True, False])
def test_determine_max_batch_size_releases_cached_memory(
*,
si_sim_state: ts.SimState,
lj_model: LennardJonesModel,
monkeypatch: pytest.MonkeyPatch,
probe_ooms: bool,
) -> None:
"""The probe must hand the device back on every exit path.

Regression test: probing deliberately grows batches until the GPU runs out
of memory. If PyTorch's caching allocator keeps holding that memory after
the probe returns, a separate allocator - such as the Warp/cudaMallocAsync
pool behind the alchemiops neighbor lists - can fail to allocate even a few
bytes on the very next call.
"""

def mock_measure(*_args: Any, **_kwargs: Any) -> float:
# Reserve a chunk and release it, so it lands in PyTorch's cache the
# way a real probe's activations do.
buffer = torch.empty(256 * 1024**2, dtype=torch.uint8, device="cuda")
del buffer
if probe_ooms:
raise RuntimeError("CUDA out of memory")
return 0.1

monkeypatch.setattr(
"torch_sim.autobatching.measure_model_memory_forward", mock_measure
)

torch.cuda.empty_cache()
baseline = torch.cuda.memory_reserved()

determine_max_batch_size(si_sim_state, lj_model, max_atoms=10_000)

assert torch.cuda.memory_reserved() <= baseline


def test_determine_max_batch_size_reraises_non_oom_error(
si_sim_state: ts.SimState,
lj_model: LennardJonesModel,
Expand Down
69 changes: 39 additions & 30 deletions torch_sim/autobatching.py
Original file line number Diff line number Diff line change
Expand Up @@ -237,6 +237,11 @@ def determine_max_batch_size(
Notes:
The function returns a batch size slightly smaller than the actual maximum
(with a safety margin) to avoid operating too close to memory limits.

On CUDA devices the caching allocator is released before returning, so that
separately-allocating backends - such as the Warp/cudaMallocAsync pool behind
the alchemiops neighbor lists - are not starved by memory the probe left
cached.
"""
# Convert oom_error_message to list if it's a string
if isinstance(oom_error_message, str):
Expand All @@ -249,29 +254,40 @@ def determine_max_batch_size(
) * state.n_atoms <= max_atoms:
sizes.append(next_size)

for sys_idx in range(len(sizes)):
n_systems = sizes[sys_idx]
concat_state = ts.concatenate_states([state] * n_systems)

try:
measure_model_memory_forward(concat_state, model)
except Exception as exc:
exc_str = str(exc)
# Check if any of the OOM error messages match
if any(msg in exc_str for msg in oom_error_message):
safe_size = sizes[max(0, sys_idx - 2)]
logger.debug(
"OOM at %d systems (%d atoms), returning safe batch size %d",
n_systems,
concat_state.n_atoms,
safe_size,
)
return safe_size
try:
for sys_idx in range(len(sizes)):
n_systems = sizes[sys_idx]
concat_state = ts.concatenate_states([state] * n_systems)

try:
measure_model_memory_forward(concat_state, model)
except Exception as exc:
exc_str = str(exc)
# Check if any of the OOM error messages match
if any(msg in exc_str for msg in oom_error_message):
safe_size = sizes[max(0, sys_idx - 2)]
logger.debug(
"OOM at %d systems (%d atoms), returning safe batch size %d",
n_systems,
concat_state.n_atoms,
safe_size,
)
return safe_size

# Not an OOM error - re-raise
raise
# Not an OOM error - re-raise
raise

return sizes[-1]
return sizes[-1]
finally:
# The probing above deliberately grows batches until an OOM, which
# leaves PyTorch's caching allocator holding most of the device memory.
# Release it on every exit path, including the OOM return above: a
# separate allocator - such as the Warp/cudaMallocAsync pool behind the
# alchemiops neighbor lists - draws from what PyTorch is not holding, so
# it can otherwise fail to allocate even a few bytes on the next call.
if torch.cuda.is_available(): # pragma: no cover
torch.cuda.synchronize()
torch.cuda.empty_cache()


def _n_edges_scalers(state: SimState, cutoff: float) -> list[float]:
Expand Down Expand Up @@ -445,18 +461,11 @@ def estimate_max_memory_scaler(
f"and {max_state.n_systems} batches, and smallest system has "
f"{min_state.n_atoms} atoms and {min_state.n_systems} batches.",
)
# determine_max_batch_size releases the caching allocator on its way out, so
# the device is already handed back before the real optimization run starts.
min_state_max_batches = determine_max_batch_size(min_state, model, **kwargs)
max_state_max_batches = determine_max_batch_size(max_state, model, **kwargs)

# The probing above deliberately grows batches until an OOM, which leaves
# PyTorch's caching allocator holding most of the device memory. Release it
# so the real optimization run - and in particular any separate allocator
# such as the Warp/cudaMallocAsync pool used by ORB's neighbor lists - is
# not starved on the first real forward pass.
if torch.cuda.is_available(): # pragma: no cover
torch.cuda.synchronize()
torch.cuda.empty_cache()

return min(
min_state_max_batches * min_metric.item(),
max_state_max_batches * max_metric.item(),
Expand Down
Loading