From 4112b71a8894b9afe1ef9656e213e9f6ec600081 Mon Sep 17 00:00:00 2001 From: Ard van Noordenne Date: Mon, 5 Oct 2026 14:58:09 +0200 Subject: [PATCH 1/2] Compute n_edges memory scalers in atom-bounded chunks _n_edges_scalers ran one neighbor-list pass over the whole state. The Alchemiops backend holds fixed-width buffers of ~4.4 KiB per atom at a 6 A cutoff, so a 10.7M-atom state needs ~45 GiB for this pass and OOMs on an A40 before the first optimization step. Run the pass over contiguous runs of whole systems holding at most N_EDGES_MAX_ATOMS_PER_PASS (250k) atoms; a larger system runs alone. Edge counts do not depend on other systems, so the scalers are unchanged. --- tests/test_autobatching.py | 45 +++++++++++++++++++++++++++++++++++++ torch_sim/autobatching.py | 46 +++++++++++++++++++++++++++++--------- 2 files changed, 81 insertions(+), 10 deletions(-) diff --git a/tests/test_autobatching.py b/tests/test_autobatching.py index 0da77371..7e0712ce 100644 --- a/tests/test_autobatching.py +++ b/tests/test_autobatching.py @@ -185,6 +185,51 @@ def test_n_edges_scalers_batched(ar_double_sim_state: ts.SimState) -> None: assert all(v >= 0 for v in result) +@pytest.fixture +def mixed_many_sim_state( + ar_supercell_sim_state: ts.SimState, si_sim_state: ts.SimState +) -> ts.SimState: + """Batched state of alternating large and small periodic systems.""" + return ts.concatenate_states( + [ar_supercell_sim_state, si_sim_state, si_sim_state, ar_supercell_sim_state], + device=si_sim_state.device, + ) + + +@pytest.mark.parametrize("max_atoms_per_pass", [1, 16, 40, 10_000]) +def test_n_edges_scalers_chunked_matches_single_pass( + mixed_many_sim_state: ts.SimState, max_atoms_per_pass: int +) -> None: + """Splitting the neighbor-list pass leaves every per-system edge count unchanged.""" + state = mixed_many_sim_state + single = _n_edges_scalers(state, cutoff=5.0, max_atoms_per_pass=state.n_atoms) + chunked = _n_edges_scalers(state, cutoff=5.0, max_atoms_per_pass=max_atoms_per_pass) + assert chunked == single + + +@pytest.mark.parametrize("max_atoms_per_pass", [1, 16, 40, 10_000]) +def test_n_edges_scalers_chunk_bounds( + mixed_many_sim_state: ts.SimState, + max_atoms_per_pass: int, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Each pass holds whole systems and exceeds the bound only for a lone system.""" + passes: list[tuple[int, int]] = [] + real_nl = ts.autobatching.torchsim_nl + + def recording_nl(**kwargs: Any) -> Any: + passes.append((int(kwargs["system_idx"].max()) + 1, kwargs["positions"].shape[0])) + return real_nl(**kwargs) + + monkeypatch.setattr(ts.autobatching, "torchsim_nl", recording_nl) + state = mixed_many_sim_state + _n_edges_scalers(state, cutoff=5.0, max_atoms_per_pass=max_atoms_per_pass) + + assert sum(n_sys for n_sys, _ in passes) == state.n_systems + assert sum(n_at for _, n_at in passes) == state.n_atoms + assert all(n_at <= max_atoms_per_pass or n_sys == 1 for n_sys, n_at in passes) + + @pytest.mark.parametrize("items", [[], {}]) def test_to_constant_volume_bins_empty_input( items: list[Any] | dict[int, float], diff --git a/torch_sim/autobatching.py b/torch_sim/autobatching.py index 86a3e4c5..9263d634 100644 --- a/torch_sim/autobatching.py +++ b/torch_sim/autobatching.py @@ -42,6 +42,11 @@ # nvalchemiops kernels used by ORB v3) raise "Failed to allocate bytes". DEFAULT_OOM_ERROR_MESSAGES = ("CUDA out of memory", "Failed to allocate") +# Atoms per neighbor-list pass when computing n_edges scalers. Alchemiops holds +# fixed-width buffers of ~4.4 KiB per atom at a 6 A cutoff, so this caps the +# pass at about 1 GiB however large the state is. +N_EDGES_MAX_ATOMS_PER_PASS = 250_000 + def to_constant_volume_bins( # noqa: C901 items: dict[int, float] | list[Any], @@ -290,16 +295,37 @@ def determine_max_batch_size( torch.cuda.empty_cache() -def _n_edges_scalers(state: SimState, cutoff: float) -> list[float]: - """Return per-system edge counts from the neighbor list as memory scalers.""" - _, system_mapping, _ = torchsim_nl( - positions=state.positions, - cell=state.cell, - pbc=state.pbc, - cutoff=cutoff, - system_idx=state.system_idx, - ) - return system_mapping.bincount(minlength=state.n_systems).float().tolist() +def _n_edges_scalers( + state: SimState, + cutoff: float, + max_atoms_per_pass: int = N_EDGES_MAX_ATOMS_PER_PASS, +) -> list[float]: + """Return per-system edge counts from the neighbor list as memory scalers. + + The neighbor list runs over contiguous runs of systems holding at most + ``max_atoms_per_pass`` atoms, so peak memory does not grow with the size of + the state. A system larger than the bound runs alone. Edge counts do not + depend on other systems, so the result equals a single pass over the state. + """ + n_atoms_per_system = state.n_atoms_per_system.tolist() + n_systems = len(n_atoms_per_system) + scalers: list[float] = [] + s0 = a0 = 0 + while s0 < n_systems: + s1, a1 = s0 + 1, a0 + n_atoms_per_system[s0] + while s1 < n_systems and a1 - a0 + n_atoms_per_system[s1] <= max_atoms_per_pass: + a1 += n_atoms_per_system[s1] + s1 += 1 + _, system_mapping, _ = torchsim_nl( + positions=state.positions[a0:a1], + cell=state.cell[s0:s1], + pbc=state.pbc, + cutoff=cutoff, + system_idx=state.system_idx[a0:a1] - s0, + ) + scalers.extend(system_mapping.bincount(minlength=s1 - s0).float().tolist()) + s0, a0 = s1, a1 + return scalers def calculate_memory_scalers( From 732e8c5f2943252a44a5868704d063195b3a2dd9 Mon Sep 17 00:00:00 2001 From: Ard van Noordenne Date: Mon, 5 Oct 2026 21:47:15 +0200 Subject: [PATCH 2/2] Rename max_atoms_per_pass to max_atoms_per_chunk --- tests/test_autobatching.py | 16 ++++++++-------- torch_sim/autobatching.py | 8 ++++---- 2 files changed, 12 insertions(+), 12 deletions(-) diff --git a/tests/test_autobatching.py b/tests/test_autobatching.py index 7e0712ce..65b7143a 100644 --- a/tests/test_autobatching.py +++ b/tests/test_autobatching.py @@ -196,21 +196,21 @@ def mixed_many_sim_state( ) -@pytest.mark.parametrize("max_atoms_per_pass", [1, 16, 40, 10_000]) +@pytest.mark.parametrize("max_atoms_per_chunk", [1, 16, 40, 10_000]) def test_n_edges_scalers_chunked_matches_single_pass( - mixed_many_sim_state: ts.SimState, max_atoms_per_pass: int + mixed_many_sim_state: ts.SimState, max_atoms_per_chunk: int ) -> None: """Splitting the neighbor-list pass leaves every per-system edge count unchanged.""" state = mixed_many_sim_state - single = _n_edges_scalers(state, cutoff=5.0, max_atoms_per_pass=state.n_atoms) - chunked = _n_edges_scalers(state, cutoff=5.0, max_atoms_per_pass=max_atoms_per_pass) + single = _n_edges_scalers(state, cutoff=5.0, max_atoms_per_chunk=state.n_atoms) + chunked = _n_edges_scalers(state, cutoff=5.0, max_atoms_per_chunk=max_atoms_per_chunk) assert chunked == single -@pytest.mark.parametrize("max_atoms_per_pass", [1, 16, 40, 10_000]) +@pytest.mark.parametrize("max_atoms_per_chunk", [1, 16, 40, 10_000]) def test_n_edges_scalers_chunk_bounds( mixed_many_sim_state: ts.SimState, - max_atoms_per_pass: int, + max_atoms_per_chunk: int, monkeypatch: pytest.MonkeyPatch, ) -> None: """Each pass holds whole systems and exceeds the bound only for a lone system.""" @@ -223,11 +223,11 @@ def recording_nl(**kwargs: Any) -> Any: monkeypatch.setattr(ts.autobatching, "torchsim_nl", recording_nl) state = mixed_many_sim_state - _n_edges_scalers(state, cutoff=5.0, max_atoms_per_pass=max_atoms_per_pass) + _n_edges_scalers(state, cutoff=5.0, max_atoms_per_chunk=max_atoms_per_chunk) assert sum(n_sys for n_sys, _ in passes) == state.n_systems assert sum(n_at for _, n_at in passes) == state.n_atoms - assert all(n_at <= max_atoms_per_pass or n_sys == 1 for n_sys, n_at in passes) + assert all(n_at <= max_atoms_per_chunk or n_sys == 1 for n_sys, n_at in passes) @pytest.mark.parametrize("items", [[], {}]) diff --git a/torch_sim/autobatching.py b/torch_sim/autobatching.py index 9263d634..6d997981 100644 --- a/torch_sim/autobatching.py +++ b/torch_sim/autobatching.py @@ -45,7 +45,7 @@ # Atoms per neighbor-list pass when computing n_edges scalers. Alchemiops holds # fixed-width buffers of ~4.4 KiB per atom at a 6 A cutoff, so this caps the # pass at about 1 GiB however large the state is. -N_EDGES_MAX_ATOMS_PER_PASS = 250_000 +N_EDGES_MAX_ATOMS_PER_CHUNK = 250_000 def to_constant_volume_bins( # noqa: C901 @@ -298,12 +298,12 @@ def determine_max_batch_size( def _n_edges_scalers( state: SimState, cutoff: float, - max_atoms_per_pass: int = N_EDGES_MAX_ATOMS_PER_PASS, + max_atoms_per_chunk: int = N_EDGES_MAX_ATOMS_PER_CHUNK, ) -> list[float]: """Return per-system edge counts from the neighbor list as memory scalers. The neighbor list runs over contiguous runs of systems holding at most - ``max_atoms_per_pass`` atoms, so peak memory does not grow with the size of + ``max_atoms_per_chunk`` atoms, so peak memory does not grow with the size of the state. A system larger than the bound runs alone. Edge counts do not depend on other systems, so the result equals a single pass over the state. """ @@ -313,7 +313,7 @@ def _n_edges_scalers( s0 = a0 = 0 while s0 < n_systems: s1, a1 = s0 + 1, a0 + n_atoms_per_system[s0] - while s1 < n_systems and a1 - a0 + n_atoms_per_system[s1] <= max_atoms_per_pass: + while s1 < n_systems and a1 - a0 + n_atoms_per_system[s1] <= max_atoms_per_chunk: a1 += n_atoms_per_system[s1] s1 += 1 _, system_mapping, _ = torchsim_nl(