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
45 changes: 45 additions & 0 deletions tests/test_autobatching.py
Original file line number Diff line number Diff line change
Expand Up @@ -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_chunk", [1, 16, 40, 10_000])
def test_n_edges_scalers_chunked_matches_single_pass(
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_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_chunk", [1, 16, 40, 10_000])
def test_n_edges_scalers_chunk_bounds(
mixed_many_sim_state: ts.SimState,
max_atoms_per_chunk: 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_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_chunk 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],
Expand Down
46 changes: 36 additions & 10 deletions torch_sim/autobatching.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,11 @@
# nvalchemiops kernels used by ORB v3) raise "Failed to allocate <n> 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_CHUNK = 250_000


def to_constant_volume_bins( # noqa: C901
items: dict[int, float] | list[Any],
Expand Down Expand Up @@ -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_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_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.
"""
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_chunk:
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(
Expand Down
Loading