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)