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
6 changes: 6 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,10 @@ releases may include breaking changes.

### Fixed

- avoided process startup for single-worker simulations ([#624])
([**@aaronleesander**])
- cleaned up tests and fixed MCWF RNG continuity ([#623])
([**@aaronleesander**])
- fixed shot readout progress bar suppression ([#622])
([**@aaronleesander**])
- fixed noise characterization iteration limits and initial loss ([#621])
Expand Down Expand Up @@ -294,6 +298,8 @@ for previous changelogs._

<!-- PR links -->

[#624]: https://github.com/munich-quantum-toolkit/yaqs/pull/624
[#623]: https://github.com/munich-quantum-toolkit/yaqs/pull/623
[#622]: https://github.com/munich-quantum-toolkit/yaqs/pull/622
[#621]: https://github.com/munich-quantum-toolkit/yaqs/pull/621
[#620]: https://github.com/munich-quantum-toolkit/yaqs/pull/620
Expand Down
21 changes: 13 additions & 8 deletions docs/examples/simulator_initialization.md
Original file line number Diff line number Diff line change
Expand Up @@ -172,14 +172,19 @@ serial runs.

## `mp_context`: multiprocessing start method

`mp_context` controls how worker processes are spawned. The default `"auto"`
picks the best option per OS:

| Value | Behaviour |
| --------- | --------------------------------------------------------------------------------------------------------------- |
| `"auto"` | `"fork"` on Linux, `"spawn"` everywhere else. |
| `"fork"` | Fastest worker startup; reuses Python state from the parent. Safe in YAQS because BLAS/OpenMP pools are capped. |
| `"spawn"` | Fresh interpreter per worker. Required on Windows/macOS; slower startup but more isolated. |
`mp_context` controls how worker processes start. The default `"auto"` selects a
start method per OS:

| Value | Behaviour |
| --------- | --------------------------------------------------------------------------- |
| `"auto"` | `"forkserver"` on Linux, `"spawn"` everywhere else. |
| `"fork"` | Copies the parent process. Avoid when the parent has active threads. |
| `"spawn"` | Fresh interpreter per worker. Used on Windows/macOS and available on Linux. |

On Linux, `"auto"` creates workers through a separate server process. The first
pool has extra startup cost, but workers do not fork the application's threaded
process. Numerical thread limits control resource use; they do not make an
explicit `"fork"` safe when the parent has active threads.

```{code-cell} ipython3
automatic_context = Simulator(mp_context="auto")
Expand Down
4 changes: 2 additions & 2 deletions src/mqt/yaqs/core/parallel_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -99,15 +99,15 @@ def get_parallel_context(mp_context: MPContext = "auto") -> multiprocessing.cont
"""Return a multiprocessing context for worker processes.

Args:
mp_context: Start method selector. ``"auto"`` uses ``"fork"`` on Linux and
mp_context: Start method selector. ``"auto"`` uses ``"forkserver"`` on Linux and
``"spawn"`` elsewhere; ``"fork"`` or ``"spawn"`` select that method explicitly.

Returns:
A :class:`~multiprocessing.context.BaseContext` for creating worker processes.
"""
if mp_context == "auto":
if sys.platform == "linux":
return multiprocessing.get_context("fork")
return multiprocessing.get_context("forkserver")
return multiprocessing.get_context("spawn")
return multiprocessing.get_context(mp_context)

Expand Down
49 changes: 25 additions & 24 deletions src/mqt/yaqs/simulator.py
Original file line number Diff line number Diff line change
Expand Up @@ -1081,22 +1081,23 @@ class Simulator:
A :class:`Simulator` owns the execution-side configuration: how trajectories
are parallelized, how many workers to use, whether to display a progress bar,
which multiprocessing context to use, and the retry policy for transient
worker errors. The physics inputs (initial state, operator, simulation
errors in process pool workers. The physics inputs (initial state, operator, simulation
parameters, optional noise model) are passed per call to :meth:`run`.

Multiple :meth:`run` calls share the same configuration. Each call constructs
its own short-lived process pool when ``parallel=True``; the pool is not
persisted across runs in the current implementation.
Multiple :meth:`run` calls share the same configuration. A call constructs
a short-lived process pool when ``parallel=True`` and both the job count
and ``max_workers`` exceed one. Other calls run in the current process.
A worker cap of one with ``parallel=True`` also limits numerical threads to one.

Attributes:
parallel: Whether to execute trajectories in parallel via a process pool.
parallel: Whether to use a process pool for multiple jobs and workers.
max_workers: Maximum number of worker processes when ``parallel=True``.
Defaults to ``max(1, available_cpus() - 1)``.
show_progress: Whether to display trajectory and shot-readout progress bars.
mp_context: Multiprocessing context: ``"auto"`` (default), ``"fork"``,
or ``"spawn"``. ``"auto"`` selects ``"fork"`` on Linux and ``"spawn"`` elsewhere.
max_retries: Maximum retry attempts for transient worker errors.
retry_exceptions: Exception types that trigger a retry.
or ``"spawn"``. ``"auto"`` selects ``"forkserver"`` on Linux and ``"spawn"`` elsewhere.
max_retries: Maximum retry attempts for transient errors in process pool workers.
retry_exceptions: Exception types that trigger a retry in a process pool.
"""

def __init__(
Expand All @@ -1112,13 +1113,13 @@ def __init__(
"""Initialize the simulator with execution-side configuration.

Args:
parallel: Boolean that enables a process pool for multi-trajectory runs.
parallel: Enable a process pool when both the job count and worker cap exceed one.
max_workers: Positive worker-process cap. ``None`` (default) resolves to
``max(1, available_cpus() - 1)``.
``max(1, available_cpus() - 1)``. A cap of one runs in the current process.
show_progress: Whether to display trajectory and shot-readout progress bars.
mp_context: Multiprocessing start method (``"auto"``, ``"fork"``, or ``"spawn"``).
max_retries: Non-negative maximum retries for transient worker errors.
retry_exceptions: Exception types that trigger a retry.
max_retries: Non-negative maximum retries for transient errors in process pool workers.
retry_exceptions: Exception types that trigger a retry in a process pool.
"""
self._execution = ExecutionConfig(
parallel=parallel,
Expand All @@ -1131,7 +1132,7 @@ def __init__(

@property
def parallel(self) -> bool:
"""Whether parallel execution is enabled."""
"""Whether process pools are enabled for multiple jobs and workers."""
return self._execution.parallel

@parallel.setter
Expand All @@ -1140,7 +1141,7 @@ def parallel(self, value: bool) -> None:

@property
def max_workers(self) -> int:
"""Effective worker count for parallel execution."""
"""Effective worker cap; one keeps execution in the current process."""
return self._execution.resolved_max_workers()

@max_workers.setter
Expand All @@ -1167,7 +1168,7 @@ def mp_context(self, value: MPContext) -> None:

@property
def max_retries(self) -> int:
"""Maximum retries per job in parallel execution."""
"""Maximum retries per job in a process pool."""
return self._execution.max_retries

@max_retries.setter
Expand All @@ -1176,7 +1177,7 @@ def max_retries(self, value: int) -> None:

@property
def retry_exceptions(self) -> tuple[type[BaseException], ...]:
"""Exception types that trigger a parallel job retry."""
"""Exception types that trigger a retry in a process pool."""
return self._execution.retry_exceptions

@retry_exceptions.setter
Expand Down Expand Up @@ -1403,7 +1404,7 @@ def consume(traj_index: int, trajectory: _ProgramTrajectory) -> None:
if trajectory_final is not None:
final_mps = trajectory_final

if self.parallel and effective_num_traj > 1:
if self.parallel and effective_num_traj > 1 and self.max_workers > 1:
for traj_index, trajectory in run_backend_parallel(
worker_fn=_program_worker,
payload=payload,
Expand All @@ -1428,7 +1429,7 @@ def consume(traj_index: int, trajectory: _ProgramTrajectory) -> None:
_program_worker,
traj_index,
payload,
n_threads=available_cpus(),
n_threads=1 if self.parallel and self.max_workers == 1 else available_cpus(),
)
consume(traj_index, trajectory)

Expand Down Expand Up @@ -1638,7 +1639,7 @@ def _run_analog(
final_psi: np.ndarray | None = None
final_rho: np.ndarray | None = None

if self.parallel and effective_num_traj > 1:
if self.parallel and effective_num_traj > 1 and self.max_workers > 1:
for i, traj_payload in run_backend_parallel(
worker_fn=worker_fn,
payload=payload,
Expand All @@ -1662,7 +1663,7 @@ def _run_analog(
else:
final_mps = cast("MPS", traj_final)
else:
n_threads = available_cpus()
n_threads = 1 if self.parallel and self.max_workers == 1 else available_cpus()

args: list[Any]
if state_rep == "vector":
Expand Down Expand Up @@ -1820,7 +1821,7 @@ def _consume(
final_mps = traj_final

try:
if self.parallel and effective_num_traj > 1:
if self.parallel and effective_num_traj > 1 and self.max_workers > 1:
for i, traj_payload in run_backend_parallel(
worker_fn=_digital_worker,
payload=payload,
Expand All @@ -1835,7 +1836,7 @@ def _consume(
traj_data, traj_diag, shot_counts, traj_final = traj_payload
_consume(i, traj_data, traj_diag, shot_counts, traj_final)
else:
n_threads = available_cpus()
n_threads = 1 if self.parallel and self.max_workers == 1 else available_cpus()
iterator = tqdm(
range(effective_num_traj),
desc="Running trajectories",
Expand Down Expand Up @@ -1977,7 +1978,7 @@ def _run_ensemble(
"operator": operator,
}

if self.parallel and len(initial_states) > 1:
if self.parallel and len(initial_states) > 1 and self.max_workers > 1:
for i, (obs_result, traj_diag, multi_time_result) in run_backend_parallel(
worker_fn=_ensemble_worker,
payload=payload,
Expand All @@ -1995,7 +1996,7 @@ def _run_ensemble(
assert multi_time_result is not None
multi_time_matrix[i] = multi_time_result
else:
n_threads = available_cpus()
n_threads = 1 if self.parallel and self.max_workers == 1 else available_cpus()
args = [(i, initial_states[i], worker_params, operator) for i in range(len(initial_states))]
iterator = tqdm(args, desc="Running unitary ensemble", ncols=80, disable=not self.show_progress)
for i, arg in enumerate(iterator):
Expand Down
4 changes: 3 additions & 1 deletion tests/analog/test_ensemble.py
Original file line number Diff line number Diff line change
Expand Up @@ -258,7 +258,9 @@ def test_list_mps_unitary_ensemble_parallel_worker_path() -> None:
dt=0.05,
multi_time_observables=[(z0, z0), (z0, z1)],
)
result = Simulator(parallel=True, show_progress=False).run(states, hamiltonian, sim_params, noise_model=None)
result = Simulator(parallel=True, max_workers=2, show_progress=False).run(
states, hamiltonian, sim_params, noise_model=None
)
assert result.expectation_values[0] is not None
assert result.multi_time_results is not None

Expand Down
28 changes: 28 additions & 0 deletions tests/characterization/noise/optimization/test_trajectories.py
Original file line number Diff line number Diff line change
Expand Up @@ -66,6 +66,34 @@ def test_simulate_observable_trajectories_shape() -> None:
assert expectations.shape == (len(observables), len(times))


def test_seeded_observable_trajectories_match_in_serial_and_parallel() -> None:
"""Noise-fitting forward trajectories agree with automatic process-pool execution."""
hamiltonian, init_state, observables, _sim_params, noise = _three_site_problem()
params = AnalogSimParams(
observables=observables,
elapsed_time=0.2,
dt=0.1,
num_traj=8,
random_seed=42,
order=1,
)
results = [
simulate_observable_trajectories(
sim_params=params,
hamiltonian=hamiltonian,
init_state=init_state,
noise_model=noise,
observables=observables,
simulator=build_simulator(ExecutionConfig(parallel=parallel, max_workers=2, show_progress=False)),
representation="mps",
)
for parallel in (False, True)
]
serial, parallel = results
np.testing.assert_array_equal(parallel[1], serial[1])
np.testing.assert_allclose(parallel[0], serial[0], atol=1e-12)


def test_ref_expectations_path_matches_simulation() -> None:
"""Precomputed expectations are accepted when shapes match the fitting set."""
hamiltonian, init_state, observables, sim_params, reference_model = _three_site_problem()
Expand Down
6 changes: 4 additions & 2 deletions tests/core/data_structures/test_mps.py
Original file line number Diff line number Diff line change
Expand Up @@ -1892,10 +1892,12 @@ def test_inplace_measure() -> None:
assert np.isclose(psi.expect(Observable("x", 2)), 1.0 if outcome == 0 else -1.0)


def test_multi_shot() -> None:
@pytest.mark.parametrize("parallel", [False, True], ids=["serial", "parallel"])
def test_multi_shot(*, parallel: bool, monkeypatch: pytest.MonkeyPatch) -> None:
"""Full-chain readout of a basis state counts every shot at the expected integer."""
monkeypatch.setattr(mps_mod, "available_cpus", lambda: 3 if parallel else 1)
psi_mps = MPS(length=3, state="ones")
assert psi_mps.measure_shots(shots=10) == {7: 10}
assert psi_mps.measure_shots(shots=10, show_progress=False) == {7: 10}


def test_norm() -> None:
Expand Down
47 changes: 46 additions & 1 deletion tests/core/test_parallel_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
import multiprocessing
import os
import sys
import threading
from types import SimpleNamespace
from typing import Any, cast

Expand All @@ -35,6 +36,50 @@
from mqt.yaqs.simulator import available_cpus as simulator_available_cpus


def _indexed_pool_worker(job_idx: int, payload: dict[str, Any] | None = None) -> tuple[int, int]:
"""Return the indexed value and worker PID after both workers reach the barrier."""
context = resolve_worker_ctx(payload)
if "barrier" in context:
context["barrier"].wait(timeout=30)
return job_idx + context["offset"], os.getpid()


@pytest.mark.filterwarnings("error:This process.*is multi-threaded.*:DeprecationWarning")
def test_auto_pool_from_threaded_parent_preserves_results() -> None:
"""Automatic pools start safely beside a live thread and preserve indexed results."""
serial = run_indexed_jobs(
_indexed_pool_worker,
payload={"offset": 10},
n_jobs=2,
config=ExecutionConfig(parallel=False, show_progress=False),
desc="test",
)
stop = threading.Event()
thread = threading.Thread(target=stop.wait)
thread.start()
try:
assert thread.is_alive()
context = get_parallel_context()
pooled = run_indexed_jobs(
_indexed_pool_worker,
payload={"offset": 10, "barrier": context.Barrier(2)},
n_jobs=2,
config=ExecutionConfig(max_workers=2, show_progress=False, max_retries=0),
desc="test",
)
finally:
stop.set()
thread.join(timeout=5)

assert not thread.is_alive()
assert {index: value for index, (value, _pid) in pooled.items()} == {
index: value for index, (value, _pid) in serial.items()
}
worker_pids = {pid for _value, pid in pooled.values()}
assert len(worker_pids) == 2
assert os.getpid() not in worker_pids


def test_available_cpus_without_slurm(monkeypatch: pytest.MonkeyPatch) -> None:
"""Without overrides, ``available_cpus`` falls back to affinity or ``cpu_count``."""
monkeypatch.delenv("YAQS_MAX_WORKERS", raising=False)
Expand Down Expand Up @@ -96,7 +141,7 @@ def test_threading_config() -> None:
"""Verify correct multiprocessing context and Numba threading configuration."""
ctx = get_parallel_context()
if sys.platform == "linux":
assert ctx.get_start_method() == "fork"
assert ctx.get_start_method() == "forkserver"
else:
assert ctx.get_start_method() == "spawn"

Expand Down
5 changes: 3 additions & 2 deletions tests/test_equivalence_checker.py
Original file line number Diff line number Diff line change
Expand Up @@ -1408,7 +1408,8 @@ def test_ensemble_trajectory_worker_uses_initialized_context() -> None:
assert result["fidelity"] == pytest.approx(1.0, abs=1e-12)


def test_seeded_serial_and_process_pool_ensembles_agree() -> None:
@pytest.mark.parametrize("mp_context", ["auto", "spawn"])
def test_seeded_serial_and_process_pool_ensembles_agree(mp_context: Literal["auto", "spawn"]) -> None:
"""Serial and process-pool workers return the same seeded MPO ensemble."""
qc = QuantumCircuit(2)
qc.h(0)
Expand All @@ -1417,7 +1418,7 @@ def test_seeded_serial_and_process_pool_ensembles_agree() -> None:
kwargs = {"noise_model": noise, "num_traj": 6, "random_seed": 0, "return_trajectories": True}

serial = EquivalenceChecker(representation="mpo", parallel=False).check(qc, qc, **kwargs)
pooled = EquivalenceChecker(representation="mpo", parallel=True, max_workers=2, mp_context="spawn").check(
pooled = EquivalenceChecker(representation="mpo", parallel=True, max_workers=2, mp_context=mp_context).check(
qc, qc, **kwargs
)
serial_fidelities = [traj["fidelity"] for traj in serial["trajectories"]]
Expand Down
Loading
Loading