diff --git a/crates/ppvm-python-native/src/interface_tableau.rs b/crates/ppvm-python-native/src/interface_tableau.rs index 9bbcda6b7..8c7f50044 100644 --- a/crates/ppvm-python-native/src/interface_tableau.rs +++ b/crates/ppvm-python-native/src/interface_tableau.rs @@ -5,7 +5,7 @@ use bnum::types::{U256, U512, U1024, U2048}; use paste::paste; use ppvm_tableau::prelude::*; use pyo3::prelude::*; -use pyo3::types::{PyComplex, PyDict}; +use pyo3::types::{PyByteArray, PyComplex, PyDict}; pub(crate) fn measurement_to_u8(m: Option) -> u8 { match m { @@ -322,16 +322,19 @@ macro_rules! create_interface { /// (wrapping mod 2⁶⁴), so results are reproducible and /// independent of the thread count; set the `RAYON_NUM_THREADS` /// environment variable to control the pool size. + /// + /// Returns the outcome codes (0/1/2 = zero/one/lost) as one flat, + /// shot-major `bytearray` plus its `(num_shots, n_measurements)` shape. #[staticmethod] #[pyo3(signature = (prog, n_qubits, min_abs_coeff = 1e-10, num_shots = 1, seed = None))] - pub fn sample( - py: Python<'_>, + pub fn sample<'py>( + py: Python<'py>, prog: &crate::stim_program::PyStimProgram, n_qubits: usize, min_abs_coeff: f64, num_shots: usize, seed: Option, - ) -> pyo3::PyResult>> { + ) -> pyo3::PyResult<(Bound<'py, PyByteArray>, (usize, usize))> { // `prog` was already validated at `StimProgram.parse()` time; // use the validated path to skip redundant re-validation. let raw = py.detach(|| { @@ -352,14 +355,19 @@ macro_rules! create_interface { }, ) }); - Ok(raw + let n_meas = prog.measurement_count(); + let flat: Vec = raw .into_iter() - .map(|shot| { - shot.into_iter() - .map(crate::interface_tableau::measurement_to_u8) - .collect() - }) - .collect()) + .flatten() + .map(crate::interface_tableau::measurement_to_u8) + .collect(); + if flat.len() != num_shots * n_meas { + return Err(pyo3::exceptions::PyRuntimeError::new_err(format!( + "expected {num_shots} x {n_meas} measurements, got {}", + flat.len() + ))); + } + Ok((PyByteArray::new(py, &flat), (num_shots, n_meas))) } /// Fork this tableau, cloning all quantum state but reinitializing the RNG. diff --git a/ppvm-python/pyproject.toml b/ppvm-python/pyproject.toml index 5189852df..094a98fca 100644 --- a/ppvm-python/pyproject.toml +++ b/ppvm-python/pyproject.toml @@ -10,6 +10,7 @@ requires-python = ">=3.10" dependencies = [ "bloqade-circuit>=0.14.1", "kirin-toolchain~=0.22.2", + "numpy>=2.2.6", ] [build-system] @@ -54,7 +55,6 @@ addopts = "--benchmark-disable" [dependency-groups] dev = [ - "numpy>=2.2.6", "pytest>=9.0.2", "pytest-benchmark>=5.2.3", ] diff --git a/ppvm-python/src/ppvm/_core.pyi b/ppvm-python/src/ppvm/_core.pyi index bd890adc7..c4017324c 100644 --- a/ppvm-python/src/ppvm/_core.pyi +++ b/ppvm-python/src/ppvm/_core.pyi @@ -171,7 +171,7 @@ class _GeneralizedTableauBase: min_abs_coeff: float = 1e-10, num_shots: int = 1, seed: int | None = None, - ) -> list[list[int]]: ... + ) -> tuple[bytearray, tuple[int, int]]: ... def fork(self, seed: int | None = None) -> _GeneralizedTableauBase: ... class StimProgram: diff --git a/ppvm-python/src/ppvm/generalized_tableau.py b/ppvm-python/src/ppvm/generalized_tableau.py index 846828f95..321dcf346 100644 --- a/ppvm-python/src/ppvm/generalized_tableau.py +++ b/ppvm-python/src/ppvm/generalized_tableau.py @@ -4,6 +4,9 @@ import enum from collections.abc import Iterable from dataclasses import InitVar, dataclass, field +from typing import Literal, overload + +import numpy as np from . import _core from ._core import StimProgram @@ -377,6 +380,43 @@ def run(self, prog: StimProgram) -> list[MeasurementResult]: # stim familiarity alias do = run + @overload + @classmethod + def sample( + cls, + prog: StimProgram, + n_qubits: int | None = ..., + min_abs_coeff: float = ..., + num_shots: int = ..., + seed: int | None = ..., + as_numpy: Literal[False] = ..., + ) -> list[list[MeasurementResult]]: ... + + @overload + @classmethod + def sample( + cls, + prog: StimProgram, + n_qubits: int | None = ..., + min_abs_coeff: float = ..., + num_shots: int = ..., + seed: int | None = ..., + *, + as_numpy: Literal[True], + ) -> np.ndarray: ... + + @overload + @classmethod + def sample( + cls, + prog: StimProgram, + n_qubits: int | None = ..., + min_abs_coeff: float = ..., + num_shots: int = ..., + seed: int | None = ..., + as_numpy: bool = ..., + ) -> list[list[MeasurementResult]] | np.ndarray: ... + @classmethod def sample( cls, @@ -385,7 +425,8 @@ def sample( min_abs_coeff: float = 1e-10, num_shots: int = 1, seed: int | None = None, - ) -> list[list[MeasurementResult]]: + as_numpy: bool = False, + ) -> list[list[MeasurementResult]] | np.ndarray: """Run ``num_shots`` shots of ``prog`` and return all measurement results. Each shot starts from a fresh tableau, so this is the right entry @@ -403,12 +444,53 @@ def sample( reproducible and independent of the number of threads. Set the ``RAYON_NUM_THREADS`` environment variable before the first call to control the pool size (it defaults to the number of logical cores). + + With ``as_numpy=True`` the results come back as a writable ``int8`` + array of shape ``(num_shots, n_measurements)`` holding the + `MeasurementResult` values (0/1/2 = zero/one/lost), which avoids + building one Python object per measurement. """ if n_qubits is None: n_qubits = max(1, prog.num_qubits) native_cls = _native_tableau_cls(n_qubits) - raw = native_cls.sample(prog, n_qubits, min_abs_coeff, num_shots, seed) - return [[_BY_VALUE[x] for x in shot] for shot in raw] + buf, (n, m) = native_cls.sample(prog, n_qubits, min_abs_coeff, num_shots, seed) + if as_numpy: + return np.frombuffer(buf, dtype=np.int8).reshape(n, m) + return [[_BY_VALUE[x] for x in buf[i * m : (i + 1) * m]] for i in range(n)] + + +@overload +def sample_stim( + prog: StimProgram, + n_qubits: int | None = ..., + min_abs_coeff: float = ..., + num_shots: int = ..., + seed: int | None = ..., + as_numpy: Literal[False] = ..., +) -> list[list[MeasurementResult]]: ... + + +@overload +def sample_stim( + prog: StimProgram, + n_qubits: int | None = ..., + min_abs_coeff: float = ..., + num_shots: int = ..., + seed: int | None = ..., + *, + as_numpy: Literal[True], +) -> np.ndarray: ... + + +@overload +def sample_stim( + prog: StimProgram, + n_qubits: int | None = ..., + min_abs_coeff: float = ..., + num_shots: int = ..., + seed: int | None = ..., + as_numpy: bool = ..., +) -> list[list[MeasurementResult]] | np.ndarray: ... def sample_stim( @@ -417,14 +499,20 @@ def sample_stim( min_abs_coeff: float = 1e-10, num_shots: int = 1, seed: int | None = None, -) -> list[list[MeasurementResult]]: + as_numpy: bool = False, +) -> list[list[MeasurementResult]] | np.ndarray: """Multi-shot sampling — module-level alias for ``GeneralizedTableau.sample``. When ``n_qubits`` is ``None`` (the default) the qubit count is inferred from the program; see `GeneralizedTableau.sample`. Shots are sampled in parallel across CPU cores with the GIL released; see `GeneralizedTableau.sample` for - seeding and ``RAYON_NUM_THREADS``. + seeding, ``RAYON_NUM_THREADS`` and the ``as_numpy`` array format. """ return GeneralizedTableau.sample( - prog, n_qubits, min_abs_coeff=min_abs_coeff, num_shots=num_shots, seed=seed + prog, + n_qubits, + min_abs_coeff=min_abs_coeff, + num_shots=num_shots, + seed=seed, + as_numpy=as_numpy, ) diff --git a/ppvm-python/test/generalized_tableau/test_basics.py b/ppvm-python/test/generalized_tableau/test_basics.py index f90464379..d238289eb 100644 --- a/ppvm-python/test/generalized_tableau/test_basics.py +++ b/ppvm-python/test/generalized_tableau/test_basics.py @@ -26,7 +26,7 @@ def test_sample_too_many_qubits_raises_clear_error(): with pytest.raises(ValueError, match=f"between 1 and {MAX_N_QUBITS}"): # `prog=None` is deliberate: n_qubits is validated before `prog` is # ever used, so the ValueError fires regardless of the program. - sample_stim(prog=None, n_qubits=MAX_N_QUBITS + 1) # ty: ignore[invalid-argument-type] + sample_stim(prog=None, n_qubits=MAX_N_QUBITS + 1) # ty: ignore[no-matching-overload] def test_measure_zero_state(): diff --git a/ppvm-python/test/generalized_tableau/test_stim.py b/ppvm-python/test/generalized_tableau/test_stim.py index 599264b79..d9c9c313d 100644 --- a/ppvm-python/test/generalized_tableau/test_stim.py +++ b/ppvm-python/test/generalized_tableau/test_stim.py @@ -2,6 +2,7 @@ import tempfile import textwrap +import numpy as np import pytest from ppvm import GeneralizedTableau, StimProgram, sample_stim @@ -259,6 +260,34 @@ def test_sample_stim_zero_shots_returns_empty(): assert sample_stim(prog, n_qubits=1, num_shots=0) == [] +def test_sample_stim_as_numpy_matches_list_output(): + # Loss + H so all three outcome codes (0/1/2) appear. + prog = StimProgram.parse("H 0 1\nI_ERROR[loss](0.3) 0 1\nM 0 1") + as_list = sample_stim(prog, n_qubits=2, num_shots=200, seed=3) + bits = sample_stim(prog, n_qubits=2, num_shots=200, seed=3, as_numpy=True) + assert isinstance(bits, np.ndarray) + assert bits.dtype == np.int8 + assert bits.shape == (200, 2) + assert bits.flags.writeable + assert bits.tolist() == [[int(r) for r in shot] for shot in as_list] + assert set(np.unique(bits)) == {0, 1, 2} + + +def test_sample_classmethod_as_numpy_equivalent(): + prog = StimProgram.parse("H 0\nM 0") + a = GeneralizedTableau.sample(prog, 1, num_shots=10, seed=0, as_numpy=True) + b = sample_stim(prog, n_qubits=1, num_shots=10, seed=0, as_numpy=True) + np.testing.assert_array_equal(a, b) + + +def test_sample_stim_as_numpy_empty_shapes(): + prog = StimProgram.parse("X 0\nM 0 0 0") + assert sample_stim(prog, n_qubits=1, num_shots=0, as_numpy=True).shape == (0, 3) + no_meas = StimProgram.parse("X 0") + assert sample_stim(no_meas, n_qubits=1, num_shots=4, as_numpy=True).shape == (4, 0) + assert sample_stim(no_meas, n_qubits=1, num_shots=4) == [[]] * 4 + + def test_sample_stim_seeded_is_reproducible_for_large_batches(): # A randomising circuit at a large shot count. Per-shot seeds are derived # from the shot index, so two runs must agree exactly regardless of whether diff --git a/ppvm-python/uv.lock b/ppvm-python/uv.lock index 88e4fa704..a94a638df 100644 --- a/ppvm-python/uv.lock +++ b/ppvm-python/uv.lock @@ -906,11 +906,11 @@ source = { editable = "." } dependencies = [ { name = "bloqade-circuit" }, { name = "kirin-toolchain" }, + { name = "numpy" }, ] [package.dev-dependencies] dev = [ - { name = "numpy" }, { name = "pytest" }, { name = "pytest-benchmark" }, ] @@ -919,11 +919,11 @@ dev = [ requires-dist = [ { name = "bloqade-circuit", specifier = ">=0.14.1" }, { name = "kirin-toolchain", specifier = "~=0.22.2" }, + { name = "numpy", specifier = ">=2.2.6" }, ] [package.metadata.requires-dev] dev = [ - { name = "numpy", specifier = ">=2.2.6" }, { name = "pytest", specifier = ">=9.0.2" }, { name = "pytest-benchmark", specifier = ">=5.2.3" }, ]