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
30 changes: 19 additions & 11 deletions crates/ppvm-python-native/src/interface_tableau.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<bool>) -> u8 {
match m {
Expand Down Expand Up @@ -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<u64>,
) -> pyo3::PyResult<Vec<Vec<u8>>> {
) -> 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(|| {
Expand All @@ -352,14 +355,19 @@ macro_rules! create_interface {
},
)
});
Ok(raw
let n_meas = prog.measurement_count();
let flat: Vec<u8> = 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.
Expand Down
2 changes: 1 addition & 1 deletion ppvm-python/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down Expand Up @@ -54,7 +55,6 @@ addopts = "--benchmark-disable"

[dependency-groups]
dev = [
"numpy>=2.2.6",
"pytest>=9.0.2",
"pytest-benchmark>=5.2.3",
]
2 changes: 1 addition & 1 deletion ppvm-python/src/ppvm/_core.pyi
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
100 changes: 94 additions & 6 deletions ppvm-python/src/ppvm/generalized_tableau.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -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
Expand All @@ -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(
Expand All @@ -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,
)
2 changes: 1 addition & 1 deletion ppvm-python/test/generalized_tableau/test_basics.py
Original file line number Diff line number Diff line change
Expand Up @@ -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():
Expand Down
29 changes: 29 additions & 0 deletions ppvm-python/test/generalized_tableau/test_stim.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
import tempfile
import textwrap

import numpy as np
import pytest

from ppvm import GeneralizedTableau, StimProgram, sample_stim
Expand Down Expand Up @@ -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
Expand Down
4 changes: 2 additions & 2 deletions ppvm-python/uv.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

Loading