Skip to content
Draft
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
176 changes: 158 additions & 18 deletions python/sbd_solver.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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).

Expand Down
Original file line number Diff line number Diff line change
@@ -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 <https://qiskit.github.io/qiskit-addon-sqd/apidocs/qiskit_addon_sqd.fermion.html>`__.

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.
Loading
Loading