From da8eff620c86b3756985f52f7d0fd6a91898169a Mon Sep 17 00:00:00 2001 From: Jim Garrison Date: Fri, 9 Oct 2026 15:34:04 -0400 Subject: [PATCH] Diagonalize the batches concurrently over groups of processes solve_sci_batch diagonalized each subspace in turn over every process. Because qiskit-addon-sqd passes all of an iteration's subspaces in one call, the processes can instead be divided into one group per subspace and the subspaces diagonalized at the same time. The division happens after the FCIDUMP is loaded, so the Hamiltonian is still read once over the whole communicator rather than once per group. Each group then runs _solve_sci_core over its own communicator, so a result is produced on each group's rank-0 process rather than only on the world's. The group leaders exchange their results and the collected list is broadcast, ordered by subspace index, so every process returns the same thing as before. Groups are equal in size, and any leftover processes are given no group rather than making one group larger. This is not tidiness: a schedule such as TrimPolicy diagonalizes these batches in order to rank them against each other, and a batch solved over more processes than its siblings reaches a slightly different floating-point energy, so the ranking would turn partly on how the processes happened to divide. Those idle processes still enter every collective, or the ones doing the work would wait on them forever. This is what trim SQD needs, and it needs no parameter. qiskit-addon-sqd's loop calls the solver once per round: TrimPolicy's screening round hands over every batch in one call, which is divided here, and its merged round hands over a single subspace, which is given every process. Dividing the processes is the solver's concern rather than the schedule's, which is what diagonalize_fermionic_hamiltonian documents for a collective sci_solver. The wavefunction dump is renamed from wavefunction.bin to wavefunction_.bin. SBD does not return the amplitudes through the binding: it writes them to this file and _solve_sci_core reads them back to build the SCIState. The directory holding it is created once and broadcast, so every group shares it, and the name used to be a constant -- safe only while the batches ran one after another, each write being consumed before the next began. Run concurrently, the groups would write one path at once and each would then read back whichever write landed last, building an SCIState from another group's subspace with no error raised. The batch index distinguishes them: every member of a group agrees on it and distinct groups differ, so the names cannot collide and no coordination is needed. With carryover_type=1 the dump is never requested, so the collision cannot arise there. Communicators are freed rather than left to the garbage collector, since a long run would otherwise exhaust them. Only the ones Split produced are freed, never the caller's own. One consequence for the process grid. adet_comm_size, bdet_comm_size and task_comm_size describe a single diagonalization, so they are relative to whichever communicator that diagonalization is handed -- a group, once the processes are divided, rather than all of them. The rule is divisibility: diag() derives the helper dimension by integer division, h_comm_size = mpi_size / (task_comm_size * base_comm_size), and TaskCommunicator checks that multiplying the four back recovers mpi_size, which fails exactly when that division truncated. So adet * bdet * task has to divide the communicator evenly. Nothing about that rule changed; which communicator it applies to did. A config sizing the grid from MPI.COMM_WORLD was only ever shorthand for "every process doing this solve", and the world stopped being that set. With 6 processes and 3 subspaces each group holds 2, so a grid asking for 6 no longer divides it -- while 2, or the default 1, would. A caller cannot size the grid to the group anyway, not knowing how the processes were divided, and one config has to serve both rounds. The new tests therefore leave the grid at its default instead of using the _sbd_config helper, which sizes it from the world. Tested on the CPU backend under OpenMPI. _split_for_batches is arithmetic over (size, rank, num_batches), so the whole decision table -- group sizes, leader selection, idle ranks, contiguity -- is checked in a single process for counts up to eleven, including the uneven divisions. Two MPI tests then run the real split: one diagonalizes several subspaces both divided and one at a time over every process and requires the energies to agree, and one runs the two-round trim shape repeatedly in one process lifetime, which is what would expose a leaked communicator or a dump read back by the wrong round. The suite passes at one through seven processes. Not yet exercised: the GPU backends, and more than one node. Assisted-by: Claude Opus 5 --- python/sbd_solver.py | 176 +++++++++- ...atch-diagonalization-7a3e91c48d02f6b5.yaml | 31 ++ test/test_sqd_integration.py | 331 ++++++++++++++++++ 3 files changed, 520 insertions(+), 18 deletions(-) create mode 100644 releasenotes/notes/concurrent-batch-diagonalization-7a3e91c48d02f6b5.yaml diff --git a/python/sbd_solver.py b/python/sbd_solver.py index e8b4508..afdf5dd 100644 --- a/python/sbd_solver.py +++ b/python/sbd_solver.py @@ -157,12 +157,19 @@ def _solve_sci_core( fcidump, device_config=None, ecore_offset: float = 0.0, + group_id: int = 0, ) -> SCIResult: """ Inner diagonalization kernel that operates on a pre-loaded FCIDUMP object. Separated from solve_sci so that solve_sci_batch can write and load the FCIDUMP only once and reuse it across all batches. + + ``mpi_comm`` is the communicator whose processes perform this one + diagonalization, which is a subset of the world when the caller has divided the + processes among several batches. ``mpi_rank`` is the rank within it, so the + result is returned by the rank-0 process *of that group*. ``group_id`` + distinguishes the groups' wavefunction dumps, which share a directory. """ strings_a, strings_b = ci_strings @@ -193,7 +200,15 @@ def _solve_sci_core( # Use .bin extension to trigger SBD's fast binary write path # (SaveMatrixFormWF in restart.h checks extension: .bin -> raw doubles) - wf_dump_file = sbd_dir / "wavefunction.bin" + # + # The name carries the group id because concurrent groups share sbd_dir: two + # groups writing the same path would overwrite each other's amplitudes, and each + # would then read whichever write happened to land last. Every member of a group + # agrees on the id, and distinct groups have distinct ids, so no coordination is + # needed to keep the names apart. Under ``skip_wf`` nothing is written at all, so + # the collision cannot arise there; the name still carries the id, since whether + # the dump is requested is not this name's concern. + wf_dump_file = sbd_dir / f"wavefunction_{group_id}.bin" if not skip_wf: sbd_data.dump_matrix_form_wf = str(wf_dump_file) @@ -530,6 +545,26 @@ def solve_sci_batch( The FCIDUMP file is loaded once and reused across all batches. + When there is more than one subspace and at least as many processes as + subspaces, the processes are divided into one group per subspace and the + subspaces are diagonalized concurrently, each over its own group. Otherwise + every subspace is diagonalized in turn over all of the processes. Either way the + returned list is ordered by subspace, and the results are the same up to the + floating-point differences that come of using a different number of processes for + a diagonalization. + + Dividing the processes requires this function to be called collectively, with the + same number of subspaces on every process. Any remainder is left idle rather than + making the groups uneven: with 10 processes and 3 subspaces, three groups of + three are formed and one process takes no part in the diagonalizations. + + This is what a two-round schedule such as ``qiskit_addon_sqd.trim.TrimPolicy`` + asks for without having to configure anything. Its screening round hands over + every batch in one call, which is divided into groups here; its merged round + hands over a single subspace, which is given every process. ``SubspacePolicy`` + describes the schedule alone, and how the processes are spread over it is the + solver's concern. + Args: ci_strings: List of (strings_a, strings_b) pairs. one_body_tensor: The one-body tensor of the Hamiltonian. @@ -563,39 +598,144 @@ def solve_sci_batch( sbd_dir, owns_sbd_dir = _make_sbd_dir(mpi_comm, mpi_rank, temp_dir) + # Divide the processes into one group per batch, so that the batches are + # diagonalized concurrently rather than each in turn over every process. The + # FCIDUMP is loaded before the split, over the whole communicator, so that it is + # read once rather than once per group. + divided, group_comm, group_id, leaders_comm = _split_for_batches(mpi_comm, len(ci_strings)) + try: fcidump, ecore_offset = _load_or_regenerate_fcidump( backend, mpi_rank, mpi_comm, sbd_dir, fcidump_path, one_body_tensor, two_body_tensor, norb, nelec, ) - return [ - _solve_sci_core( - ci_strs, - norb=norb, - nelec=nelec, - spin_sq=spin_sq, - mpi_comm=mpi_comm, - mpi_rank=mpi_rank, - sbd_config=sbd_config, - sbd_dir=sbd_dir, - backend=backend, - fcidump=fcidump, - device_config=device_config, - ecore_offset=ecore_offset, + # Each group owns one batch when the processes were divided; otherwise this + # process owns them all and runs them one after another. A process left + # without a group owns none, and only takes part in the exchange below. + if not divided: + owned = list(enumerate(ci_strings)) + elif group_comm is None: + owned = [] + else: + owned = [(group_id, ci_strings[group_id])] + + local = [ + ( + index, + _solve_sci_core( + ci_strs, + norb=norb, + nelec=nelec, + spin_sq=spin_sq, + mpi_comm=group_comm, + mpi_rank=group_comm.Get_rank(), + sbd_config=sbd_config, + sbd_dir=sbd_dir, + backend=backend, + fcidump=fcidump, + device_config=device_config, + ecore_offset=ecore_offset, + group_id=index, + ), ) - for ci_strs in ci_strings + for index, ci_strs in owned ] + + if not divided: + return [result for _, result in local] + + # Each group's result exists on its own leader, so the leaders exchange them + # and then every process is given the collected list. The broadcast is over + # the whole communicator rather than over a group, so that a process left + # without a group receives the results too, and so that the leaders' exchange + # is not repeated once per group. The order follows the batch index rather + # than the order in which the groups finished. + if leaders_comm is not None: + gathered = [pair for chunk in leaders_comm.allgather(local) for pair in chunk] + else: + gathered = None + gathered = mpi_comm.bcast(gathered, root=0) + return [result for _, result in sorted(gathered, key=lambda pair: pair[0])] finally: + # Communicators are a finite resource, so the ones created here are released + # rather than left for the garbage collector. + if divided: + if group_comm is not None: + group_comm.Free() + if leaders_comm is not None: + leaders_comm.Free() if clean_temp_dir and owns_sbd_dir and mpi_rank == 0: shutil.rmtree(sbd_dir, ignore_errors=True) +def _split_for_batches(mpi_comm, num_batches): + """Divide a communicator into one group per batch. + + Returns ``(divided, group_comm, group_id, leaders_comm)``: + + - ``divided`` is whether the processes were divided at all. When they were not -- + because there is a single batch, or fewer processes than batches -- the other + values are ``(mpi_comm, 0, None)`` and the caller diagonalizes the batches one + after another over the whole communicator. + - ``group_comm`` spans the processes assigned to this process's batch, and + ``group_id`` is that batch's index. Both are ``None`` on a process that was + left without a group, which happens when the process count is not a multiple of + the batch count. + - ``leaders_comm`` spans the rank-0 process of every group, and is ``None`` on + every other process. The rank-0 process of ``mpi_comm`` always belongs to it, + being the leader of the first group, so results collected here can be broadcast + from that rank. + + Splitting is collective: every process of ``mpi_comm`` must call this with the + same ``num_batches``, which holds because the caller receives the same subspaces + on every process. Note that a process without a group still enters both splits, + since a process that skipped them would leave the others waiting. + """ + size = mpi_comm.Get_size() + rank = mpi_comm.Get_rank() + if num_batches < 2 or size < num_batches: + return False, mpi_comm, 0, None + + # Contiguous groups of equal size, so that every batch is diagonalized with the + # same number of processes. When the process count is not a multiple of the batch + # count, the leftover processes at the end are given no group rather than making + # one group larger than the others. + # + # Equal sizes matter beyond tidiness: a schedule such as + # ``qiskit_addon_sqd.trim.TrimPolicy`` diagonalizes these batches in order to rank + # them against each other and keep the highest-weight configurations of each. A + # batch given more processes than its siblings is solved to a different + # floating-point result, so the comparison would turn partly on how the processes + # happened to divide. + group_size = size // num_batches + group_id = rank // group_size if rank < group_size * num_batches else None + group_comm = mpi_comm.Split( + color=MPI.UNDEFINED if group_id is None else group_id, + key=rank, + ) + if group_comm == MPI.COMM_NULL: + # This process takes no part in the diagonalizations. It still has to enter + # the second split below, and the collectives that follow, so that the + # processes that do are not left waiting on it. + group_comm = None + is_leader = group_comm is not None and group_comm.Get_rank() == 0 + leaders_comm = mpi_comm.Split( + color=0 if is_leader else MPI.UNDEFINED, + key=group_id if is_leader else 0, + ) + if leaders_comm == MPI.COMM_NULL: + leaders_comm = None + return True, group_comm, group_id, leaders_comm + + def _make_sbd_dir(mpi_comm, mpi_rank, temp_dir): """Create a per-run tempdir on rank 0 and broadcast its path. - Used to hold the wavefunction.bin written by rank 0 (and the - regenerated fcidump.txt when ``fcidump_path`` is not provided). + Used to hold the wavefunction dumps written by each group's rank-0 process + (and the regenerated fcidump.txt when ``fcidump_path`` is not provided). The + directory is shared by every group, so the dumps are named per group; see + ``_solve_sci_core``. Returns (path, owns_sbd_dir) where owns_sbd_dir is True on rank 0 (so the caller knows to rmtree it on exit). diff --git a/releasenotes/notes/concurrent-batch-diagonalization-7a3e91c48d02f6b5.yaml b/releasenotes/notes/concurrent-batch-diagonalization-7a3e91c48d02f6b5.yaml new file mode 100644 index 0000000..6037a65 --- /dev/null +++ b/releasenotes/notes/concurrent-batch-diagonalization-7a3e91c48d02f6b5.yaml @@ -0,0 +1,31 @@ +--- +features: + - | + :func:`~sbd.sbd_solver.solve_sci_batch` now divides the processes into one group + per subspace and diagonalizes the subspaces concurrently, instead of giving every + subspace all of the processes in turn. The division happens when there is more than + one subspace and at least as many processes as subspaces; otherwise the previous + behavior is unchanged. The returned list is ordered by subspace either way. + + The FCIDUMP is still read once over the whole communicator, before the processes + are divided, rather than once per group. Any remainder is left idle instead of + making the groups uneven, so that every subspace is diagonalized with the same + number of processes: with 10 processes and 3 subspaces, three groups of three are + formed and one process takes no part in the diagonalizations. + + This is what a two-round schedule such as ``qiskit_addon_sqd.trim.TrimPolicy`` + needs, and it requires no configuration. That schedule's screening round hands over + every batch in a single call, which is divided here, and its merged round hands over + one subspace, which is given every process. Dividing the processes is the solver's + concern rather than the schedule's, as described under + `diagonalize_fermionic_hamiltonian `__. + + One thing to be aware of when setting the process grid. SBD's + ``adet_comm_size``, ``bdet_comm_size`` and ``task_comm_size`` describe a single + diagonalization, so their product has to divide the communicator that + diagonalization is given -- which, once the processes are divided, is a group + rather than all of them. An ``sbd_config`` that sizes the grid from + ``MPI.COMM_WORLD`` therefore no longer fits: with 6 processes and 3 subspaces each + group holds 2, and a grid asking for 6 does not divide it. Leaving the grid at its + default works at any size, since the derived helper dimension absorbs whatever the + group turns out to be. diff --git a/test/test_sqd_integration.py b/test/test_sqd_integration.py index 19bc00a..0587e38 100644 --- a/test/test_sqd_integration.py +++ b/test/test_sqd_integration.py @@ -522,3 +522,334 @@ def test_sbd_carryover_is_weight_ordered_mpi( _check_sbd_carryover_is_weight_ordered( data_dir, device_config, tmp_path, carryover_type ) + + +# --- dividing the processes among the batches ---------------------------------------- +# +# ``_split_for_batches`` decides which batch each process works on, and its decision is +# arithmetic over ``(size, rank, num_batches)`` alone. Those can be supplied, so the +# whole decision table is checkable in a single process, for process counts larger than +# a test run would ever launch. ``MPI.Split`` is what the arithmetic is handed to; the +# tests below stand in for it so that what is under test is the division rather than +# mpi4py. +# +# The MPI tests further down then run the real thing, which is what confirms the +# arithmetic was handed to Split correctly. + + +class _FakeComm: + """Enough of a communicator to drive ``_split_for_batches`` in one process. + + ``Split`` records the colors it is given rather than forming a communicator, and + returns a stand-in whose rank is this process's position among the ranks sharing its + color -- which is what ``Split`` guarantees for the ascending keys the caller passes. + """ + + def __init__(self, size, rank, num_batches=1): + self._size = size + self._rank = rank + self._num_batches = num_batches + self.colors = [] + + def Get_size(self): + return self._size + + def Get_rank(self): + return self._rank + + def Split(self, color, key): # pylint: disable=unused-argument + from mpi4py import MPI + + self.colors.append(color) + if color == MPI.UNDEFINED: + return MPI.COMM_NULL + # Split orders a new communicator by ascending key, and the caller passes the + # rank as the key, so the lowest-ranked member of a color becomes its rank 0. + # The members of this color are known from the division rule itself. + group_size = self._size // self._num_batches + first = color * group_size + return _FakeComm(group_size, self._rank - first, num_batches=self._num_batches) + + +def _divide(size, num_batches): + """Run ``_split_for_batches`` for every rank of a ``size``-process communicator. + + Returns a list of ``(divided, group_id, is_leader)``, one entry per rank, with the + group membership recovered from the colors passed to ``Split`` rather than from a + real communicator. + """ + from mpi4py import MPI + + from sbd.sbd_solver import _split_for_batches + + table = [] + for rank in range(size): + comm = _FakeComm(size, rank, num_batches=num_batches) + divided, _, group_id, _ = _split_for_batches(comm, num_batches) + if not divided: + table.append((False, group_id, None)) + continue + group_color, leader_color = comm.colors + in_group = group_color != MPI.UNDEFINED + table.append((True, group_id, in_group and leader_color != MPI.UNDEFINED)) + return table + + +@pytest.mark.parametrize("num_batches", [1, 2, 5]) +def test_single_process_is_never_divided(num_batches): + """One process runs the batches in turn, whatever their number.""" + (entry,) = _divide(1, num_batches) + assert entry == (False, 0, None) + + +@pytest.mark.parametrize("size", [1, 2, 8]) +def test_one_batch_is_never_divided(size): + """A lone subspace is given every process, which is the merged round of trim SQD.""" + assert all(divided is False for divided, _, _ in _divide(size, 1)) + + +@pytest.mark.parametrize(("size", "num_batches"), [(2, 3), (3, 4), (1, 2)]) +def test_fewer_processes_than_batches_is_not_divided(size, num_batches): + """With too few processes to go around, the batches run in turn over all of them.""" + assert all(divided is False for divided, _, _ in _divide(size, num_batches)) + + +@pytest.mark.parametrize( + ("size", "num_batches", "expected_group_sizes"), + [ + (4, 2, [2, 2]), + (6, 3, [2, 2, 2]), + (8, 2, [4, 4]), + (3, 3, [1, 1, 1]), + (9, 3, [3, 3, 3]), + ], +) +def test_equal_division_assigns_every_process(size, num_batches, expected_group_sizes): + """When the division is exact, every process joins a group and the groups match.""" + table = _divide(size, num_batches) + assert all(divided for divided, _, _ in table) + assert [group for _, group, _ in table].count(None) == 0 + for batch, expected in enumerate(expected_group_sizes): + assert sum(1 for _, group, _ in table if group == batch) == expected + + +@pytest.mark.parametrize( + ("size", "num_batches", "expected_idle"), + [ + (10, 3, 1), + (5, 2, 1), + (7, 3, 1), + (11, 3, 2), + (7, 2, 1), + ], +) +def test_leftover_processes_are_left_idle(size, num_batches, expected_idle): + """A remainder is left out rather than making one group larger than its siblings. + + Equal groups are what makes the batches comparable, which is the point of a + screening round: a batch solved over more processes reaches a slightly different + floating-point energy, so the ranking would depend on how the processes divided. + """ + table = _divide(size, num_batches) + assert all(divided for divided, _, _ in table) + + idle = [rank for rank, (_, group, _) in enumerate(table) if group is None] + assert len(idle) == expected_idle + # The idle processes are the ones at the end, so a group is always contiguous. + assert idle == list(range(size - expected_idle, size)) + + group_sizes = [ + sum(1 for _, group, _ in table if group == batch) for batch in range(num_batches) + ] + assert group_sizes == [size // num_batches] * num_batches + + +@pytest.mark.parametrize(("size", "num_batches"), [(4, 2), (6, 3), (10, 3), (9, 3)]) +def test_every_group_has_exactly_one_leader(size, num_batches): + """Each group contributes one process to the leaders' exchange, and rank 0 is one. + + The results are broadcast from rank 0 of the whole communicator, so it has to be a + leader; it is, being the first process of the first group. + """ + table = _divide(size, num_batches) + leaders = [rank for rank, (_, _, is_leader) in enumerate(table) if is_leader] + assert len(leaders) == num_batches + assert 0 in leaders + # One leader per group, and each leads a different one. + assert sorted(table[rank][1] for rank in leaders) == list(range(num_batches)) + + +# --- the grouped path against the sequential one ------------------------------------- +# +# The division is only correct if it does not change the answer, so the tests below +# diagonalize the same several subspaces twice over: once letting the processes divide, +# and once one subspace at a time over every process, which is the path that predates +# the division. The energies have to agree. +# +# They agree to a tolerance rather than exactly. A subspace solved over a group of two +# processes and the same subspace solved over all four sum their contributions in a +# different order, so the Davidson iterations differ in the last bits. SOLVER_CONFIG +# converges to 1e-10, well inside the 1e-8 asserted here. +# +# These tests pass SOLVER_CONFIG rather than going through ``_sbd_config``, which sizes +# the alpha dimension from ``MPI.COMM_WORLD``. That is right for an undivided call, where +# the world is the set of processes performing the diagonalization, and wrong here, where +# a group is: SBD's grid describes one diagonalization, so it is relative to whichever +# communicator that diagonalization is handed. +# +# The rule itself is just divisibility. ``diag()`` derives the helper dimension by +# integer division, ``h_comm_size = mpi_size / (task_comm_size * base_comm_size)``, and +# ``TaskCommunicator`` then checks that multiplying the four back recovers ``mpi_size`` +# (chemistry/tpb/sbdiag.h and chemistry/tpb/helper.h) -- which fails exactly when the +# division truncated. So ``adet * bdet * task`` has to divide the communicator evenly. +# ``adet_comm_size=2`` would be fine on groups of two; it is 6, taken from a world the +# groups are no longer the same size as, that is not. +# +# Leaving the grid at its 1x1x1 default avoids having to know: the helper dimension +# absorbs however many processes the group turns out to have. A caller cannot size the +# grid to the group anyway, not knowing how the processes were divided, and one config +# has to serve both of trim SQD's rounds, whose communicators differ in size by +# construction. + + +def _subspaces_for_batches(data_dir, num_batches): + """``num_batches`` distinct subspaces drawn from the h2o selection. + + Each takes a different slice of the alpha determinants, so the subspaces differ and + a result mistakenly carried from the wrong group would show up as a wrong energy. + Every slice starts at the Hartree-Fock determinant, the first in the file, so each + subspace is a sensible one to diagonalize rather than an arbitrary set. + """ + strings = _read_alpha_determinants(data_dir / "h2o" / "h2o-1em3-alpha.txt") + subspaces = [] + for batch in range(num_batches): + taken = np.concatenate([strings[:1], strings[1 + batch : 24 + batch]]) + subspaces.append((taken, taken)) + return subspaces + + +def _check_grouped_matches_sequential(data_dir, device_config, num_batches): + from mpi4py import MPI + + from sbd.sbd_solver import solve_sci_batch + + comm = MPI.COMM_WORLD + hcore, eri, _ = _load_hamiltonian(data_dir) + subspaces = _subspaces_for_batches(data_dir, num_batches) + + def diagonalize(batch): + return solve_sci_batch( + batch, + hcore, + eri, + norb=NORB, + nelec=NELEC, + sbd_config=SOLVER_CONFIG, + device_config=device_config, + ) + + # All the subspaces in one call: divided into groups when there are enough + # processes, and run in turn when there are not. + grouped = diagonalize(subspaces) + assert len(grouped) == num_batches + + # One subspace per call, so every call is given the whole communicator. This is + # what the division has to reproduce. + sequential = [diagonalize([subspace])[0] for subspace in subspaces] + + if comm.Get_rank() != 0: + return + + # Ordered by subspace, not by whichever group finished first. The subspaces differ, + # so a misordered or misattributed result fails here. + for from_group, from_sequence in zip(grouped, sequential): + assert from_group.energy == pytest.approx(from_sequence.energy, abs=1e-8) + _assert_result_is_consistent(from_group) + + # The subspaces are distinct, so their energies should be too -- without this, the + # comparison above would also pass if every group had solved the same subspace. + energies = [result.energy for result in grouped] + assert len(set(energies)) == num_batches + + +@pytest.mark.parametrize("num_batches", [2, 3]) +def test_grouped_matches_sequential_standalone(data_dir, device_config, num_batches): + """In one process nothing is divided, and the two paths are the same code.""" + _check_grouped_matches_sequential(data_dir, device_config, num_batches) + + +@pytest.mark.mpi +@pytest.mark.parametrize("num_batches", [2, 3]) +def test_grouped_matches_sequential_mpi(data_dir, device_config, num_batches): + """Dividing the launched ranks among the subspaces gives the same energies. + + Whether the division actually happens depends on the process count the suite was + launched with: at or above ``num_batches`` processes it does, below that the call + falls back to running them in turn. Both are worth exercising, and which one runs + is reported by ``tox -e mpi``'s header rather than asserted here. + """ + _check_grouped_matches_sequential(data_dir, device_config, num_batches) + + +def _check_two_phase_rounds(data_dir, device_config): + """The trim SQD shape: a divided screening round, then an undivided merged one. + + Both rounds happen in one process lifetime, as they do inside + ``diagonalize_fermionic_hamiltonian``. That is what makes this more than the sum of + the two cases above: the screening round creates communicators and per-group + wavefunction dumps, and the merged round that follows must not inherit either. A + communicator left unfreed would eventually exhaust the supply over many iterations, + and a dump left behind under a name the merged round reuses would be read back as + if it were the merged round's own amplitudes. + """ + from mpi4py import MPI + + from sbd.sbd_solver import solve_sci_batch + + comm = MPI.COMM_WORLD + hcore, eri, _ = _load_hamiltonian(data_dir) + + def diagonalize(batch): + return solve_sci_batch( + batch, + hcore, + eri, + norb=NORB, + nelec=NELEC, + sbd_config=SOLVER_CONFIG, + device_config=device_config, + ) + + # Several iterations, so that a communicator leaked once per round would accumulate + # rather than merely occur. + for _ in range(3): + screened = diagonalize(_subspaces_for_batches(data_dir, 3)) + assert len(screened) == 3 + + # The merged round: one subspace built from the screening round, given every + # process. Merging the inputs keeps this independent of what the solver chose + # to carry over, which is a separate concern tested above. + screened_alpha = [a for a, _ in _subspaces_for_batches(data_dir, 3)] + merged_a = np.unique(np.concatenate(screened_alpha)) + (merged,) = diagonalize([(merged_a, merged_a)]) + + if comm.Get_rank() != 0: + continue + + _assert_result_is_consistent(merged) + # The merged subspace contains each screened subspace, so its energy is at or + # below every one of theirs. This would fail if the merged round had read back + # a screening round's wavefunction dump instead of its own. + for result in screened: + assert merged.energy <= result.energy + 1e-8 + + +def test_two_phase_rounds_standalone(data_dir, device_config): + """The two-round shape runs in a single process.""" + _check_two_phase_rounds(data_dir, device_config) + + +@pytest.mark.mpi +def test_two_phase_rounds_mpi(data_dir, device_config): + """The screening round divides the launched ranks; the merged round gets them all.""" + _check_two_phase_rounds(data_dir, device_config)