diff --git a/CHANGELOG.md b/CHANGELOG.md index 5470441c..49d0d018 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -10,16 +10,25 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Added - **QEC Stim MPP Import**: Added utilities for building `StabilizerCode` inputs from unsigned Stim `MPP` layers, including sparse Stim qubit id mapping, coordinate import, multi-layer selection, detector/logical-observable import, and the optional `graphqomb[stim]` extra. Signed products using inverted Pauli targets are rejected because stabilizer signs are not retained. -- **Stim Circuit Import**: Added `stim_file_to_pattern()`, `stim_text_to_pattern()`, and `stim_circuit_to_pattern()` for converting supported Stim circuits into GraphQOMB patterns. The importer handles Clifford unitary blocks and `TICK`-separated Pauli measurements (`M`/`MZ`, `MX`, `MY`, `MXX`, `MYY`, `MZZ`, and `MPP`), assigns single-qubit measurement bases directly to their graph nodes, terminates each qubit lifetime at its direct measurement while allowing disjoint qubits to continue, validates that same-block MPP products commute, supports Type I and Type II Y foliation, causally lowers commuting MPP groups, composes each MPP unit through a separate unmeasured output layer, and resolves detector/logical-observable records once across the full flattened circuit. Circuit-level noise and measurement-error probabilities are intentionally omitted because GraphQOMB uses an MBQC-specific noise model; reset, measurement-reset, and feedback instructions remain unsupported. +- **Stim Circuit Import**: Added `stim_file_to_pattern()`, `stim_text_to_pattern()`, and `stim_circuit_to_pattern()` for converting supported Stim circuits into GraphQOMB patterns. The importer handles initial Pauli resets (`R`/`RZ`, `RX`, and `RY`), Clifford unitary blocks, and `TICK`-separated Pauli measurements (`M`/`MZ`, `MX`, `MY`, `MXX`, `MYY`, `MZZ`, and `MPP`), assigns single-qubit measurement bases directly to their graph nodes, terminates each qubit lifetime at its direct measurement while allowing disjoint qubits to continue, validates that same-block MPP products commute, supports Type I and Type II Y foliation, keeps MPP ancilla nodes out of the X-correction flow, derives the complete Z-correction flow from the composed data flow without Pauli simplification, places gate and MPP blocks on explicit temporal Z layers while preserving the first two Stim coordinate components, aligns parallel gate outputs across different transpiled depths, relocates idle I/O nodes without adding graph nodes or edges, and resolves detector/logical-observable records once across the full flattened circuit. Circuit-level noise and measurement-error probabilities are intentionally omitted because GraphQOMB uses an MBQC-specific noise model; mid-circuit reset, measurement-reset, and feedback instructions remain unsupported. +- **Input Initialization Bases**: Added per-input positive Pauli eigenstate initialization via `GraphState.register_input(..., init_axis=...)`, with `X+`, `Y+`, and `Z+` propagated through `qompile()`, `Pattern`, `PatternSimulator`, and Stim export. ### Changed - **Development Tooling**: Use uv as the default dependency manager for local development, CI, documentation builds, and publishing workflows. - **Graph State API**: Replaced legacy `physical_*` graph methods/properties with standard graph-style `nodes`, `edges`, `add_node()`, `add_edge()`, `remove_node()`, `remove_edge()`, and count/query helpers. - **Pattern Simulator**: Materialize pending output Pauli-frame corrections when returning output statevectors or explicit output measurement results. +- **Pattern Simulator**: Sample all measurements from exact Born probabilities by default, with `calc_prob=False` retaining the legacy 50/50 assumption for non-output measurements. +- **PTN Format**: Bumped exported `.ptn` files to format version 2 to record non-default input initialization bases with `.input_basis`; version 1 files remain readable and default inputs to `X+`. +- **Stim Compiler**: Input reset instructions now follow the stored initialization basis (`RX`, `RY`, or `R`) while non-input preparations remain `RX`. +- **Stim Coordinate Import**: Reject `QUBIT_COORDS` whose first two components collide across qubits, because graph nodes are placed by the XY projection and colliding projections produce coincident nodes. ### Fixed +- **Pattern Simulator**: Raise `TypeError` for unsupported commands instead of recursively redispatching them. +- **QEC Graph-State Builder**: Add the required ancilla CZ edge when shared data qubits contain an odd number of oppositely ordered stabilizer-interaction pairs, using the same `Z`-before-`Y`-before-`X` rule for Type I and Type II foliation. +- **Stim Circuit Import**: Preserve the existing data-lane endpoint and its coordinate when importing a terminal single-qubit measurement. +- **Stim Circuit Import Coordinates**: Generate consistent 3D spacetime coordinates across gate, idle, and MPP fragments instead of mixing raw 2D gate coordinates with MPP Z layers. - **Stim Compiler**: Preserve minus-signed axis measurement bases using inverted Stim measurement targets. - **State Vector Array Conversion**: Convert to a real NumPy dtype without warnings when every amplitude is real-valued, and reject the conversion when it would discard nonzero imaginary amplitudes. diff --git a/docs/source/ptn_format.rst b/docs/source/ptn_format.rst index bd6b1f14..77aeebf9 100644 --- a/docs/source/ptn_format.rst +++ b/docs/source/ptn_format.rst @@ -1,6 +1,24 @@ Pattern Text Format =================== +GraphQOMB writes pattern files using format version 2. Version 2 adds optional +per-input positive Pauli eigenstate initialization: + +.. code-block:: text + + .version 2 + .input 0:0 1:1 2:2 + .input_basis 1:Y 2:Z + +Each ``.input_basis`` entry has the form ``node:X``, ``node:Y``, or ``node:Z`` +and must reference a node declared by ``.input``. ``X`` initialization is the +default and is omitted when serializing, so the example initializes nodes 0, +1, and 2 as ``X+``, ``Y+``, and ``Z+`` respectively. + +Version 1 files remain readable. Inputs in a version 1 file, or in a version 2 +file without a corresponding ``.input_basis`` entry, are initialized as +``X+``. The ``.input_basis`` directive is rejected in version 1 files. + :mod:`graphqomb.ptn_format` module ++++++++++++++++++++++++++++++++++ diff --git a/docs/source/qec.rst b/docs/source/qec.rst index 45052d01..7a640577 100644 --- a/docs/source/qec.rst +++ b/docs/source/qec.rst @@ -16,6 +16,15 @@ support, followed by an output layer; qubits without Y support retain the Type I two-measurement-layer layout. Ancilla support edges only touch measurement layers, and the output of one composed unit becomes the input of the next. +Both foliation variants use the local stabilizer-interaction order +``Z -> Y -> X`` on each shared data qubit. For a pair of stabilizers, the +builder adds a CZ edge between their ancillas when an odd number of shared +data-qubit pairs reverse that order: one qubit applies stabilizer ``a`` before +``b`` while the other applies ``b`` before ``a``. Equal-Pauli overlaps have no +strict order and do not contribute. The builder tracks only the parity of the +two directions and considers only stabilizer pairs that share a data qubit, so +sparse codes do not require an all-pairs stabilizer scan. + .. automodule:: graphqomb.qec.qeccode :members: :show-inheritance: diff --git a/docs/source/simulator.rst b/docs/source/simulator.rst index e44fb3e5..f24b2fa2 100644 --- a/docs/source/simulator.rst +++ b/docs/source/simulator.rst @@ -1,6 +1,17 @@ Statevector Simulator ===================== +``PatternSimulator`` samples measurement results from their exact Born +probabilities by default. This is required when inputs use ``Y+`` or ``Z+`` +initialization, because non-output measurements are not necessarily uniformly +random. + +For compatibility with the previous faster approximation, pass +``calc_prob=False``. In that mode, non-output measurements are sampled 50/50; +output measurements still use their exact probabilities. This approximation is +only appropriate when the pattern guarantees uniformly random non-output +measurements. + .. automodule:: graphqomb.simulator :members: :show-inheritance: diff --git a/docs/source/stim_importer.rst b/docs/source/stim_importer.rst index 479ad4fa..1473b7ce 100644 --- a/docs/source/stim_importer.rst +++ b/docs/source/stim_importer.rst @@ -8,17 +8,47 @@ Install the optional Stim integration before importing this module: uv add "graphqomb[stim]" The circuit importer converts supported Stim circuits into GraphQOMB -measurement patterns. It accepts Clifford unitary blocks and Pauli measurement -blocks separated by ``TICK``. The supported Pauli measurement instructions are -``M``/``MZ``, ``MX``, ``MY``, ``MXX``, ``MYY``, ``MZZ``, and ``MPP``. +measurement patterns. It accepts initial Pauli resets, Clifford unitary blocks, +and Pauli measurement blocks separated by ``TICK``. The supported Pauli +measurement instructions are ``M``/``MZ``, ``MX``, ``MY``, ``MXX``, ``MYY``, +``MZZ``, and ``MPP``. + +Initial reset instructions +-------------------------- + +Leading reset instructions determine the positive Pauli eigenstate used for an +input: ``R``/``RZ`` initializes ``Z+``, ``RX`` initializes ``X+``, and ``RY`` +initializes ``Y+``. A reset is initial when its target qubit has not previously +participated in a unitary or measurement operation. Operations on other qubits +do not prevent a later initial reset. Repeated leading resets are accepted, and +the last reset on each qubit determines its initialization state. + +Stim canonicalizes the Z-axis aliases internally, so the importer accepts both +``M`` and ``MZ`` for Z measurement and both ``R`` and ``RZ`` for Z reset. Stim +export uses ``MZ`` for Z measurement and ``R`` for Z reset. + +Mid-circuit resets are rejected because GraphQOMB patterns do not currently +represent multiple lifetimes for one logical qubit. Combined measurement-reset +instructions (``MR``/``MRZ``, ``MRX``, and ``MRY``) remain unsupported. Single-qubit measurements assign an ``AxisMeasBasis`` directly to the measured -graph node. They do not create an ``MPP`` extraction or an ancillary parity -measurement node. Inverted single-qubit measurement targets select the minus -sign of that node's basis. A direct single-qubit measurement terminates that -qubit's lifetime: a later quantum operation on the same qubit is rejected, -while operations on other qubits may continue. Reset and qubit reuse are not -supported. +data-lane endpoint without replacing that node or its coordinate. They do not +create an ``MPP`` extraction or an ancillary parity measurement node. Inverted +single-qubit measurement targets select the minus sign of that node's basis. A +direct single-qubit measurement terminates that qubit's lifetime: a later +quantum operation on the same qubit is rejected, while operations on other +qubits may continue. A measured qubit cannot begin a new lifetime later in the +circuit. + +The first two components of ``QUBIT_COORDS`` are used as the fixed spatial +``(x, y)`` position of each data lane. The importer supplies the temporal ``z`` +component. Every unitary ``TICK`` block is transpiled before placement. Its +input layer starts at the preceding block's output ``z``, and all of its output +nodes share the maximum transpiled depth of the block. A shorter data-wire +chain is spread across that same interval. A live qubit with no operation in +the block remains a single input/output node and is relocated directly to the +common output layer; this adds no graph node or edge and does not change the +circuit semantics. Two-qubit measurements are parity measurements and are lowered to equivalent unsigned ``MPP`` products. Inverted targets in ``MXX``, ``MYY``, ``MZZ``, and @@ -27,14 +57,29 @@ corresponding parity offset. All ``MPP`` instructions within one ``TICK`` block are represented by one combined extraction and are validated to commute. Anticommuting products in -the same block are rejected. The importer uses one compact -stabilizer-measurement unit when its correction flow is causal. If the compact -unit has a cyclic Pauli flow, the commuting products are lowered to equivalent -sequential units instead. Non-commuting measurements must be separated by -``TICK`` in the source circuit. Each unit has a distinct unmeasured output -layer, which is composed with the next unitary or measurement fragment by qubit -index. Pass ``y_foliation=YFoliation.TYPE_II`` to any of the three import entry -points to use the three-layer Y-measurement construction; Type I is the default. +the same block are rejected. Within a combined block, local stabilizer +interactions are ordered ``Z -> Y -> X`` on each shared data qubit. If an odd +number of shared-data-qubit pairs reverse the order of the same two +stabilizers, the graph-state builder adds the required CZ edge between their +ancillas. This rule is applied automatically for both Type I and Type II +foliation. + +Only data-wire nodes contribute to the X-correction flow; MPP ancilla nodes do +not produce X corrections. After composing all graph fragments, the importer +derives the Z-correction flow from the odd neighborhood of the complete +X-correction flow. Both correction maps are passed directly to ``qompile()`` +without Pauli simplification or an importer-specific fallback. Non-commuting +measurements must be separated by ``TICK`` in the source circuit. Each unit has +a distinct unmeasured output layer, which is composed with the next unitary or +measurement fragment by qubit index. Pass +``y_foliation=YFoliation.TYPE_II`` to any of the three import entry points to use +the three-layer Y-measurement construction; Type I is the default. + +An MPP block starts at the preceding gate or MPP output layer and ends two +``z`` units later. Live lanes not used by that MPP block are relocated to the +same output layer without adding nodes. Consequently, a composed output and +the next active fragment input have the same ``z`` coordinate, and imported +patterns do not mix 2D spatial coordinates with 3D spacetime coordinates. The flattened ideal circuit is analyzed once for measurement records. Records from single-qubit measurements, pair measurements, ``MPP``, and ideal-zero @@ -56,11 +101,6 @@ to Pauli measurements are also omitted while retaining the ideal measurement. Heralded noise records are retained as ideal zero-valued record positions so that later ``DETECTOR`` and ``OBSERVABLE_INCLUDE`` references remain aligned. -Reset instructions (``R``/``RZ``, ``RX``, and ``RY``) and combined -measurement-reset instructions (``MR``/``MRZ``, ``MRX``, and ``MRY``) are not -handled by this importer. Consequently, a directly measured qubit cannot begin -a new lifetime later in the circuit. - .. code-block:: python from graphqomb.qec.qeccode import YFoliation @@ -68,6 +108,8 @@ a new lifetime later in the circuit. result = stim_text_to_pattern( """ + RY 0 + R 1 2 MX 0 MYY 1 2 DETECTOR rec[-2] rec[-1] diff --git a/graphqomb/euler.py b/graphqomb/euler.py index 0b4e0bb9..70c78797 100644 --- a/graphqomb/euler.py +++ b/graphqomb/euler.py @@ -83,7 +83,7 @@ def bloch_sphere_coordinates(vector: NDArray[np.complex128]) -> tuple[float, flo Returns ------- - `tuple`\[`float`, `float`] + `tuple`\[`float`, `float`\] Bloch sphere coordinates (:math:`\theta`, :math:`\phi`) """ # normalize @@ -223,7 +223,7 @@ def meas_basis_info(vector: NDArray[np.complex128]) -> tuple[Plane, float]: Returns ------- - `tuple`\[`Plane`, `float`] + `tuple`\[`Plane`, `float`\] measurement plane and angle Raises diff --git a/graphqomb/feedforward.py b/graphqomb/feedforward.py index 395b0f01..be3e2e47 100644 --- a/graphqomb/feedforward.py +++ b/graphqomb/feedforward.py @@ -124,13 +124,15 @@ def check_dag(dag: Mapping[int, Iterable[int]]) -> None: Raises ------ ValueError - If the flowlike object is not causal with respect to the graph state + If the graph contains a cycle """ - for node, children in dag.items(): - for child in children: - if node in dag[child]: - msg = f"Cycle detected in the graph: {node} -> {child}" - raise ValueError(msg) + inv_dag = inverse_dag_from_dag(dag) + try: + tuple(TopologicalSorter(inv_dag).static_order()) + except CycleError as exc: + cycle = " -> ".join(map(str, exc.args[1])) + msg = f"Cycle detected in the graph: {cycle}" + raise ValueError(msg) from exc def inverse_dag_from_dag( @@ -223,7 +225,7 @@ def signal_shifting( Returns ------- - `tuple`\[`dict`\[`int`, `set`\[`int`\]\], `dict`\[`int`, `set`\[`int`\]\]] + `tuple`\[`dict`\[`int`, `set`\[`int`\]\], `dict`\[`int`, `set`\[`int`\]\]\] Updated correction maps for X and Z after signal shifting. """ if zflow is None: @@ -266,7 +268,7 @@ def propagate_correction_map( # noqa: C901, PLR0912 Returns ------- - `tuple`\[`dict`\[`int`, `set`\[`int`\]\], `dict`\[`int`, `set`\[`int`\]\]] + `tuple`\[`dict`\[`int`, `set`\[`int`\]\], `dict`\[`int`, `set`\[`int`\]\]\] Updated correction maps for X and Z after measurement at the target node. Raises @@ -349,7 +351,7 @@ def pauli_simplification( # noqa: C901, PLR0912 Returns ------- - `tuple`\[`dict`\[`int`, `set`\[`int`\]\], `dict`\[`int`, `set`\[`int`\]\]] + `tuple`\[`dict`\[`int`, `set`\[`int`\]\], `dict`\[`int`, `set`\[`int`\]\]\] Updated correction maps for X and Z after simplification. """ if zflow is None: diff --git a/graphqomb/graphstate.py b/graphqomb/graphstate.py index 37c5853a..d40649e1 100644 --- a/graphqomb/graphstate.py +++ b/graphqomb/graphstate.py @@ -26,7 +26,7 @@ import typing_extensions -from graphqomb.common import MeasBasis, Plane, PlannerMeasBasis +from graphqomb.common import Axis, MeasBasis, Plane, PlannerMeasBasis from graphqomb.euler import update_lc_basis, update_lc_lc if TYPE_CHECKING: @@ -49,6 +49,17 @@ def input_node_indices(self) -> dict[int, int]: qubit indices map of input nodes. """ + @property + def input_initialization_axes(self) -> dict[int, Axis]: + r"""Input initialization Pauli axes. + + Returns + ------- + `dict`\[`int`, `Axis`\] + map of input nodes to Pauli initialization axes. + """ + return dict.fromkeys(self.input_node_indices, Axis.X) + @property @abc.abstractmethod def output_node_indices(self) -> dict[int, int]: @@ -174,7 +185,7 @@ def number_of_edges(self) -> int: return len(self.edges) @abc.abstractmethod - def register_input(self, node: int, q_index: int) -> None: + def register_input(self, node: int, q_index: int, *, init_axis: Axis = Axis.X) -> None: """Mark the node as an input node. Parameters @@ -183,6 +194,8 @@ def register_input(self, node: int, q_index: int) -> None: node index q_index : `int` logical qubit index + init_axis : `Axis`, optional + Pauli axis for positive-eigenstate initialization, by default Axis.X """ @abc.abstractmethod @@ -244,6 +257,7 @@ class GraphState(BaseGraphState): """Minimal implementation of GraphState.""" __input_node_indices: dict[int, int] + __input_initialization_axes: dict[int, Axis] __output_node_indices: dict[int, int] __nodes: set[int] __neighbors: dict[int, set[int]] @@ -257,6 +271,7 @@ class GraphState(BaseGraphState): def __init__(self) -> None: self.__input_node_indices = {} + self.__input_initialization_axes = {} self.__output_node_indices = {} self.__nodes = set() self.__neighbors = {} @@ -278,6 +293,18 @@ def input_node_indices(self) -> dict[int, int]: """ return self.__input_node_indices.copy() + @property + @typing_extensions.override + def input_initialization_axes(self) -> dict[int, Axis]: + r"""Input initialization Pauli axes. + + Returns + ------- + `dict`\[`int`, `Axis`\] + map of input nodes to Pauli initialization axes. + """ + return self.__input_initialization_axes.copy() + @property @typing_extensions.override def output_node_indices(self) -> dict[int, int]: @@ -514,7 +541,7 @@ def remove_edge(self, node1: int, node2: int) -> None: self.__neighbors[node2] -= {node1} @typing_extensions.override - def register_input(self, node: int, q_index: int) -> None: + def register_input(self, node: int, q_index: int, *, init_axis: Axis = Axis.X) -> None: """Mark the node as an input node. Parameters @@ -523,9 +550,13 @@ def register_input(self, node: int, q_index: int) -> None: node index q_index : `int` logical qubit index + init_axis : `Axis`, optional + Pauli axis for positive-eigenstate initialization, by default Axis.X Raises ------ + TypeError + If ``init_axis`` is not an `Axis` value. ValueError If the node is already registered as an input node. """ @@ -536,7 +567,11 @@ def register_input(self, node: int, q_index: int) -> None: if q_index in self.input_node_indices.values(): msg = "The q_index already exists in input qubit indices" raise ValueError(msg) + if not isinstance(init_axis, Axis): + msg = "Input initialization axis must be one of Axis.X, Axis.Y, Axis.Z" + raise TypeError(msg) self.__input_node_indices[node] = q_index + self.__input_initialization_axes[node] = init_axis @typing_extensions.override def register_output(self, node: int, q_index: int) -> None: @@ -677,14 +712,18 @@ def _expand_input_local_cliffords(self) -> dict[int, LocalCliffordExpansion]: """ node_index_addition_map: dict[int, LocalCliffordExpansion] = {} new_input_indices: dict[int, int] = {} + new_input_initialization_axes: dict[int, Axis] = {} for input_node, q_index in self.input_node_indices.items(): + init_axis = self.input_initialization_axes[input_node] lc = self._pop_local_clifford(input_node) if lc is None: new_input_indices[input_node] = q_index + new_input_initialization_axes[input_node] = init_axis continue new_node_index0 = self.add_node() new_input_indices[new_node_index0] = q_index + new_input_initialization_axes[new_node_index0] = init_axis new_node_index1 = self.add_node() new_node_index2 = self.add_node() @@ -701,8 +740,13 @@ def _expand_input_local_cliffords(self) -> dict[int, LocalCliffordExpansion]: ) self.__input_node_indices = {} + self.__input_initialization_axes = {} for new_input_index, q_index in new_input_indices.items(): - self.register_input(new_input_index, q_index) + self.register_input( + new_input_index, + q_index, + init_axis=new_input_initialization_axes[new_input_index], + ) return node_index_addition_map @@ -754,6 +798,7 @@ def from_graph( # noqa: C901, PLR0912, PLR0913, PLR0917 outputs: Sequence[NodeT] | None = None, meas_bases: Mapping[NodeT, MeasBasis] | None = None, coordinates: Mapping[NodeT, tuple[float, ...]] | None = None, + input_initialization_axes: Mapping[NodeT, Axis] | None = None, ) -> tuple[GraphState, dict[NodeT, int]]: r"""Create a graph state from nodes and edges with arbitrary node types. @@ -778,6 +823,8 @@ def from_graph( # noqa: C901, PLR0912, PLR0913, PLR0917 Default is None (no bases assigned initially). coordinates : `collections.abc.Mapping`\[NodeT, `tuple`\[`float`, ...\]\] | `None`, optional Coordinates for nodes (2D or 3D). Default is None (no coordinates). + input_initialization_axes : `collections.abc.Mapping`\[NodeT, `Axis`\] | `None`, optional + Pauli initialization axes for input nodes. Default is None (all inputs use Axis.X). Returns ------- @@ -788,7 +835,8 @@ def from_graph( # noqa: C901, PLR0912, PLR0913, PLR0917 Raises ------ ValueError - If duplicate nodes, invalid edges, or invalid input/output nodes. + If duplicate nodes, invalid edges, invalid input/output nodes, or + initialization axes are specified for non-input nodes. """ # Convert nodes to list to preserve order nodes_list = list(nodes) @@ -806,6 +854,15 @@ def from_graph( # noqa: C901, PLR0912, PLR0913, PLR0917 if input_node not in node_set: msg = f"Input node {input_node} not in nodes collection" raise ValueError(msg) + input_set: set[NodeT] = set() if inputs is None else set(inputs) + if input_initialization_axes is not None: + non_input_initialization_nodes = set(input_initialization_axes) - input_set + if non_input_initialization_nodes: + msg = ( + "Input initialization axes specified for non-input node(s): " + f"{sorted(non_input_initialization_nodes, key=repr)}" + ) + raise ValueError(msg) # Validate outputs if outputs is not None: @@ -844,7 +901,10 @@ def from_graph( # noqa: C901, PLR0912, PLR0913, PLR0917 # Register inputs with sequential qubit indices if inputs is not None: for q_index, input_node in enumerate(inputs): - graph_state.register_input(node_map[input_node], q_index) + init_axis = ( + Axis.X if input_initialization_axes is None else input_initialization_axes.get(input_node, Axis.X) + ) + graph_state.register_input(node_map[input_node], q_index, init_axis=init_axis) # Register outputs with sequential qubit indices if outputs is not None: @@ -908,7 +968,11 @@ def from_base_graph_state( # Register inputs with same qubit indices for input_node, q_index in base.input_node_indices.items(): - graph_state.register_input(node_map[input_node], q_index) + graph_state.register_input( + node_map[input_node], + q_index, + init_axis=base.input_initialization_axes.get(input_node, Axis.X), + ) # Register outputs with same qubit indices for output_node, q_index in base.output_node_indices.items(): @@ -1018,11 +1082,19 @@ def compose( # noqa: C901, PLR0912 # Register input nodes with preserved qindices for input_node, q_index in graph1.input_node_indices.items(): - composed_graph.register_input(node_map1[input_node], q_index) + composed_graph.register_input( + node_map1[input_node], + q_index, + init_axis=graph1.input_initialization_axes.get(input_node, Axis.X), + ) for input_node, q_index in graph2.input_node_indices.items(): if q_index not in target_q_indices: - composed_graph.register_input(node_map2[input_node], q_index) + composed_graph.register_input( + node_map2[input_node], + q_index, + init_axis=graph2.input_initialization_axes.get(input_node, Axis.X), + ) # Register output nodes with preserved qindices for output_node, q_index in graph1.output_node_indices.items(): diff --git a/graphqomb/matrix.py b/graphqomb/matrix.py index 4d42a82f..226abfcf 100644 --- a/graphqomb/matrix.py +++ b/graphqomb/matrix.py @@ -8,22 +8,20 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any, TypeVar +from typing import TYPE_CHECKING, Any import numpy as np if TYPE_CHECKING: from numpy.typing import NDArray -T = TypeVar("T", bound=np.number[Any]) # can be removed >= 3.10 - -def is_unitary(mat: NDArray[T]) -> bool: +def is_unitary(mat: NDArray[np.number[Any]]) -> bool: r"""Check if a matrix is unitary. Parameters ---------- - mat : `numpy.typing.NDArray`\[T\] + mat : `numpy.typing.NDArray`\[`numpy.number`\] matrix to check Returns @@ -36,12 +34,12 @@ def is_unitary(mat: NDArray[T]) -> bool: return np.allclose(np.eye(mat.shape[0]), mat @ mat.T.conj()) -def is_hermitian(mat: NDArray[T]) -> bool: +def is_hermitian(mat: NDArray[np.number[Any]]) -> bool: r"""Check if a matrix is Hermitian. Parameters ---------- - mat : `numpy.typing.NDArray`\[T\] + mat : `numpy.typing.NDArray`\[`numpy.number`\] matrix to check Returns diff --git a/graphqomb/pattern.py b/graphqomb/pattern.py index ef36f5d8..1821002a 100644 --- a/graphqomb/pattern.py +++ b/graphqomb/pattern.py @@ -16,6 +16,7 @@ from typing import TYPE_CHECKING from graphqomb.command import TICK, Command, E, M, N +from graphqomb.common import Axis if TYPE_CHECKING: from collections.abc import Callable, Iterator @@ -39,13 +40,16 @@ class Pattern(Sequence[Command]): Pauli frame of the pattern to track the Pauli state of each node input_coordinates : `dict`\[`int`, `tuple`\[`float`, ...\]\] Coordinates for input nodes (2D or 3D) + input_initialization_axes : `dict`\[`int`, `Axis`\] + Pauli initialization axes for input nodes. Missing inputs default to Axis.X. """ input_node_indices: dict[int, int] output_node_indices: dict[int, int] commands: tuple[Command, ...] pauli_frame: PauliFrame - input_coordinates: dict[int, tuple[float, ...]] = dataclasses.field(default_factory=dict) + input_coordinates: dict[int, tuple[float, ...]] = dataclasses.field(default_factory=dict[int, tuple[float, ...]]) + input_initialization_axes: dict[int, Axis] = dataclasses.field(default_factory=dict[int, Axis]) def __len__(self) -> int: return len(self.commands) diff --git a/graphqomb/ptn_format.py b/graphqomb/ptn_format.py index 8765e667..544101f5 100644 --- a/graphqomb/ptn_format.py +++ b/graphqomb/ptn_format.py @@ -38,7 +38,8 @@ if TYPE_CHECKING: from collections.abc import Sequence -PTN_VERSION = 1 +PTN_VERSION = 2 +SUPPORTED_PTN_VERSIONS = frozenset({1, PTN_VERSION}) # Angle formatting/parsing lookup tables _ANGLE_TO_STR: dict[float, str] = { @@ -178,6 +179,13 @@ def _write_header(out: StringIO, pattern: Pattern) -> None: f"{node}:{qidx}" for node, qidx in sorted(pattern.input_node_indices.items(), key=operator.itemgetter(1)) ] out.write(f".input {' '.join(input_parts)}\n") + input_basis_parts = [ + f"{node}:{axis.name}" + for node, _qidx in sorted(pattern.input_node_indices.items(), key=operator.itemgetter(1)) + if (axis := pattern.input_initialization_axes.get(node, Axis.X)) is not Axis.X + ] + if input_basis_parts: + out.write(f".input_basis {' '.join(input_basis_parts)}\n") if pattern.output_node_indices: output_parts = [ @@ -381,6 +389,38 @@ def _parse_node_qubit_pairs(parts: Sequence[str]) -> dict[int, int]: return result +def _parse_node_axis_pairs(parts: Sequence[str]) -> dict[int, Axis]: + r"""Parse node:axis pairs from string parts. + + Returns + ------- + `dict`\[`int`, `Axis`\] + Mapping from node to Pauli axis. + + Raises + ------ + ValueError + If any pair is malformed, duplicated, or uses an invalid axis. + """ + result: dict[int, Axis] = {} + for part in parts: + pair = part.split(":") + if len(pair) != 2: # ruff:ignore[magic-value-comparison] + msg = f"Invalid node:axis pair: {part!r}" + raise ValueError(msg) + node_str, axis_str = pair + node = _parse_int(node_str, "node") + if node in result: + msg = f"Duplicate input basis node: {node}" + raise ValueError(msg) + try: + result[node] = Axis[axis_str] + except KeyError as exc: + msg = f"Invalid input basis axis: {axis_str!r}" + raise ValueError(msg) from exc + return result + + def _parse_node_set(parts: Sequence[str], label: str) -> set[int]: r"""Parse a non-empty set of node ids. @@ -432,61 +472,6 @@ def _parse_arrow_mapping(line: str, label: str) -> tuple[int, set[int]]: return source, targets -def _empty_node_index_map() -> dict[int, int]: - r"""Return an empty node-to-qubit-index map. - - Returns - ------- - `dict`\[`int`, `int`\] - Empty node-to-qubit-index map. - """ - return {} - - -def _empty_coordinates() -> dict[int, tuple[float, ...]]: - r"""Return an empty coordinate map. - - Returns - ------- - `dict`\[`int`, `tuple`\[`float`, ...\]\] - Empty coordinate map. - """ - return {} - - -def _empty_commands() -> list[Command]: - r"""Return an empty command list. - - Returns - ------- - `list`\[`Command`\] - Empty command list. - """ - return [] - - -def _empty_node_set_map() -> dict[int, set[int]]: - r"""Return an empty node-to-node-set map. - - Returns - ------- - `dict`\[`int`, `set`\[`int`\]\] - Empty node-to-node-set map. - """ - return {} - - -def _empty_node_groups() -> list[set[int]]: - r"""Return an empty node group list. - - Returns - ------- - `list`\[`set`\[`int`\]\] - Empty node group list. - """ - return [] - - @dataclass(slots=True) class _PatternData: """Container for parsed pattern data from .ptn format. @@ -499,6 +484,8 @@ class _PatternData: Mapping from node to qubit index for output nodes. input_coordinates : `dict`[`int`, `tuple`[`float`, ...]] Coordinates for input nodes. + input_initialization_axes : `dict`[`int`, `Axis`] + Pauli initialization axes for input nodes. commands : `list`[`Command`] List of quantum commands. xflow : `dict`[`int`, `set`[`int`]] @@ -509,14 +496,15 @@ class _PatternData: Parity check groups for error detection. """ - input_node_indices: dict[int, int] = field(default_factory=_empty_node_index_map) - output_node_indices: dict[int, int] = field(default_factory=_empty_node_index_map) - input_coordinates: dict[int, tuple[float, ...]] = field(default_factory=_empty_coordinates) - commands: list[Command] = field(default_factory=_empty_commands) - xflow: dict[int, set[int]] = field(default_factory=_empty_node_set_map) - zflow: dict[int, set[int]] = field(default_factory=_empty_node_set_map) - parity_check_groups: list[set[int]] = field(default_factory=_empty_node_groups) - logical_observables: dict[int, set[int]] = field(default_factory=_empty_node_set_map) + input_node_indices: dict[int, int] = field(default_factory=dict[int, int]) + output_node_indices: dict[int, int] = field(default_factory=dict[int, int]) + input_coordinates: dict[int, tuple[float, ...]] = field(default_factory=dict[int, tuple[float, ...]]) + input_initialization_axes: dict[int, Axis] = field(default_factory=dict[int, Axis]) + commands: list[Command] = field(default_factory=list[Command]) + xflow: dict[int, set[int]] = field(default_factory=dict[int, set[int]]) + zflow: dict[int, set[int]] = field(default_factory=dict[int, set[int]]) + parity_check_groups: list[set[int]] = field(default_factory=list[set[int]]) + logical_observables: dict[int, set[int]] = field(default_factory=dict[int, set[int]]) @dataclass(slots=True) @@ -529,6 +517,7 @@ class _LoadedGraphState(BaseGraphState): _edges: set[tuple[int, int]] _meas_bases: dict[int, MeasBasis] _coordinates: dict[int, tuple[float, ...]] + _input_initialization_axes: dict[int, Axis] _neighbors: dict[int, set[int]] = field(init=False, repr=False) def __post_init__(self) -> None: @@ -538,6 +527,7 @@ def __post_init__(self) -> None: self._edges = {(node1, node2) if node1 < node2 else (node2, node1) for node1, node2 in self._edges} self._meas_bases = dict(self._meas_bases) self._coordinates = dict(self._coordinates) + self._input_initialization_axes = dict(self._input_initialization_axes) self._neighbors: dict[int, set[int]] = {node: set() for node in self._nodes} for node1, node2 in self._edges: self._neighbors.setdefault(node1, set()).add(node2) @@ -547,6 +537,10 @@ def __post_init__(self) -> None: def input_node_indices(self) -> dict[int, int]: return self._input_node_indices.copy() + @property + def input_initialization_axes(self) -> dict[int, Axis]: + return self._input_initialization_axes.copy() + @property def output_node_indices(self) -> dict[int, int]: return self._output_node_indices.copy() @@ -575,7 +569,7 @@ def add_edge(self, node1: int, node2: int) -> None: msg = "Loaded .ptn graph states are read-only" raise NotImplementedError(msg) - def register_input(self, node: int, q_index: int) -> None: + def register_input(self, node: int, q_index: int, *, init_axis: Axis = Axis.X) -> None: msg = "Loaded .ptn graph states are read-only" raise NotImplementedError(msg) @@ -615,6 +609,27 @@ def _command_nodes(cmd: Command) -> set[int]: return set() +def _input_initialization_axes_from_data(data: _PatternData) -> dict[int, Axis]: + r"""Validate and normalize parsed input initialization axes. + + Returns + ------- + `dict`\[`int`, `Axis`\] + Normalized input initialization axes. + + Raises + ------ + ValueError + If an input basis is specified for a non-input node. + """ + non_input_basis_nodes = set(data.input_initialization_axes) - set(data.input_node_indices) + if non_input_basis_nodes: + msg = f"Input basis specified for non-input node(s): {sorted(non_input_basis_nodes)}" + raise ValueError(msg) + + return {node: data.input_initialization_axes.get(node, Axis.X) for node in data.input_node_indices} + + def _build_pattern(data: _PatternData) -> Pattern: """Build a Pattern from parsed .ptn data. @@ -628,6 +643,8 @@ def _build_pattern(data: _PatternData) -> Pattern: ValueError If parsed commands contain invalid graph structure. """ + input_initialization_axes = _input_initialization_axes_from_data(data) + nodes: set[int] = set(data.input_node_indices) | set(data.output_node_indices) | set(data.input_coordinates) edges: set[tuple[int, int]] = set() meas_bases: dict[int, MeasBasis] = {} @@ -665,6 +682,7 @@ def _build_pattern(data: _PatternData) -> Pattern: _edges=edges, _meas_bases=meas_bases, _coordinates=coordinates, + _input_initialization_axes=input_initialization_axes, ) pauli_frame = PauliFrame( graphstate, @@ -679,6 +697,7 @@ def _build_pattern(data: _PatternData) -> Pattern: commands=tuple(data.commands), pauli_frame=pauli_frame, input_coordinates=dict(data.input_coordinates), + input_initialization_axes=input_initialization_axes, ) @@ -688,7 +707,7 @@ class _Parser: def __init__(self) -> None: self.result = _PatternData() self.current_timeslice = -1 - self.version_found = False + self.version: int | None = None def parse(self, s: str) -> Pattern: r"""Parse the input string and return Pattern. @@ -711,9 +730,12 @@ def parse(self, s: str) -> Pattern: for line_num, raw_line in enumerate(s.splitlines(), 1): self._parse_line(line_num, raw_line) - if not self.version_found: + if self.version is None: msg = "Missing .version directive" raise ValueError(msg) + if self.version == 1 and self.result.input_initialization_axes: + msg = ".input_basis requires .ptn version 2 or later" + raise ValueError(msg) return _build_pattern(self.result) @@ -759,6 +781,8 @@ def _parse_directive(self, line: str) -> None: self._handle_version(content) elif directive == ".input": self.result.input_node_indices = _parse_node_qubit_pairs(content.split()) + elif directive == ".input_basis": + self.result.input_initialization_axes = _parse_node_axis_pairs(content.split()) elif directive == ".output": self.result.output_node_indices = _parse_node_qubit_pairs(content.split()) elif directive == ".coord": @@ -787,10 +811,11 @@ def _handle_version(self, content: str) -> None: If the version is unsupported. """ version = _parse_int(content, "version") - if version != PTN_VERSION: - msg = f"Unsupported .ptn version: {version} (expected {PTN_VERSION})" + if version not in SUPPORTED_PTN_VERSIONS: + supported = ", ".join(str(supported_version) for supported_version in sorted(SUPPORTED_PTN_VERSIONS)) + msg = f"Unsupported .ptn version: {version} (supported: {supported})" raise ValueError(msg) - self.version_found = True + self.version = version def _handle_coord(self, content: str) -> None: """Handle .coord directive. diff --git a/graphqomb/qec/_stim.py b/graphqomb/qec/_stim.py index 1259368f..2c6773ce 100644 --- a/graphqomb/qec/_stim.py +++ b/graphqomb/qec/_stim.py @@ -21,25 +21,25 @@ @dataclass(frozen=True) class StimMppExtraction: - """Stabilizer-code data extracted from Stim MPP products. + r"""Stabilizer-code data extracted from Stim MPP products. Attributes ---------- code : StabilizerCode Dense-column stabilizer code using the ``[Hx | Hz]`` convention. - stim_to_column : dict[int, int] + stim_to_column : `dict`\[`int`, `int`\] Mapping from original Stim qubit ids to dense matrix columns. - column_to_stim : dict[int, int] + column_to_stim : `dict`\[`int`, `int`\] Inverse dense-column mapping. - supports : tuple[PauliSupport, ...] + supports : `tuple`\[``PauliSupport``, ...\] Original Stim Pauli supports, one support per stabilizer row. - detector_rows : tuple[frozenset[int], ...] + detector_rows : `tuple`\[`frozenset`\[`int`\], ...\] Detector groups as selected-MPP stabilizer row indices. - logical_observable_rows : dict[int, frozenset[int]] + logical_observable_rows : `dict`\[`int`, `frozenset`\[`int`\]\] Logical observables as selected-MPP stabilizer row indices. - detector_record_indices : tuple[frozenset[int], ...] + detector_record_indices : `tuple`\[`frozenset`\[`int`\], ...\] Absolute Stim measurement-record indices for selected detectors. - logical_observable_record_indices : dict[int, frozenset[int]] + logical_observable_record_indices : `dict`\[`int`, `frozenset`\[`int`\]\] Absolute Stim record indices for selected logical observables. """ @@ -53,21 +53,21 @@ class StimMppExtraction: logical_observable_record_indices: dict[int, frozenset[int]] = field(default_factory=dict) def detector_groups(self, ancilla_nodes: Mapping[int, int]) -> list[set[int]]: - """Return detector groups mapped to graph node ids for ``qompile``. + r"""Return detector groups mapped to graph node ids for ``qompile``. Returns ------- - list[set[int]] + `list`\[`set`\[`int`\]\] Detector groups suitable for ``qompile``. """ return [_map_rows_to_nodes(rows, ancilla_nodes, "detector") for rows in self.detector_rows] def logical_observables(self, ancilla_nodes: Mapping[int, int]) -> dict[int, set[int]]: - """Return logical observables mapped to graph node ids for ``qompile``. + r"""Return logical observables mapped to graph node ids for ``qompile``. Returns ------- - dict[int, set[int]] + `dict`\[`int`, `set`\[`int`\]\] Logical-observable node groups keyed by Stim observable index. """ return { @@ -149,9 +149,12 @@ def extract_qubit_coordinates( Raises ------ ValueError - If a coordinate has fewer dimensions than requested. + If a coordinate has fewer dimensions than requested, or if two qubits + share the same XY projection. Graph nodes are placed by the first two + coordinate components, so XY collisions produce coincident nodes. """ coordinates: dict[int, Coordinate] = {} + stim_ids_by_xy: dict[Coordinate, list[int]] = {} for stim_id, values in circuit.get_final_qubit_coordinates().items(): if len(values) < coord_dims: msg = ( @@ -159,7 +162,14 @@ def extract_qubit_coordinates( f"fewer than requested coord_dims={coord_dims}." ) raise ValueError(msg) - coordinates[int(stim_id)] = tuple(float(value) for value in values[:coord_dims]) + coordinate = tuple(float(value) for value in values[:coord_dims]) + coordinates[int(stim_id)] = coordinate + stim_ids_by_xy.setdefault(coordinate[:2], []).append(int(stim_id)) + duplicates = {xy: stim_ids for xy, stim_ids in stim_ids_by_xy.items() if len(stim_ids) > 1} + if duplicates: + described = "; ".join(f"qubits {sorted(stim_ids)} share {xy}" for xy, stim_ids in sorted(duplicates.items())) + msg = f"QUBIT_COORDS must have distinct XY projections: {described}." + raise ValueError(msg) return coordinates diff --git a/graphqomb/qec/qeccode.py b/graphqomb/qec/qeccode.py index 22c25b73..79b2bce4 100644 --- a/graphqomb/qec/qeccode.py +++ b/graphqomb/qec/qeccode.py @@ -2,7 +2,9 @@ from __future__ import annotations +from collections import defaultdict from enum import Enum, auto +from itertools import combinations from typing import TYPE_CHECKING, Any, NamedTuple from scipy.sparse import csr_array @@ -11,7 +13,7 @@ from graphqomb.graphstate import GraphState if TYPE_CHECKING: - from collections.abc import Mapping + from collections.abc import Mapping, Sequence _TYPE_II_CHAIN_LENGTH = 3 @@ -94,7 +96,7 @@ def build_graph_state( data_as_io: bool = False, qubit_indices: Mapping[int, int] | None = None, ) -> StabilizerGraphStateBuildResult: - """Build a graph-state unit from a stabilizer code. + r"""Build a graph-state unit from a stabilizer code. Parameters ---------- @@ -106,12 +108,17 @@ def build_graph_state( uses two measurement layers; Type II uses three for Y support. When ``data_as_io`` is enabled, a separate output layer is appended. y_foliation : `YFoliation`, optional - Foliation variant. Type II uses a three-node Y-measured data chain only - for qubits that have an Hx=Hz=1 support in at least one stabilizer row. + Foliation variant. Type I measures each stabilizer ancilla in X when + its row has an even number of Y supports, including zero, and in Y + when the number is odd. Type II uses a three-node Y-measured data chain + only for qubits that have an Hx=Hz=1 support in at least one stabilizer + row. Both variants add a CZ between stabilizer ancillas when an odd + number of shared-data-qubit pairs have opposite local interaction + order under Z-before-Y-before-X ordering. data_as_io : `bool`, optional Whether to register the first stabilizer-measurement data nodes as inputs and append separate unmeasured output nodes, by default False. - qubit_indices : collections.abc.Mapping[int, int] | None, optional + qubit_indices : `collections.abc.Mapping`\[`int`, `int`\] | `None`, optional Mapping from stabilizer-code qubit columns to graph qindices when ``data_as_io`` is enabled. If omitted, code qubit columns are used. @@ -272,23 +279,22 @@ def _add_ancilla_nodes( Mapping from stabilizer row index to graph node. """ ancilla_nodes: dict[int, int] = {} - hx = code.hx.copy() - hz = code.hz.copy() - hx.eliminate_zeros() - hz.eliminate_zeros() - for stabilizer in range(code.num_stabilizers): + supports = _stabilizer_supports(code) + y_meas_basis = AxisMeasBasis(Axis.Y, Sign.PLUS) + for stabilizer, support in enumerate(supports): explicit_ancilla_coord = _explicit_ancilla_coordinate(code, stabilizer) ancilla_node = graph.add_node(coordinate=explicit_ancilla_coord) - graph.assign_meas_basis(ancilla_node, meas_basis) + has_odd_y_support = len(support.hx & support.hz) % 2 == 1 + ancilla_meas_basis = ( + y_meas_basis if data_layer_plan.y_foliation is YFoliation.TYPE_I and has_odd_y_support else meas_basis + ) + graph.assign_meas_basis(ancilla_node, ancilla_meas_basis) ancilla_nodes[stabilizer] = ancilla_node connected_data_nodes = _connect_stabilizer_support( graph, ancilla_node=ancilla_node, - support=_StabilizerSupport( - hx=set(_row_support(hx, stabilizer)), - hz=set(_row_support(hz, stabilizer)), - ), + support=support, data_nodes=data_nodes, data_layer_plan=data_layer_plan, ) @@ -298,9 +304,53 @@ def _add_ancilla_nodes( if inferred_coord is not None: graph.set_coordinate(ancilla_node, inferred_coord) + for left_stabilizer, right_stabilizer in _twisted_stabilizer_pairs(supports): + graph.add_edge(ancilla_nodes[left_stabilizer], ancilla_nodes[right_stabilizer]) + return ancilla_nodes +def _twisted_stabilizer_pairs(supports: Sequence[_StabilizerSupport]) -> list[tuple[int, int]]: + """Return stabilizer pairs whose ancillas require a CZ edge. + + Local interactions on each shared data qubit are ordered Z, Y, then X. + For a stabilizer pair, let ``forward`` and ``reverse`` be the numbers of + shared qubits with the two possible strict orders. The number of twisted + qubit pairs is ``forward * reverse``, so only the parity of each direction + is required. Equal-Pauli overlaps have no strict order and are ignored. + + Candidate stabilizer pairs are generated from per-qubit incidence lists, + avoiding a scan over all stabilizer pairs when supports are sparse. + + Returns + ------- + `list`[`tuple`[`int`, `int`]] + Sorted stabilizer-row pairs requiring an ancilla CZ edge. + """ + incidences_by_qubit: dict[int, list[tuple[int, int]]] = defaultdict(list) + for stabilizer, support in enumerate(supports): + qubits_by_order = ( + support.hz - support.hx, + support.hx & support.hz, + support.hx - support.hz, + ) + for order, qubits in enumerate(qubits_by_order): + for qubit in qubits: + incidences_by_qubit[qubit].append((stabilizer, order)) + + # Incidence lists are in ascending stabilizer order, so combinations + # always yield stabilizer_a < stabilizer_b. + parities: dict[tuple[int, int], list[bool]] = defaultdict(lambda: [False, False]) + for incidences in incidences_by_qubit.values(): + for (stabilizer_a, order_a), (stabilizer_b, order_b) in combinations(incidences, 2): + if order_a == order_b: + continue + direction = 0 if order_a < order_b else 1 + parities[stabilizer_a, stabilizer_b][direction] ^= True + + return sorted(pair for pair, (forward_odd, reverse_odd) in parities.items() if forward_odd and reverse_odd) + + def _connect_stabilizer_support( graph: GraphState, *, @@ -339,15 +389,31 @@ def _type_ii_support_layer(layers: tuple[int, ...], *, has_x: bool, has_z: bool) return layers[2] -def _qubits_with_y_support(code: StabilizerCode) -> set[int]: +def _stabilizer_supports(code: StabilizerCode) -> list[_StabilizerSupport]: + """Return sparse X/Z support sets for every stabilizer row. + + Returns + ------- + `list`[`_StabilizerSupport`] + Support sets indexed by stabilizer row. + """ hx = code.hx.copy() hz = code.hz.copy() hx.eliminate_zeros() hz.eliminate_zeros() + return [ + _StabilizerSupport( + hx=set(_row_support(hx, stabilizer)), + hz=set(_row_support(hz, stabilizer)), + ) + for stabilizer in range(code.num_stabilizers) + ] + +def _qubits_with_y_support(code: StabilizerCode) -> set[int]: y_qubits: set[int] = set() - for stabilizer in range(code.num_stabilizers): - y_qubits.update(set(_row_support(hx, stabilizer)) & set(_row_support(hz, stabilizer))) + for support in _stabilizer_supports(code): + y_qubits.update(support.hx & support.hz) return y_qubits diff --git a/graphqomb/qompiler.py b/graphqomb/qompiler.py index d49dec39..14858a92 100644 --- a/graphqomb/qompiler.py +++ b/graphqomb/qompiler.py @@ -149,4 +149,5 @@ def _qompile( commands=tuple(commands), pauli_frame=pauli_frame, input_coordinates=input_coords, + input_initialization_axes=graph.input_initialization_axes, ) diff --git a/graphqomb/simulator.py b/graphqomb/simulator.py index 1d523dbb..12830955 100644 --- a/graphqomb/simulator.py +++ b/graphqomb/simulator.py @@ -16,13 +16,15 @@ import numpy as np from graphqomb.command import TICK, E, M, N -from graphqomb.common import MeasBasis, Plane +from graphqomb.common import Axis, MeasBasis, Plane from graphqomb.gates import MultiGate, SingleGate, TwoQubitGate from graphqomb.pattern import is_runnable from graphqomb.rng import ensure_rng from graphqomb.statevec import StateVector if TYPE_CHECKING: + from numpy.typing import NDArray + from graphqomb.circuit import BaseCircuit from graphqomb.command import Command from graphqomb.gates import Gate @@ -31,6 +33,11 @@ _X_MATRIX = np.asarray([[0, 1], [1, 0]], dtype=np.complex128) _Z_MATRIX = np.asarray([[1, 0], [0, -1]], dtype=np.complex128) +_INPUT_STATE_VECTORS: dict[Axis, NDArray[np.complex128]] = { + Axis.X: np.asarray([1.0, 1.0], dtype=np.complex128) / np.sqrt(2), + Axis.Y: np.asarray([1.0, 1.0j], dtype=np.complex128) / np.sqrt(2), + Axis.Z: np.asarray([1.0, 0.0], dtype=np.complex128), +} class SimulatorBackend(Enum): @@ -118,7 +125,8 @@ class PatternSimulator: output_results : `dict`\[`int`, `bool`\] Measurement results for output nodes, keyed by logical output index. calc_prob : `bool` - Whether to calculate probabilities. + Whether to sample every measurement from its exact Born probability. + If False, non-output measurements use the legacy 50/50 assumption. """ state: BaseFullStateSimulator @@ -133,7 +141,7 @@ def __init__( pattern: Pattern, backend: SimulatorBackend, *, - calc_prob: bool = False, + calc_prob: bool = True, ) -> None: self.node_indices = list(pattern.input_node_indices.keys()) self.results = {} @@ -146,8 +154,11 @@ def __init__( is_runnable(self.__pattern) if backend == SimulatorBackend.StateVector: - # Note: deterministic check skipped for now - self.state = StateVector.from_num_qubits(len(self.__pattern.input_node_indices)) + input_states = [ + _INPUT_STATE_VECTORS[self.__pattern.input_initialization_axes.get(node, Axis.X)] + for node in self.node_indices + ] + self.state = StateVector.from_product_states(input_states) elif backend == SimulatorBackend.DensityMatrix: raise NotImplementedError else: @@ -155,7 +166,9 @@ def __init__( raise ValueError(msg) @functools.singledispatchmethod - def apply_cmd(self, cmd: Command, *, rng: np.random.Generator) -> None: + def apply_cmd( # ruff:ignore[no-self-use] + self, cmd: Command, *, rng: np.random.Generator + ) -> None: """Apply a command to the state. Parameters @@ -164,8 +177,15 @@ def apply_cmd(self, cmd: Command, *, rng: np.random.Generator) -> None: The command to apply. rng : `numpy.random.Generator` Random number generator to use. + + Raises + ------ + TypeError + If the command type is not supported by the simulator. """ - self.apply_cmd(cmd, rng=rng) + _ = rng + msg = f"Unsupported command for pattern simulation: {type(cmd).__name__}" + raise TypeError(msg) @apply_cmd.register def _(self, cmd: N, *, rng: np.random.Generator) -> None: # noqa: ARG002 @@ -201,20 +221,6 @@ def _updated_measurement_basis(self, cmd: M) -> MeasBasis: return basis - def _sample_measurement_result( - self, - node_id: int, - meas_basis: MeasBasis, - rng: np.random.Generator, - ) -> bool: - state = self.state.state() - norm_sq = float(np.real(np.vdot(state, state))) - basis_vector = meas_basis.vector() - projected = np.tensordot(basis_vector.conjugate(), state, axes=(0, node_id)) - prob_false = float(np.real(np.vdot(projected, projected)) / norm_sq) - prob_false = min(1.0, max(0.0, prob_false)) - return bool(rng.uniform() >= prob_false) - def _apply_output_pauli_frame(self, node: int) -> None: node_id = self.node_indices.index(node) if self.__pattern.pauli_frame.x_pauli[node]: @@ -224,18 +230,13 @@ def _apply_output_pauli_frame(self, node: int) -> None: @apply_cmd.register def _(self, cmd: M, *, rng: np.random.Generator) -> None: - if self.calc_prob: - raise NotImplementedError - node_id = self.node_indices.index(cmd.node) - if cmd.node in self.__pattern.output_node_indices: - meas_basis = self._updated_measurement_basis(cmd) - result = self._sample_measurement_result(node_id, meas_basis, rng) + meas_basis = self._updated_measurement_basis(cmd) + if self.calc_prob or cmd.node in self.__pattern.output_node_indices: + result = self.state.sample_measure(node_id, meas_basis, rng) else: - meas_basis = self._updated_measurement_basis(cmd) result = rng.uniform() < 1 / 2 - - self.state.measure(node_id, meas_basis, result) + self.state.measure(node_id, meas_basis, result) self.results[cmd.node] = result self.node_indices.remove(cmd.node) diff --git a/graphqomb/simulator_backend.py b/graphqomb/simulator_backend.py index dbc29632..e1add065 100644 --- a/graphqomb/simulator_backend.py +++ b/graphqomb/simulator_backend.py @@ -215,6 +215,33 @@ def state(self) -> NDArray[np.complex128]: The current state vector. """ + @abc.abstractmethod + def sample_measure( + self, + qubit: int, + meas_basis: MeasBasis, + rng: np.random.Generator, + ) -> bool: + """Sample and apply a measurement. + + Implementations should reuse the sampled projection when collapsing + the state when possible. + + Parameters + ---------- + qubit : `int` + The qubit to measure. + meas_basis : `MeasBasis` + The measurement basis to use. + rng : `numpy.random.Generator` + Random number generator used to sample the outcome. + + Returns + ------- + `bool` + The sampled measurement result. + """ + @abc.abstractmethod def norm(self) -> float: r"""Get the current state vector norm. diff --git a/graphqomb/statevec.py b/graphqomb/statevec.py index ffa31bd1..1a287269 100644 --- a/graphqomb/statevec.py +++ b/graphqomb/statevec.py @@ -129,6 +129,34 @@ def from_num_qubits(num_qubits: int) -> StateVector: raise ValueError(msg) return StateVector(np.full((2,) * num_qubits, 1 / math.sqrt(2**num_qubits), dtype=np.complex128)) + @staticmethod + def from_product_states(states: Sequence[ArrayLike]) -> StateVector: + r"""Create a state vector from one-qubit product states. + + Parameters + ---------- + states : sequence of array-like + One-qubit state vectors in external qubit order. + + Returns + ------- + `StateVector` + The resulting product state vector. + + Raises + ------ + ValueError + If any one-qubit state does not have exactly two amplitudes. + """ + result = np.asarray(1.0, dtype=np.complex128) + for state in states: + one_qubit_state = np.asarray(state, dtype=np.complex128) + if one_qubit_state.size != 2: # ruff:ignore[magic-value-comparison] + msg = "Each product-state factor must have exactly two amplitudes." + raise ValueError(msg) + result = np.asarray(np.kron(result, one_qubit_state.reshape(2)), dtype=np.complex128) + return StateVector(result) + @staticmethod def tensor_product(a: StateVector, b: StateVector) -> StateVector: """Tensor product with other state vector, self ⊗ other. @@ -199,6 +227,47 @@ def measure(self, qubit: int, meas_basis: MeasBasis, result: int) -> None: self.normalize() + @typing_extensions.override + def sample_measure( + self, + qubit: int, + meas_basis: MeasBasis, + rng: np.random.Generator, + ) -> bool: + """Sample and apply a measurement without repeating the sampled projection. + + Parameters + ---------- + qubit : `int` + The qubit to measure. + meas_basis : `MeasBasis` + The measurement basis to use. + rng : `numpy.random.Generator` + Random number generator used to sample the outcome. + + Returns + ------- + `bool` + The sampled measurement result. + """ + internal_qubit = self.__qindex_mng.external_to_internal(qubit) + norm_sq = float(np.real(np.vdot(self.__state, self.__state))) + basis_vector = meas_basis.vector() + projected = np.tensordot(basis_vector.conjugate(), self.__state, axes=(0, internal_qubit)) + projected_norm_sq = float(np.real(np.vdot(projected, projected))) + prob_false = projected_norm_sq / norm_sq + prob_false = min(1.0, max(0.0, prob_false)) + result = bool(rng.uniform() >= prob_false) + + if result: + basis_vector = meas_basis.flip().vector() + projected = np.tensordot(basis_vector.conjugate(), self.__state, axes=(0, internal_qubit)) + projected_norm_sq = float(np.real(np.vdot(projected, projected))) + + self.__state = np.asarray(projected / math.sqrt(projected_norm_sq), dtype=np.complex128) + self.__qindex_mng.remove_qubit(internal_qubit) + return result + @typing_extensions.override def add_node(self, num_qubits: int) -> None: """Add plus state to the end of state vector. diff --git a/graphqomb/stim_compiler.py b/graphqomb/stim_compiler.py index 9b93b767..1821031e 100644 --- a/graphqomb/stim_compiler.py +++ b/graphqomb/stim_compiler.py @@ -79,7 +79,8 @@ def _emit_input_nodes(self) -> None: coordinates = self._pattern.input_coordinates if self._emit_qubit_coords else None for node in self._pattern.input_node_indices: coord = coordinates.get(node) if coordinates else None - self._process_prepare(node, coord, is_input=True) + axis = self._pattern.input_initialization_axes.get(node, Axis.X) + self._process_prepare(node, coord, is_input=True, init_axis=axis) def _process_commands(self) -> None: for cmd in self._pattern: @@ -95,7 +96,14 @@ def _process_commands(self) -> None: msg = f"Unsupported command for stim compilation: {type(cmd).__name__}" raise TypeError(msg) - def _process_prepare(self, node: int, coordinate: tuple[float, ...] | None, *, is_input: bool) -> None: + def _process_prepare( + self, + node: int, + coordinate: tuple[float, ...] | None, + *, + is_input: bool, + init_axis: Axis = Axis.X, + ) -> None: event = PrepareEvent(time=self._tick, node=self._node_info(node), is_input=is_input) ops = self._validate_ops_for_event(event, self._collect_noise_ops_from_models(lambda m: m.on_prepare(event))) default_placement = default_noise_placement(event) @@ -104,7 +112,8 @@ def _process_prepare(self, node: int, coordinate: tuple[float, ...] | None, *, i coord = coordinate if self._emit_qubit_coords else None if coord is not None: self._stim_io.write(f"QUBIT_COORDS({', '.join(str(c) for c in coord)}) {node}\n") - self._stim_io.write(f"RX {node}\n") + reset_instr = {Axis.X: "RX", Axis.Y: "RY", Axis.Z: "R"}[init_axis] + self._stim_io.write(f"{reset_instr} {node}\n") self._rec_index += self._emit_noise_ops(ops, NoisePlacement.AFTER, default_placement) self._alive_nodes.add(node) diff --git a/graphqomb/stim_importer.py b/graphqomb/stim_importer.py index 1b64966d..b207b8f2 100644 --- a/graphqomb/stim_importer.py +++ b/graphqomb/stim_importer.py @@ -4,7 +4,6 @@ import math from dataclasses import dataclass -from graphlib import CycleError, TopologicalSorter from itertools import combinations, pairwise from pathlib import Path from typing import TYPE_CHECKING @@ -13,9 +12,8 @@ from graphqomb.circuit import Circuit, CircuitScheduleStrategy, circuit2graph from graphqomb.common import Axis, AxisMeasBasis, Sign -from graphqomb.feedforward import dag_from_flow from graphqomb.gates import CNOT, CZ, SWAP, Gate, H, Rz, S, X, Y, Z -from graphqomb.graphstate import GraphState, compose +from graphqomb.graphstate import GraphState, compose, odd_neighbors from graphqomb.qec._stim import ( PauliSupport, StimMppExtraction, @@ -37,6 +35,9 @@ _UNITARY_GATES = frozenset({"H", "S", "SQRT_Z", "S_DAG", "SQRT_Z_DAG", "X", "Y", "Z", "CX", "CNOT", "CZ", "SWAP"}) +# Stim canonicalizes Z-axis aliases when parsing circuits: RZ becomes R and +# MZ becomes M before instructions reach the importer. +_RESET_AXES = {"R": Axis.Z, "RX": Axis.X, "RY": Axis.Y} _SINGLE_PAULI_MEASUREMENT_AXES = {"M": Axis.Z, "MX": Axis.X, "MY": Axis.Y} _PAIR_PAULI_MEASUREMENT_AXES = {"MXX": "X", "MYY": "Y", "MZZ": "Z"} _PAULI_PRODUCT_MEASUREMENT_GATES = frozenset({"MPP", *_PAIR_PAULI_MEASUREMENT_AXES}) @@ -79,7 +80,9 @@ class _Fragment: @dataclass(frozen=True) class _ImportContext: stim_to_qubit: Mapping[int, int] + qubit_to_stim: Mapping[int, int] coordinate_by_stim_id: Mapping[int, tuple[float, ...]] + input_initialization_axes: Mapping[int, Axis] detector_record_indices: Sequence[frozenset[int]] logical_observable_record_indices: Mapping[int, frozenset[int]] schedule_strategy: CircuitScheduleStrategy @@ -96,6 +99,15 @@ class _IdealizedCircuit: class _AnalyzedInstruction: instruction: stim.CircuitInstruction record_indices: tuple[int, ...] + qubit_ids: frozenset[int] + + +@dataclass(frozen=True) +class _DirectMeasurement: + stim_id: int + record_index: int + axis: Axis + sign: Sign @dataclass(frozen=True) @@ -104,6 +116,9 @@ class _CircuitAnalysis: detector_record_indices: tuple[frozenset[int], ...] logical_observable_record_indices: dict[int, frozenset[int]] measurement_count: int + stim_ids: frozenset[int] + input_initialization_axes: dict[int, Axis] + direct_measurements: tuple[_DirectMeasurement, ...] def stim_file_to_pattern( @@ -159,12 +174,16 @@ def stim_circuit_to_pattern( ) -> StimImportResult: """Import a supported Stim circuit into a GraphQOMB pattern. - The importer supports Clifford unitary blocks and Pauli measurement blocks. - Stim noise instructions and measurement-error probabilities are omitted - because circuit-level noise is outside the GraphQOMB import model. Pauli - measurement blocks must be separated from unitary blocks by TICK. A direct - single-qubit measurement terminates that qubit's lifetime; other qubits may - continue, but the measured qubit cannot be used by a later operation. + The importer supports initial Pauli resets, Clifford unitary blocks, and + Pauli measurement blocks. Stim ``R``, ``RX``, and ``RY`` instructions are + imported as positive Z-, X-, and Y-eigenstate input initialization, + respectively, when they occur before any other quantum operation on the + target qubit. Stim noise instructions and measurement-error probabilities + are omitted because circuit-level noise is outside the GraphQOMB import + model. Pauli measurement blocks must be separated from unitary blocks by + TICK. A direct single-qubit measurement terminates that qubit's lifetime; + other qubits may continue, but the measured qubit cannot be used by a later + operation. Returns ------- @@ -189,11 +208,13 @@ def stim_circuit_to_pattern( msg = "DETECTOR and OBSERVABLE_INCLUDE require at least one imported measurement instruction." raise ValueError(msg) coordinate_by_stim_id = extract_qubit_coordinates(idealized.circuit, coord_dims=coord_dims) - stim_to_qubit = _stim_to_qubit_map(idealized.circuit) + stim_to_qubit = {stim_id: qubit for qubit, stim_id in enumerate(sorted(analysis.stim_ids))} qubit_to_stim = {qubit: stim_id for stim_id, qubit in stim_to_qubit.items()} context = _ImportContext( stim_to_qubit=stim_to_qubit, + qubit_to_stim=qubit_to_stim, coordinate_by_stim_id=coordinate_by_stim_id, + input_initialization_axes=analysis.input_initialization_axes, detector_record_indices=analysis.detector_record_indices, logical_observable_record_indices=analysis.logical_observable_record_indices, schedule_strategy=schedule_strategy, @@ -202,6 +223,7 @@ def stim_circuit_to_pattern( fragments = _fragments_from_blocks(analysis.blocks, context=context) fragment = _compose_fragments(fragments) + fragment = _apply_single_measurements(fragment, analysis.direct_measurements, context=context) parity_check_groups, logical_observables = _measurement_annotations_from_analysis( analysis, record_nodes=fragment.record_nodes, @@ -210,7 +232,7 @@ def stim_circuit_to_pattern( pattern = qompile( fragment.graph, fragment.xflow, - None, + {node: odd_neighbors(targets, fragment.graph) for node, targets in fragment.xflow.items()}, parity_check_group=parity_check_groups, logical_observables=logical_observables, ) @@ -271,66 +293,234 @@ def _idealize_circuit(circuit: stim.Circuit) -> _IdealizedCircuit: def _analyze_circuit(circuit: stim.Circuit) -> _CircuitAnalysis: - """Index measurement records and split a flattened circuit at TICKs. + """Validate and analyze a flattened circuit in one instruction pass. Returns ------- `_CircuitAnalysis` TICK-separated instructions and whole-circuit record annotations. - - Raises - ------ - TypeError - If a flattened circuit unexpectedly contains a repeat block. """ - blocks: list[tuple[_AnalyzedInstruction, ...]] = [] - current_block: list[_AnalyzedInstruction] = [] - detector_record_indices: list[frozenset[int]] = [] - logical_record_indices: dict[int, set[int]] = {} - measurement_count = 0 + return _CircuitAnalyzer().analyze(circuit) + + +class _CircuitAnalyzer: + """Mutable state for the single-pass Stim circuit analysis.""" + + def __init__(self) -> None: + self.blocks: list[tuple[_AnalyzedInstruction, ...]] = [] + self.current_block: list[_AnalyzedInstruction] = [] + self.detector_record_indices: list[frozenset[int]] = [] + self.logical_record_indices: dict[int, set[int]] = {} + self.measurement_count = 0 + self.stim_ids: set[int] = set() + self.input_initialization_axes: dict[int, Axis] = {} + self.direct_measurements: list[_DirectMeasurement] = [] + self.used_qubits: set[int] = set() + self.measured_qubits: set[int] = set() + self.block_has_unitary = False + self.block_has_pauli_measurement = False + + def analyze(self, circuit: stim.Circuit) -> _CircuitAnalysis: + """Analyze all instructions and return immutable analysis data. + + Returns + ------- + `_CircuitAnalysis` + Validated blocks, measurement metadata, and qubit lifetime data. + + Raises + ------ + TypeError + If a flattened circuit unexpectedly contains a repeat block. + """ + for instruction in circuit: + if not isinstance(instruction, stim.CircuitInstruction): + msg = "Flattened Stim circuit unexpectedly contains a repeat block." + raise TypeError(msg) + self._process_instruction(instruction) + self.measurement_count += instruction.num_measurements + self._finish_block() + return self._result() + + def _process_instruction(self, instruction: stim.CircuitInstruction) -> None: + instruction_qubits = self._tracked_qubits(instruction) + self.stim_ids.update(instruction_qubits) - for instruction in circuit: - if not isinstance(instruction, stim.CircuitInstruction): - msg = "Flattened Stim circuit unexpectedly contains a repeat block." - raise TypeError(msg) if instruction.name == "TICK": - blocks.append(tuple(current_block)) - current_block = [] - continue - if instruction.name == "QUBIT_COORDS": - continue - if instruction.name == "DETECTOR": - detector_record_indices.append( + self._finish_block() + elif instruction.name == "QUBIT_COORDS": + return + elif instruction.name == "DETECTOR": + self.detector_record_indices.append( record_targets_to_absolute_indices( instruction.targets_copy(), - measurement_count=measurement_count, + measurement_count=self.measurement_count, instruction_name=instruction.name, ) ) elif instruction.name == "OBSERVABLE_INCLUDE": - logical_idx = observable_index(instruction) - logical_record_indices.setdefault(logical_idx, set()).symmetric_difference_update( - record_targets_to_absolute_indices( - instruction.targets_copy(), - measurement_count=measurement_count, - instruction_name=f"OBSERVABLE_INCLUDE({logical_idx})", + self._record_logical_observable(instruction) + elif instruction.name != "MPAD": + self._record_operation(instruction, instruction_qubits) + + @staticmethod + def _tracked_qubits(instruction: stim.CircuitInstruction) -> set[int]: + if instruction.name in {"TICK", "DETECTOR", "OBSERVABLE_INCLUDE", "MPAD"}: + return set() + return _instruction_qubit_ids(instruction) + + def _finish_block(self) -> None: + self.blocks.append(tuple(self.current_block)) + self.current_block = [] + self.block_has_unitary = False + self.block_has_pauli_measurement = False + + def _record_logical_observable(self, instruction: stim.CircuitInstruction) -> None: + logical_idx = observable_index(instruction) + self.logical_record_indices.setdefault(logical_idx, set()).symmetric_difference_update( + record_targets_to_absolute_indices( + instruction.targets_copy(), + measurement_count=self.measurement_count, + instruction_name=f"OBSERVABLE_INCLUDE({logical_idx})", + ) + ) + + def _record_operation(self, instruction: stim.CircuitInstruction, instruction_qubits: set[int]) -> None: + _validate_supported_instruction(instruction) + record_indices = tuple(range(self.measurement_count, self.measurement_count + instruction.num_measurements)) + analyzed = _AnalyzedInstruction(instruction, record_indices, frozenset(instruction_qubits)) + self.current_block.append(analyzed) + + is_unitary = instruction.name in _UNITARY_GATES + is_pauli_measurement = instruction.name == "MPP" or instruction.name in _SINGLE_PAULI_MEASUREMENT_AXES + self._validate_block_separation( + is_unitary=is_unitary, + is_pauli_measurement=is_pauli_measurement, + ) + self._record_lifetime( + instruction, + instruction_qubits, + is_quantum_operation=is_unitary or is_pauli_measurement, + ) + + if instruction.name in _SINGLE_PAULI_MEASUREMENT_AXES: + self.direct_measurements.extend(_direct_measurements_from_instruction(analyzed)) + self.measured_qubits.update(instruction_qubits) + + def _validate_block_separation(self, *, is_unitary: bool, is_pauli_measurement: bool) -> None: + self.block_has_unitary |= is_unitary + self.block_has_pauli_measurement |= is_pauli_measurement + if self.block_has_unitary and self.block_has_pauli_measurement: + msg = "Pauli measurement instructions must be separated from unitary gate instructions by TICK." + raise ValueError(msg) + + def _record_lifetime( + self, + instruction: stim.CircuitInstruction, + instruction_qubits: set[int], + *, + is_quantum_operation: bool, + ) -> None: + reset_axis = _RESET_AXES.get(instruction.name) + if reset_axis is not None: + reset_after_use = self.used_qubits & instruction_qubits + if reset_after_use: + msg = ( + f"Stim reset instruction {instruction.name} targets qubit(s) {sorted(reset_after_use)} " + "after another quantum operation; only initial resets are supported." ) + raise ValueError(msg) + self.input_initialization_axes.update(dict.fromkeys(instruction_qubits, reset_axis)) + return + if not is_quantum_operation: + return + + reused_qubits = self.measured_qubits & instruction_qubits + if reused_qubits: + msg = ( + f"Stim qubit(s) {sorted(reused_qubits)} are used after a single-qubit measurement; " + "single-qubit measurements terminate those qubit lifetimes." ) - elif instruction.name != "MPAD": - record_indices = tuple(range(measurement_count, measurement_count + instruction.num_measurements)) - current_block.append(_AnalyzedInstruction(instruction, record_indices)) - - measurement_count += instruction.num_measurements - - blocks.append(tuple(current_block)) - return _CircuitAnalysis( - blocks=tuple(blocks), - detector_record_indices=tuple(detector_record_indices), - logical_observable_record_indices={ - logical_idx: frozenset(records) for logical_idx, records in sorted(logical_record_indices.items()) - }, - measurement_count=measurement_count, - ) + raise ValueError(msg) + self.used_qubits.update(instruction_qubits) + + def _result(self) -> _CircuitAnalysis: + return _CircuitAnalysis( + blocks=tuple(self.blocks), + detector_record_indices=tuple(self.detector_record_indices), + logical_observable_record_indices={ + logical_idx: frozenset(records) for logical_idx, records in sorted(self.logical_record_indices.items()) + }, + measurement_count=self.measurement_count, + stim_ids=frozenset(self.stim_ids), + input_initialization_axes=self.input_initialization_axes, + direct_measurements=tuple(self.direct_measurements), + ) + + +def _instruction_qubit_ids(instruction: stim.CircuitInstruction) -> set[int]: + """Return all qubit ids targeted by a Stim instruction. + + Returns + ------- + `set`[`int`] + Targeted qubit ids. + """ + return {int(target.qubit_value) for target in instruction.targets_copy() if target.qubit_value is not None} + + +def _validate_supported_instruction(instruction: stim.CircuitInstruction) -> None: + """Reject instructions outside the supported operation set. + + Raises + ------ + ValueError + If the instruction is unsupported. + """ + if ( + instruction.name in _UNITARY_GATES + or instruction.name in _RESET_AXES + or instruction.name in _SINGLE_PAULI_MEASUREMENT_AXES + or instruction.name == "MPP" + ): + return + msg = f"Unsupported Stim instruction(s): {instruction.name}." + raise ValueError(msg) + + +def _direct_measurements_from_instruction( + analyzed: _AnalyzedInstruction, +) -> list[_DirectMeasurement]: + """Normalize one direct-measurement instruction into per-qubit events. + + Returns + ------- + `list`[`_DirectMeasurement`] + One normalized event per measured qubit. + + Raises + ------ + ValueError + If target and record counts differ or a qubit target is repeated. + """ + instruction = analyzed.instruction + targets = instruction.targets_copy() + if len(targets) != len(analyzed.record_indices): + msg = f"{instruction.name} target count does not match its measurement-record count." + raise ValueError(msg) + + measurements: list[_DirectMeasurement] = [] + seen_qubits: set[int] = set() + axis = _SINGLE_PAULI_MEASUREMENT_AXES[instruction.name] + for target, record_index in zip(targets, analyzed.record_indices, strict=True): + stim_id = plain_qubit_target(target, instruction.name) + if stim_id in seen_qubits: + msg = f"{instruction.name} measures qubit {stim_id} more than once in one instruction." + raise ValueError(msg) + seen_qubits.add(stim_id) + sign = Sign.MINUS if target.is_inverted_result_target else Sign.PLUS + measurements.append(_DirectMeasurement(stim_id, record_index, axis, sign)) + return measurements def _append_ideal_pauli_measurements( @@ -374,224 +564,191 @@ def _fragments_from_blocks( *, context: _ImportContext, ) -> list[_Fragment]: - _validate_blocks(blocks) - _validate_single_measurement_lifetimes(blocks) - fragments = [_identity_fragment(context)] - mpp_layer_index = 0 + z_base = 0 + live_stim_ids = set(context.stim_to_qubit) for block in blocks: + directly_measured_stim_ids = { + stim_id + for analyzed in block + if analyzed.instruction.name in _SINGLE_PAULI_MEASUREMENT_AXES + for stim_id in analyzed.qubit_ids + } unitary_instructions = tuple( analyzed.instruction for analyzed in block if analyzed.instruction.name in _UNITARY_GATES ) if unitary_instructions: - fragments.append(_unitary_fragment(unitary_instructions, context=context)) + fragment, z_base = _unitary_fragment( + unitary_instructions, + live_stim_ids=live_stim_ids, + z_base=z_base, + context=context, + ) + fragments.append(fragment) else: - measurement_fragments = _measurement_fragments_from_block( + mpp_stim_ids = { + stim_id for analyzed in block if analyzed.instruction.name == "MPP" for stim_id in analyzed.qubit_ids + } + mpp_fragment = _mpp_fragment_from_block( block, - mpp_layer_index=mpp_layer_index, + z_base=z_base, + io_stim_ids=(live_stim_ids - directly_measured_stim_ids) | mpp_stim_ids, context=context, ) - fragments.extend(measurement_fragments) - if any(analyzed.instruction.name == "MPP" for analyzed in block): - mpp_layer_index += sum( - len(analyzed.record_indices) for analyzed in block if analyzed.instruction.name == "MPP" - ) - - return fragments - + if mpp_fragment is not None: + fragments.append(mpp_fragment) + z_base += 2 + elif directly_measured_stim_ids: + continuing_stim_ids = live_stim_ids - directly_measured_stim_ids + if continuing_stim_ids: + z_base += 1 + fragments.append(_relocation_fragment(continuing_stim_ids, z=z_base, context=context)) -def _validate_blocks(blocks: Sequence[Sequence[_AnalyzedInstruction]]) -> None: - """Validate supported instructions and required TICK separation. + live_stim_ids.difference_update(directly_measured_stim_ids) - Raises - ------ - ValueError - If an instruction is unsupported or a block mixes unitary gates with - Pauli measurements. - """ - for block in blocks: - unsupported = [ - analyzed.instruction.name - for analyzed in block - if ( - analyzed.instruction.name not in _UNITARY_GATES - and analyzed.instruction.name not in _SINGLE_PAULI_MEASUREMENT_AXES - and analyzed.instruction.name != "MPP" - ) - ] - if unsupported: - msg = f"Unsupported Stim instruction(s): {', '.join(sorted(set(unsupported)))}." - raise ValueError(msg) - - has_unitary = any(analyzed.instruction.name in _UNITARY_GATES for analyzed in block) - has_pauli_measurement = any( - analyzed.instruction.name == "MPP" or analyzed.instruction.name in _SINGLE_PAULI_MEASUREMENT_AXES - for analyzed in block - ) - if has_unitary and has_pauli_measurement: - msg = "Pauli measurement instructions must be separated from unitary gate instructions by TICK." - raise ValueError(msg) - - -def _validate_single_measurement_lifetimes( - blocks: Sequence[Sequence[_AnalyzedInstruction]], -) -> None: - """Reject quantum operations after a directly measured qubit terminates. - - Raises - ------ - ValueError - If a quantum operation reuses a directly measured qubit. - """ - measured_qubits: set[int] = set() - for block in blocks: - for analyzed in block: - instruction = analyzed.instruction - if ( - instruction.name not in _UNITARY_GATES - and instruction.name != "MPP" - and instruction.name not in _SINGLE_PAULI_MEASUREMENT_AXES - ): - continue - - instruction_qubits = { - int(target.qubit_value) for target in instruction.targets_copy() if target.qubit_value is not None - } - reused_qubits = measured_qubits & instruction_qubits - if reused_qubits: - msg = ( - f"Stim qubit(s) {sorted(reused_qubits)} are used after a single-qubit measurement; " - "single-qubit measurements terminate those qubit lifetimes." - ) - raise ValueError(msg) - if instruction.name in _SINGLE_PAULI_MEASUREMENT_AXES: - measured_qubits.update(instruction_qubits) + return fragments -def _measurement_fragments_from_block( +def _mpp_fragment_from_block( block: Sequence[_AnalyzedInstruction], *, - mpp_layer_index: int, + z_base: int, + io_stim_ids: set[int], context: _ImportContext, -) -> list[_Fragment]: - """Build direct and product measurement fragments in record order. +) -> _Fragment | None: + """Build the combined MPP fragment in a measurement block. Returns ------- - `list`[`_Fragment`] - Measurement fragments in source order. + `_Fragment` | `None` + Combined MPP fragment, or `None` when the block has no MPP instruction. """ - fragments: list[_Fragment] = [] mpp_items = tuple(analyzed for analyzed in block if analyzed.instruction.name == "MPP") - mpp_added = False - for analyzed in block: - instruction = analyzed.instruction - if instruction.name == "MPP" and not mpp_added: - fragments.append( - _mpp_fragment( - mpp_items, - mpp_layer_index=mpp_layer_index, - context=context, - ) - ) - mpp_added = True - elif instruction.name in _SINGLE_PAULI_MEASUREMENT_AXES: - fragments.append( - _single_measurement_fragment( - instruction, - record_indices=analyzed.record_indices, - context=context, - ) - ) - return fragments - - -def _single_measurement_fragment( - instruction: stim.CircuitInstruction, - *, - record_indices: Sequence[int], - context: _ImportContext, -) -> _Fragment: - """Build a fragment by assigning a basis directly to each measured node. + if not mpp_items: + return None + return _mpp_fragment( + mpp_items, + z_base=z_base, + io_stim_ids=io_stim_ids, + context=context, + ) - Returns - ------- - `_Fragment` - Direct-measurement graph fragment and record-to-node mapping. - - Raises - ------ - ValueError - If target counts differ or a qubit is repeated in the instruction. - """ - targets = instruction.targets_copy() - if len(targets) != len(record_indices): - msg = f"{instruction.name} target count does not match its measurement-record count." - raise ValueError(msg) +def _identity_fragment(context: _ImportContext) -> _Fragment: graph = GraphState() - record_nodes: dict[int, int] = {} - seen_qubits: set[int] = set() - axis = _SINGLE_PAULI_MEASUREMENT_AXES[instruction.name] - for target, record_index in zip(targets, record_indices, strict=True): - stim_id = plain_qubit_target(target, instruction.name) - if stim_id in seen_qubits: - msg = f"{instruction.name} measures qubit {stim_id} more than once in one instruction." - raise ValueError(msg) - seen_qubits.add(stim_id) - - node = graph.add_node(coordinate=context.coordinate_by_stim_id.get(stim_id)) - qubit_index = context.stim_to_qubit[stim_id] - graph.register_input(node, qubit_index) + for stim_id, qubit_index in sorted(context.stim_to_qubit.items()): + coord = context.coordinate_by_stim_id.get(stim_id) + node = graph.add_node(coordinate=_coordinate_at_z(coord, 0) if coord is not None else None) + graph.register_input( + node, + qubit_index, + init_axis=context.input_initialization_axes.get(stim_id, Axis.X), + ) graph.register_output(node, qubit_index) - sign = Sign.MINUS if target.is_inverted_result_target else Sign.PLUS - graph.assign_meas_basis(node, AxisMeasBasis(axis, sign)) - record_nodes[record_index] = node - return _Fragment( graph=graph, xflow={}, - record_nodes=record_nodes, + record_nodes={}, ) -def _identity_fragment(context: _ImportContext) -> _Fragment: - graph = GraphState() - for stim_id, qubit_index in sorted(context.stim_to_qubit.items()): - node = graph.add_node(coordinate=context.coordinate_by_stim_id.get(stim_id)) - graph.register_input(node, qubit_index) - graph.register_output(node, qubit_index) - return _Fragment(graph=graph, xflow={}, record_nodes={}) - - def _unitary_fragment( block: Sequence[stim.CircuitInstruction], *, + live_stim_ids: set[int], + z_base: int, context: _ImportContext, -) -> _Fragment: - active_stim_ids = sorted( - {plain_qubit_target(target, instruction.name) for instruction in block for target in instruction.targets_copy()} - ) - stim_to_local = {stim_id: local_index for local_index, stim_id in enumerate(active_stim_ids)} +) -> tuple[_Fragment, int]: + ordered_stim_ids = sorted(live_stim_ids) + stim_to_local = {stim_id: local_index for local_index, stim_id in enumerate(ordered_stim_ids)} local_to_global = {local_index: context.stim_to_qubit[stim_id] for stim_id, local_index in stim_to_local.items()} - circuit = Circuit(len(active_stim_ids)) + circuit = Circuit(len(ordered_stim_ids)) for instruction in block: _append_unitary_instruction(circuit, instruction, stim_to_local) - local_graph, local_xflow, _scheduler = circuit2graph(circuit, schedule_strategy=context.schedule_strategy) + local_graph, local_xflow, _ = circuit2graph(circuit, schedule_strategy=context.schedule_strategy) graph, node_map = _copy_graph_with_qindices(local_graph, local_to_global) - _apply_stim_coordinates( + xflow = _remap_flow(local_xflow, node_map) + z_span = _apply_unitary_coordinates( graph, - stim_to_qubit=context.stim_to_qubit, - coordinate_by_stim_id=context.coordinate_by_stim_id, + xflow, + z_base=z_base, + context=context, ) - return _Fragment( - graph=graph, - xflow=_remap_flow(local_xflow, node_map), - record_nodes={}, + return ( + _Fragment( + graph=graph, + xflow=xflow, + record_nodes={}, + ), + z_base + z_span, ) +def _apply_unitary_coordinates( + graph: GraphState, + xflow: Mapping[int, set[int]], + *, + z_base: int, + context: _ImportContext, +) -> int: + """Place a transpiled gate block between common input and output Z layers. + + Qubit chains with gates are stretched across the block's maximum transpiled + depth. An idle qubit remains a single input/output node and is relocated to + the common output layer without adding graph nodes or edges. + + Returns + ------- + `int` + Z span occupied by the gate block. + + Raises + ------ + ValueError + If the transpiled unitary graph does not consist of data-wire chains. + """ + output_node_by_qubit = {qubit: node for node, qubit in graph.output_node_indices.items()} + chains: dict[int, list[int]] = {} + chain_nodes: set[int] = set() + for input_node, qubit in graph.input_node_indices.items(): + output_node = output_node_by_qubit[qubit] + chain = [input_node] + visited = {input_node} + while chain[-1] != output_node: + targets = xflow.get(chain[-1]) + if targets is None or len(targets) != 1: + msg = f"Transpiled unitary wire for qubit {qubit} is not a single X-flow chain." + raise ValueError(msg) + next_node = next(iter(targets)) + if next_node in visited: + msg = f"Transpiled unitary wire for qubit {qubit} contains an X-flow cycle." + raise ValueError(msg) + chain.append(next_node) + visited.add(next_node) + chains[qubit] = chain + chain_nodes.update(chain) + + if chain_nodes != graph.nodes: + msg = "Transpiled unitary graph contains a node outside its data-wire X-flow chains." + raise ValueError(msg) + + max_chain_depth = max((len(chain) - 1 for chain in chains.values()), default=0) + z_span = max(1, max_chain_depth) + for qubit, chain in chains.items(): + coord = context.coordinate_by_stim_id.get(context.qubit_to_stim[qubit]) + if coord is None: + continue + depth = len(chain) - 1 + if depth == 0: + graph.set_coordinate(chain[0], _coordinate_at_z(coord, z_base + z_span)) + continue + for layer, node in enumerate(chain): + z = z_base + z_span * layer / depth + graph.set_coordinate(node, _coordinate_at_z(coord, z)) + return z_span + + def _copy_graph_with_qindices( graph: GraphState, local_to_global: Mapping[int, int], @@ -620,6 +777,10 @@ def _copy_graph_with_qindices( return copied, node_map +# TODO(masa10-f): Cancel repeated CZ pairs within one TICK block(#233) +# CZ*CZ = identity, but repeated pairs currently fail in graph construction with +# "Edge already exists" when neither wire advances between the two CZ instructions. +# Deferred to the full Stim parser PR. def _append_unitary_instruction( circuit: Circuit, instruction: stim.CircuitInstruction, @@ -642,7 +803,8 @@ def _append_unitary_instruction( def _mpp_fragment( block: Sequence[_AnalyzedInstruction], *, - mpp_layer_index: int, + z_base: int, + io_stim_ids: set[int], context: _ImportContext, ) -> _Fragment: supports = tuple( @@ -657,32 +819,14 @@ def _mpp_fragment( detector_record_indices=context.detector_record_indices, logical_observable_record_indices=context.logical_observable_record_indices, ) - z_base = 2 * mpp_layer_index fragment = _mpp_graph_fragment( extraction, record_indices=record_indices, z_base=z_base, + io_stim_ids=io_stim_ids, context=context, ) - if _has_causal_flow(fragment): - return _with_mpp_extraction(fragment, extraction) - - serialized_fragments = [ - _mpp_graph_fragment( - stim_mpp_extraction_from_records( - (support,), - (record_index,), - coordinate_by_stim_id=context.coordinate_by_stim_id, - detector_record_indices=context.detector_record_indices, - logical_observable_record_indices=context.logical_observable_record_indices, - ), - record_indices=(record_index,), - z_base=z_base + 2 * row, - context=context, - ) - for row, (support, record_index) in enumerate(zip(supports, record_indices, strict=True)) - ] - return _with_mpp_extraction(_compose_fragments(serialized_fragments), extraction) + return _with_mpp_extraction(fragment, extraction) def _validate_commuting_mpp_supports(supports: Sequence[PauliSupport]) -> None: @@ -699,6 +843,7 @@ def _mpp_graph_fragment( *, record_indices: Sequence[int], z_base: int, + io_stim_ids: set[int], context: _ImportContext, ) -> _Fragment: qubit_indices = {column: context.stim_to_qubit[stim_id] for column, stim_id in extraction.column_to_stim.items()} @@ -709,6 +854,13 @@ def _mpp_graph_fragment( data_as_io=True, qubit_indices=qubit_indices, ) + active_stim_ids = set(extraction.stim_to_column) + _add_relocated_io_nodes( + result.graph, + io_stim_ids - active_stim_ids, + z=z_base + 2, + context=context, + ) xflow = _mpp_flow(result) if len(record_indices) != len(result.ancilla_nodes): msg = "Imported MPP record count does not match the generated ancilla-node count." @@ -720,12 +872,29 @@ def _mpp_graph_fragment( ) -def _has_causal_flow(fragment: _Fragment) -> bool: - try: - tuple(TopologicalSorter(dag_from_flow(fragment.graph, fragment.xflow)).static_order()) - except CycleError: - return False - return True +def _relocation_fragment(stim_ids: set[int], *, z: int, context: _ImportContext) -> _Fragment: + graph = GraphState() + _add_relocated_io_nodes(graph, stim_ids, z=z, context=context) + return _Fragment(graph=graph, xflow={}, record_nodes={}) + + +def _add_relocated_io_nodes( + graph: GraphState, + stim_ids: set[int], + *, + z: int, + context: _ImportContext, +) -> None: + for stim_id in sorted(stim_ids): + coord = context.coordinate_by_stim_id.get(stim_id) + node = graph.add_node(coordinate=_coordinate_at_z(coord, z) if coord is not None else None) + qubit = context.stim_to_qubit[stim_id] + graph.register_input(node, qubit) + graph.register_output(node, qubit) + + +def _coordinate_at_z(coord: tuple[float, ...], z: float) -> tuple[float, float, float]: + return (float(coord[0]), float(coord[1]), float(z)) def _with_mpp_extraction(fragment: _Fragment, extraction: StimMppExtraction) -> _Fragment: @@ -741,25 +910,11 @@ def _mpp_flow( result: StabilizerGraphStateBuildResult, ) -> dict[int, set[int]]: xflow: dict[int, set[int]] = {} - measured_nodes_by_qubit: dict[int, list[int]] = {} for qubit in sorted({key[0] for key in result.data_nodes}): layer_nodes = [node for (data_qubit, _layer), node in sorted(result.data_nodes.items()) if data_qubit == qubit] - measured_nodes_by_qubit[qubit] = [node for node in layer_nodes if node in result.graph.meas_bases] for current_node, next_node in pairwise(layer_nodes): if current_node in result.graph.meas_bases: xflow[current_node] = {next_node} - for ancilla_node in result.ancilla_nodes.values(): - correction_nodes = {ancilla_node} - for measured_nodes in measured_nodes_by_qubit.values(): - for earlier_node, later_node in pairwise(measured_nodes): - # Type I Y support touches both data-measurement layers. Including - # the later data stabilizer cancels the backward dependency in - # the automatically derived odd-neighborhood zflow. - if result.graph.has_edge(ancilla_node, earlier_node) and result.graph.has_edge( - ancilla_node, later_node - ): - correction_nodes.add(later_node) - xflow[ancilla_node] = correction_nodes return xflow @@ -777,6 +932,36 @@ def _compose_fragments(fragments: Sequence[_Fragment]) -> _Fragment: return current +def _apply_single_measurements( + fragment: _Fragment, + direct_measurements: Sequence[_DirectMeasurement], + *, + context: _ImportContext, +) -> _Fragment: + """Assign direct measurements to existing data-lane output nodes. + + Returns + ------- + `_Fragment` + Fragment with terminal measurement bases and record mappings applied. + + """ + output_node_by_qubit = {qubit: node for node, qubit in fragment.graph.output_node_indices.items()} + record_nodes = dict(fragment.record_nodes) + + for measurement in direct_measurements: + node = output_node_by_qubit[context.stim_to_qubit[measurement.stim_id]] + fragment.graph.assign_meas_basis(node, AxisMeasBasis(measurement.axis, measurement.sign)) + record_nodes[measurement.record_index] = node + + return _Fragment( + graph=fragment.graph, + xflow=fragment.xflow, + record_nodes=record_nodes, + mpp_extractions=fragment.mpp_extractions, + ) + + def _measurement_annotations_from_analysis( analysis: _CircuitAnalysis, *, @@ -828,32 +1013,3 @@ def _remap_node_set(nodes: set[int], node_map: Mapping[int, int]) -> set[int]: def _remap_record_nodes(record_nodes: Mapping[int, int], node_map: Mapping[int, int]) -> dict[int, int]: return {record_index: node_map[node] for record_index, node in record_nodes.items()} - - -def _apply_stim_coordinates( - graph: GraphState, - *, - stim_to_qubit: Mapping[int, int], - coordinate_by_stim_id: Mapping[int, tuple[float, ...]], -) -> None: - qubit_to_stim = {qubit: stim_id for stim_id, qubit in stim_to_qubit.items()} - for node, q_index in graph.input_node_indices.items() | graph.output_node_indices.items(): - stim_id = qubit_to_stim[q_index] - coord = coordinate_by_stim_id.get(stim_id) - if coord is not None: - graph.set_coordinate(node, coord) - - -def _stim_to_qubit_map(circuit: stim.Circuit) -> dict[int, int]: - stim_ids: set[int] = set() - for instruction in circuit: - if not isinstance(instruction, stim.CircuitInstruction): - msg = "Flattened Stim circuit unexpectedly contains a repeat block." - raise TypeError(msg) - if instruction.name in {"TICK", "DETECTOR", "OBSERVABLE_INCLUDE", "MPAD"}: - continue - for target in instruction.targets_copy(): - qubit_value = target.qubit_value - if qubit_value is not None: - stim_ids.add(int(qubit_value)) - return {stim_id: qubit for qubit, stim_id in enumerate(sorted(stim_ids))} diff --git a/tests/test_feedforward.py b/tests/test_feedforward.py index 1cd0ce29..57212937 100644 --- a/tests/test_feedforward.py +++ b/tests/test_feedforward.py @@ -95,6 +95,13 @@ def test_dag_from_flow_cycle_detection() -> None: check_dag(dag) +def test_check_dag_detects_cycle_longer_than_two_nodes() -> None: + dag = {0: {1}, 1: {2}, 2: {0}} + + with pytest.raises(ValueError, match="Cycle detected in the graph:"): + check_dag(dag) + + def test_check_flow_false_for_cycle() -> None: graphstate, node1, node2 = two_node_graph() cyclic_flow = {node1: node2, node2: node1} diff --git a/tests/test_graph_compose.py b/tests/test_graph_compose.py index 510a8e00..d0a7e5ac 100644 --- a/tests/test_graph_compose.py +++ b/tests/test_graph_compose.py @@ -4,7 +4,7 @@ import pytest -from graphqomb.common import Plane, PlannerMeasBasis +from graphqomb.common import Axis, Plane, PlannerMeasBasis from graphqomb.graphstate import BaseGraphState, GraphState, compose @@ -170,6 +170,36 @@ def test_compose_preserves_measurement_bases() -> None: assert composed.meas_bases[mapped_node2] == meas_basis +def test_compose_preserves_surviving_input_initialization_axes() -> None: + """Composition preserves initialization axes for inputs that remain inputs.""" + graph1 = GraphState() + g1_in = graph1.add_node() + g1_out = graph1.add_node() + graph1.add_edge(g1_in, g1_out) + graph1.register_input(g1_in, 0, init_axis=Axis.Y) + graph1.register_output(g1_out, 1) + graph1.assign_meas_basis(g1_in, PlannerMeasBasis(Plane.XY, 0.0)) + + graph2 = GraphState() + g2_in_connected = graph2.add_node() + g2_in_survives = graph2.add_node() + g2_out = graph2.add_node() + graph2.add_edge(g2_in_connected, g2_out) + graph2.add_edge(g2_in_survives, g2_out) + graph2.register_input(g2_in_connected, 1, init_axis=Axis.Z) + graph2.register_input(g2_in_survives, 2, init_axis=Axis.Z) + graph2.register_output(g2_out, 3) + graph2.assign_meas_basis(g2_in_connected, PlannerMeasBasis(Plane.XY, 0.0)) + graph2.assign_meas_basis(g2_in_survives, PlannerMeasBasis(Plane.XY, 0.0)) + + composed, node_map1, node_map2 = compose(graph1, graph2) + + assert composed.input_initialization_axes == { + node_map1[g1_in]: Axis.Y, + node_map2[g2_in_survives]: Axis.Z, + } + + def test_compose_full_connection() -> None: """Test compose where all outputs of graph1 connect to inputs of graph2.""" # Create graph1: input [0] -> output [1, 2] diff --git a/tests/test_graphstate.py b/tests/test_graphstate.py index acd80f72..acead690 100644 --- a/tests/test_graphstate.py +++ b/tests/test_graphstate.py @@ -3,11 +3,12 @@ from __future__ import annotations import math +from typing import Any import numpy as np import pytest -from graphqomb.common import Plane, PlannerMeasBasis +from graphqomb.common import Axis, Plane, PlannerMeasBasis from graphqomb.euler import LocalClifford from graphqomb.graphstate import GraphState, bipartite_edges, compose, odd_neighbors @@ -102,6 +103,33 @@ def test_add_node_input_output(graph: GraphState) -> None: assert graph.output_node_indices[node_index] == q_index +def test_register_input_defaults_to_x_initialization_axis(graph: GraphState) -> None: + """Input nodes default to positive X-basis initialization.""" + node_index = graph.add_node() + + graph.register_input(node_index, 0) + + assert graph.input_initialization_axes == {node_index: Axis.X} + + +def test_register_input_accepts_pauli_initialization_axis(graph: GraphState) -> None: + """Input nodes can be initialized in a positive Pauli eigenstate.""" + node_index = graph.add_node() + + graph.register_input(node_index, 0, init_axis=Axis.Y) + + assert graph.input_initialization_axes == {node_index: Axis.Y} + + +def test_register_input_rejects_non_axis_initialization(graph: GraphState) -> None: + """Input registration rejects values outside the Axis enum.""" + node_index = graph.add_node() + invalid_axis: Any = "X" + + with pytest.raises(TypeError, match="Input initialization axis must be one of"): + graph.register_input(node_index, 0, init_axis=invalid_axis) + + def test_ensure_node_exists_raises(graph: GraphState) -> None: """Test ensuring a node exists in the graph.""" with pytest.raises(ValueError, match="Node does not exist node=1"): diff --git a/tests/test_graphstate_bulk_init.py b/tests/test_graphstate_bulk_init.py index 17631110..d14acc3b 100644 --- a/tests/test_graphstate_bulk_init.py +++ b/tests/test_graphstate_bulk_init.py @@ -6,7 +6,7 @@ import pytest -from graphqomb.common import Plane, PlannerMeasBasis +from graphqomb.common import Axis, Plane, PlannerMeasBasis from graphqomb.graphstate import GraphState @@ -76,6 +76,43 @@ def test_from_graph_with_multiple_inputs_outputs() -> None: assert gs.output_node_indices == {node_map[outputs[0]]: 0, node_map[outputs[1]]: 1} +def test_from_graph_with_input_initialization_axes() -> None: + """Test from_graph() preserves specified input initialization axes.""" + nodes = ["a", "b", "c"] + edges = [("a", "c"), ("b", "c")] + inputs = ["a", "b"] + + gs, node_map = GraphState.from_graph( + nodes=nodes, + edges=edges, + inputs=inputs, + input_initialization_axes={"a": Axis.Y, "b": Axis.Z}, + ) + + assert gs.input_initialization_axes == {node_map["a"]: Axis.Y, node_map["b"]: Axis.Z} + + +def test_from_graph_rejects_initialization_axis_for_non_input_node() -> None: + """Input initialization axes cannot silently target non-input nodes.""" + with pytest.raises(ValueError, match="specified for non-input node"): + GraphState.from_graph( + nodes=["input", "not-input"], + edges=[], + inputs=["input"], + input_initialization_axes={"not-input": Axis.Z}, + ) + + +def test_from_graph_rejects_initialization_axes_without_inputs() -> None: + """Initialization axes require corresponding input registrations.""" + with pytest.raises(ValueError, match="specified for non-input node"): + GraphState.from_graph( + nodes=["node"], + edges=[], + input_initialization_axes={"node": Axis.Y}, + ) + + def test_from_graph_with_meas_bases() -> None: """Test from_graph() with measurement bases.""" nodes = ["a", "b", "c"] @@ -213,6 +250,19 @@ def test_from_base_graph_state_preserves_indices() -> None: assert dst.output_node_indices[node_map[n1]] == 10 +def test_from_base_graph_state_preserves_input_initialization_axes() -> None: + """Test from_base_graph_state() preserves input initialization axes.""" + src = GraphState() + n0 = src.add_node() + n1 = src.add_node() + src.register_input(n0, 0, init_axis=Axis.Y) + src.register_input(n1, 1, init_axis=Axis.Z) + + dst, node_map = GraphState.from_base_graph_state(src) + + assert dst.input_initialization_axes == {node_map[n0]: Axis.Y, node_map[n1]: Axis.Z} + + def test_from_base_graph_state_independence() -> None: """Test that modifications to copied graph don't affect original.""" # Create source graph diff --git a/tests/test_ptn_format.py b/tests/test_ptn_format.py index ad14ec6d..60521566 100644 --- a/tests/test_ptn_format.py +++ b/tests/test_ptn_format.py @@ -104,6 +104,7 @@ def assert_pattern_equivalent(actual: Pattern, expected: Pattern) -> None: assert actual.input_node_indices == expected.input_node_indices assert actual.output_node_indices == expected.output_node_indices assert actual.input_coordinates == expected.input_coordinates + assert actual.input_initialization_axes == expected.input_initialization_axes assert actual.pauli_frame.xflow == expected.pauli_frame.xflow assert actual.pauli_frame.zflow == expected.pauli_frame.zflow assert actual.pauli_frame.parity_check_group == expected.pauli_frame.parity_check_group @@ -116,13 +117,102 @@ def test_dumps_basic() -> None: pattern = create_simple_pattern() ptn_str = dumps(pattern) - assert ".version 1" in ptn_str + assert ".version 2" in ptn_str assert ".input" in ptn_str assert ".output" in ptn_str assert "#======== QUANTUM ========" in ptn_str assert "#======== CLASSICAL ========" in ptn_str +def test_dumps_omits_default_input_basis() -> None: + """Default X-basis input initialization is implicit in .ptn output.""" + pattern = create_simple_pattern() + + assert ".input_basis" not in dumps(pattern) + + +def test_ptn_roundtrip_preserves_input_initialization_axes() -> None: + """Non-default input initialization axes survive .ptn roundtrip.""" + graph = GraphState() + in0 = graph.add_node() + in1 = graph.add_node() + out0 = graph.add_node() + out1 = graph.add_node() + + graph.register_input(in0, 0, init_axis=Axis.Y) + graph.register_input(in1, 1, init_axis=Axis.Z) + graph.register_output(out0, 0) + graph.register_output(out1, 1) + graph.add_edge(in0, out0) + graph.add_edge(in1, out1) + graph.assign_meas_basis(in0, AxisMeasBasis(Axis.X, Sign.PLUS)) + graph.assign_meas_basis(in1, AxisMeasBasis(Axis.X, Sign.PLUS)) + + pattern = qompile(graph, {in0: {out0}, in1: {out1}}) + ptn_str = dumps(pattern) + loaded = loads(ptn_str) + + assert ".input_basis" in ptn_str + assert loaded.input_initialization_axes == {in0: Axis.Y, in1: Axis.Z} + + +def test_ptn_loads_legacy_input_without_input_basis_as_x() -> None: + """Existing .ptn files without .input_basis default inputs to X initialization.""" + ptn_str = """# GraphQOMB Pattern Format v1 +.version 1 +.input 0:0 +.output 1:0 + +[0] +E 0 1 +M 0 X + + +.xflow 0 -> 1 +""" + + loaded = loads(ptn_str) + + assert loaded.input_initialization_axes == {0: Axis.X} + + +def test_ptn_load_rejects_input_basis_for_non_input_node() -> None: + """Input basis directives can only reference input nodes.""" + ptn_str = """# GraphQOMB Pattern Format v2 +.version 2 +.input 0:0 +.input_basis 1:Z +.output 1:0 + +[0] +E 0 1 +M 0 X + + +.xflow 0 -> 1 +""" + + with pytest.raises(ValueError, match="Input basis specified for non-input node"): + loads(ptn_str) + + +def test_ptn_load_rejects_input_basis_in_legacy_version() -> None: + """Version 1 .ptn files cannot use the version 2 input-basis directive.""" + ptn_str = """# GraphQOMB Pattern Format v1 +.version 1 +.input 0:0 +.input_basis 0:Z +.output 1:0 + +[0] +E 0 1 +M 0 X + + +.xflow 0 -> 1 +""" + + with pytest.raises(ValueError, match=r"\.input_basis requires \.ptn version 2"): + loads(ptn_str) + + def test_dumps_contains_commands() -> None: """Test that dumps includes all command types.""" pattern = create_simple_pattern() diff --git a/tests/test_qec.py b/tests/test_qec.py index 1d2b6bb7..abdbfdfd 100644 --- a/tests/test_qec.py +++ b/tests/test_qec.py @@ -64,6 +64,84 @@ def test_build_graph_state_assigns_x_measurement_to_all_nodes() -> None: _assert_axis_meas_basis(result.graph.meas_bases[node], Axis.X) +@pytest.mark.parametrize( + ("stabilizer_row", "expected_axis"), + [ + ([1, 0, 0, 0, 0, 0], Axis.X), # Zero Y supports. + ([1, 0, 0, 1, 0, 0], Axis.Y), # One Y support. + ([1, 1, 0, 1, 1, 0], Axis.X), # Two Y supports. + ([1, 1, 1, 1, 1, 1], Axis.Y), # Three Y supports. + ], +) +def test_build_graph_state_type_i_selects_ancilla_basis_from_y_support_parity( + stabilizer_row: list[int], + expected_axis: Axis, +) -> None: + code = StabilizerCode(_matrix([stabilizer_row])) + + result = build_graph_state(code, y_foliation=YFoliation.TYPE_I) + + _assert_axis_meas_basis(result.graph.meas_bases[result.ancilla_nodes[0]], expected_axis) + + +@pytest.mark.parametrize("y_foliation", [YFoliation.TYPE_I, YFoliation.TYPE_II]) +@pytest.mark.parametrize( + "stabilizer_matrix", + [ + pytest.param( + [ + [1, 0, 0, 1], # X0 Z1. + [0, 1, 1, 0], # Z0 X1. + ], + id="xz-order", + ), + pytest.param( + [ + [1, 0, 1, 1], # Y0 Z1. + [0, 1, 1, 1], # Z0 Y1. + ], + id="yz-order", + ), + ], +) +def test_build_graph_state_connects_ancillas_for_odd_twisted_order_pairs( + y_foliation: YFoliation, + stabilizer_matrix: list[list[int]], +) -> None: + result = build_graph_state(StabilizerCode(_matrix(stabilizer_matrix)), y_foliation=y_foliation) + + assert result.graph.has_edge(result.ancilla_nodes[0], result.ancilla_nodes[1]) + + +@pytest.mark.parametrize("y_foliation", [YFoliation.TYPE_I, YFoliation.TYPE_II]) +@pytest.mark.parametrize( + "stabilizer_matrix", + [ + pytest.param( + [ + [0, 0, 1, 1], # Z0 Z1. + [1, 1, 0, 0], # X0 X1: the same order occurs twice. + ], + id="same-order-twice", + ), + pytest.param( + [ + [0, 0, 1, 1, 1, 1, 0, 0], # Z0 Z1 X2 X3. + [1, 1, 0, 0, 0, 0, 1, 1], # X0 X1 Z2 Z3: two twists in each direction. + ], + id="even-twisted-pairs", + ), + ], +) +def test_build_graph_state_omits_ancilla_edge_without_odd_twisted_order_pairs( + y_foliation: YFoliation, + stabilizer_matrix: list[list[int]], +) -> None: + result = build_graph_state(StabilizerCode(_matrix(stabilizer_matrix)), y_foliation=y_foliation) + + assert not result.graph.has_edge(result.ancilla_nodes[0], result.ancilla_nodes[1]) + + def test_build_graph_state_returns_index_to_node_maps() -> None: code = StabilizerCode(_matrix([[1, 0, 0, 1]])) diff --git a/tests/test_simulator.py b/tests/test_simulator.py index 24313636..a67b6bf1 100644 --- a/tests/test_simulator.py +++ b/tests/test_simulator.py @@ -2,9 +2,12 @@ from __future__ import annotations +from typing import Any + import numpy as np +import pytest -from graphqomb.command import M +from graphqomb.command import TICK, E, M, N from graphqomb.common import Axis, AxisMeasBasis, Sign from graphqomb.graphstate import GraphState from graphqomb.pattern import Pattern @@ -13,10 +16,15 @@ from graphqomb.statevec import StateVector -def _single_output_pattern(*, measured: bool, axis: Axis = Axis.Z) -> tuple[Pattern, int]: +def _single_output_pattern( + *, + measured: bool, + axis: Axis = Axis.Z, + init_axis: Axis = Axis.X, +) -> tuple[Pattern, int]: graph = GraphState() node = graph.add_node() - graph.register_input(node, 0) + graph.register_input(node, 0, init_axis=init_axis) graph.register_output(node, 0) pauli_frame = PauliFrame(graph, xflow={}, zflow={}) @@ -26,10 +34,145 @@ def _single_output_pattern(*, measured: bool, axis: Axis = Axis.Z) -> tuple[Patt output_node_indices=graph.output_node_indices, commands=commands, pauli_frame=pauli_frame, + input_initialization_axes=graph.input_initialization_axes, ) return pattern, node +def _deterministic_non_output_measurement_pattern() -> tuple[Pattern, int]: + """Create a pattern whose Z-initialized input has deterministic Z outcome.""" + graph = GraphState() + input_node = graph.add_node() + output_node = graph.add_node() + graph.add_edge(input_node, output_node) + graph.register_input(input_node, 0, init_axis=Axis.Z) + graph.register_output(output_node, 0) + graph.assign_meas_basis(input_node, AxisMeasBasis(Axis.Z, Sign.PLUS)) + + pattern = Pattern( + input_node_indices=graph.input_node_indices, + output_node_indices=graph.output_node_indices, + commands=( + N(output_node), + E((input_node, output_node)), + M(input_node, graph.meas_bases[input_node]), + TICK(), + ), + pauli_frame=PauliFrame(graph, xflow={}, zflow={input_node: {output_node}}), + input_initialization_axes=graph.input_initialization_axes, + ) + return pattern, input_node + + +def test_pattern_simulator_rejects_unsupported_command() -> None: + """Unsupported commands fail explicitly instead of recursing.""" + pattern, _ = _single_output_pattern(measured=False) + simulator = PatternSimulator(pattern, SimulatorBackend.StateVector) + unsupported_command: Any = None + + with pytest.raises(TypeError, match="Unsupported command for pattern simulation: NoneType"): + simulator.apply_cmd(unsupported_command, rng=np.random.default_rng(0)) + + +def test_pattern_simulator_initializes_input_in_x_basis() -> None: + """PatternSimulator initializes X-axis inputs as |+>.""" + pattern, _ = _single_output_pattern(measured=False, init_axis=Axis.X) + simulator = PatternSimulator(pattern, SimulatorBackend.StateVector) + + simulator.simulate() + + np.testing.assert_allclose(simulator.state.state(), np.asarray([1.0, 1.0]) / np.sqrt(2)) + + +def test_pattern_simulator_initializes_input_in_y_basis() -> None: + """PatternSimulator initializes Y-axis inputs as |Y+>.""" + pattern, _ = _single_output_pattern(measured=False, init_axis=Axis.Y) + simulator = PatternSimulator(pattern, SimulatorBackend.StateVector) + + simulator.simulate() + + np.testing.assert_allclose(simulator.state.state(), np.asarray([1.0, 1.0j]) / np.sqrt(2)) + + +def test_pattern_simulator_initializes_input_in_z_basis() -> None: + """PatternSimulator initializes Z-axis inputs as |0>.""" + pattern, _ = _single_output_pattern(measured=False, init_axis=Axis.Z) + simulator = PatternSimulator(pattern, SimulatorBackend.StateVector) + + simulator.simulate() + + np.testing.assert_allclose(simulator.state.state(), np.asarray([1.0, 0.0])) + + +def test_pattern_simulator_reorders_mixed_input_axes_by_logical_qindex() -> None: + """Mixed input states are returned in logical output-qubit order.""" + graph = GraphState() + y_input = graph.add_node() + z_input = graph.add_node() + graph.register_input(y_input, 1, init_axis=Axis.Y) + graph.register_input(z_input, 0, init_axis=Axis.Z) + graph.register_output(y_input, 1) + graph.register_output(z_input, 0) + pattern = Pattern( + input_node_indices=graph.input_node_indices, + output_node_indices=graph.output_node_indices, + commands=(), + pauli_frame=PauliFrame(graph, xflow={}, zflow={}), + input_initialization_axes=graph.input_initialization_axes, + ) + simulator = PatternSimulator(pattern, SimulatorBackend.StateVector) + + simulator.simulate() + + expected = np.asarray([1.0, 1.0j, 0.0, 0.0], dtype=np.complex128).reshape(2, 2) / np.sqrt(2) + assert simulator.state.state().shape == (2, 2) + assert simulator.state.state().dtype == np.complex128 + np.testing.assert_allclose(simulator.state.state(), expected) + + +def test_pattern_simulator_samples_non_output_from_exact_probability_by_default() -> None: + """Non-output measurements use the current state instead of a 50/50 assumption.""" + pattern, input_node = _deterministic_non_output_measurement_pattern() + simulator = PatternSimulator(pattern, SimulatorBackend.StateVector) + + simulator.simulate(rng=np.random.default_rng(2)) + + assert simulator.results == {input_node: False} + np.testing.assert_allclose(simulator.state.state(), np.asarray([1.0, 1.0]) / np.sqrt(2)) + + +def test_pattern_simulator_samples_y_initialized_non_output_exactly() -> None: + """A Y-initialized input measured in Y has a deterministic positive outcome.""" + graph = GraphState() + input_node = graph.add_node() + graph.register_input(input_node, 0, init_axis=Axis.Y) + graph.assign_meas_basis(input_node, AxisMeasBasis(Axis.Y, Sign.PLUS)) + pattern = Pattern( + input_node_indices=graph.input_node_indices, + output_node_indices={}, + commands=(M(input_node, graph.meas_bases[input_node]),), + pauli_frame=PauliFrame(graph, xflow={}, zflow={}), + input_initialization_axes=graph.input_initialization_axes, + ) + simulator = PatternSimulator(pattern, SimulatorBackend.StateVector) + + simulator.simulate(rng=np.random.default_rng(2)) + + assert simulator.results == {input_node: False} + np.testing.assert_allclose(simulator.state.state(), np.asarray(1.0)) + + +def test_pattern_simulator_can_use_legacy_uniform_non_output_sampling() -> None: + """calc_prob=False preserves the legacy 50/50 non-output sampling behavior.""" + pattern, input_node = _deterministic_non_output_measurement_pattern() + simulator = PatternSimulator(pattern, SimulatorBackend.StateVector, calc_prob=False) + + simulator.simulate(rng=np.random.default_rng(2)) + + assert simulator.results == {input_node: True} + np.testing.assert_allclose(simulator.state.state(), np.asarray([1.0, -1.0]) / np.sqrt(2)) + + def test_pattern_simulator_applies_output_x_frame_to_statevector() -> None: """An unmeasured output statevector should include pending X frame corrections.""" pattern, node = _single_output_pattern(measured=False) diff --git a/tests/test_statevec.py b/tests/test_statevec.py index 3e7e022c..99975091 100644 --- a/tests/test_statevec.py +++ b/tests/test_statevec.py @@ -5,7 +5,7 @@ import numpy as np import pytest -from graphqomb.common import Plane, PlannerMeasBasis +from graphqomb.common import Axis, AxisMeasBasis, Plane, PlannerMeasBasis, Sign from graphqomb.statevec import StateVector if TYPE_CHECKING: @@ -122,6 +122,46 @@ def test_measure(state_vector: StateVector) -> None: assert np.allclose(state_vector.state().flatten(), expected_state) +@pytest.mark.parametrize( + ("initial_state", "expected_result", "expected_projection_count"), + [ + ([1.0, 0.0], False, 1), + ([0.0, 1.0], True, 2), + ], +) +def test_sample_measure_reuses_sampled_projection( + monkeypatch: pytest.MonkeyPatch, + initial_state: list[float], + expected_result: bool, + expected_projection_count: int, +) -> None: + """Reuse the sampled projection and collapse onto the sampled basis state.""" + state_vector = StateVector(initial_state) + original_tensordot = np.tensordot + projection_count = 0 + + def counting_tensordot( + a: NDArray[np.complex128], + b: NDArray[np.complex128], + axes: tuple[int, int], + ) -> NDArray[np.complex128]: + nonlocal projection_count + projection_count += 1 + return np.asarray(original_tensordot(a, b, axes=axes), dtype=np.complex128) + + monkeypatch.setattr(np, "tensordot", counting_tensordot) + + result = state_vector.sample_measure( + 0, + AxisMeasBasis(Axis.Z, Sign.PLUS), + np.random.default_rng(0), + ) + + assert result is expected_result + assert projection_count == expected_projection_count + np.testing.assert_allclose(state_vector.state(), np.asarray(1.0)) + + def test_tensor_product(state_vector: StateVector) -> None: expected_state = np.asarray([i // 2 for i in range(2 ** (state_vector.num_qubits + 1))]) / np.sqrt(2) other_vector = StateVector.from_num_qubits(1) @@ -131,6 +171,19 @@ def test_tensor_product(state_vector: StateVector) -> None: assert np.allclose(result.state().flatten(), expected_state) +def test_from_product_states_preserves_external_qubit_order() -> None: + """Distinct product-state factors retain their external qubit order.""" + y_plus = np.asarray([1.0, 1.0j], dtype=np.complex128) / np.sqrt(2) + z_plus = np.asarray([1.0, 0.0], dtype=np.complex128) + + result = StateVector.from_product_states((y_plus, z_plus)) + + expected = np.asarray([1.0, 0.0, 1.0j, 0.0], dtype=np.complex128) / np.sqrt(2) + assert result.state().shape == (2, 2) + assert result.state().dtype == np.complex128 + np.testing.assert_allclose(result.state().reshape(-1), expected) + + def test_normalize(state_vector: StateVector) -> None: state_vector.normalize() expected_norm = 1.0 diff --git a/tests/test_stim_compiler.py b/tests/test_stim_compiler.py index 116c76e3..40ddfc33 100644 --- a/tests/test_stim_compiler.py +++ b/tests/test_stim_compiler.py @@ -136,6 +136,52 @@ def test_stim_compile_basic_pattern() -> None: assert stim_str.count("\n") > 0 +@pytest.mark.parametrize( + ("init_axis", "expected_reset"), + [ + (Axis.X, "RX"), + (Axis.Y, "RY"), + (Axis.Z, "R"), + ], +) +def test_stim_compile_uses_input_initialization_axis(init_axis: Axis, expected_reset: str) -> None: + """Input initialization axes choose the corresponding Stim reset instruction.""" + graph = GraphState() + in_node = graph.add_node() + out_node = graph.add_node() + + graph.register_input(in_node, 0, init_axis=init_axis) + graph.register_output(out_node, 0) + graph.add_edge(in_node, out_node) + graph.assign_meas_basis(in_node, AxisMeasBasis(Axis.X, Sign.PLUS)) + + pattern = qompile(graph, {in_node: {out_node}}) + stim_lines = stim_compile(pattern).splitlines() + + assert f"{expected_reset} {in_node}" in stim_lines + + +def test_stim_compile_keeps_non_input_preparations_in_x_basis() -> None: + """Only input reset instructions use the input initialization axis.""" + graph = GraphState() + in_node = graph.add_node() + mid_node = graph.add_node() + out_node = graph.add_node() + + graph.register_input(in_node, 0, init_axis=Axis.Z) + graph.register_output(out_node, 0) + graph.add_edge(in_node, mid_node) + graph.add_edge(mid_node, out_node) + graph.assign_meas_basis(in_node, AxisMeasBasis(Axis.X, Sign.PLUS)) + graph.assign_meas_basis(mid_node, AxisMeasBasis(Axis.X, Sign.PLUS)) + + pattern = qompile(graph, {in_node: {mid_node}, mid_node: {out_node}}) + stim_lines = stim_compile(pattern).splitlines() + + assert f"R {in_node}" in stim_lines + assert f"RX {mid_node}" in stim_lines + + def test_stim_compile_x_measurement() -> None: """Test X measurement compilation.""" pattern, meas_node, in_node = create_simple_pattern_x_measurement() diff --git a/tests/test_stim_importer.py b/tests/test_stim_importer.py index 1f56758a..4b592dea 100644 --- a/tests/test_stim_importer.py +++ b/tests/test_stim_importer.py @@ -64,7 +64,44 @@ def test_stim_text_to_pattern_preserves_sparse_qubit_coordinates() -> None: assert result.stim_to_qubit == {10: 0, 99: 1} assert result.pattern.input_coordinates - assert set(result.pattern.input_coordinates.values()) == {(1.0, 2.0), (3.0, 4.0)} + assert set(result.pattern.input_coordinates.values()) == {(1.0, 2.0, 1.0), (3.0, 4.0, 0.0)} + + +def test_stim_text_to_pattern_aligns_parallel_gate_outputs_with_different_depths() -> None: + result = stim_text_to_pattern( + """ + QUBIT_COORDS(0, 0) 0 + QUBIT_COORDS(1, 0) 1 + H 0 + S 1 + """ + ) + graph = result.pattern.pauli_frame.graphstate + output_coordinates = {qubit: graph.coordinates[node] for node, qubit in graph.output_node_indices.items()} + lane_0_z = sorted(coord[2] for coord in graph.coordinates.values() if np.isclose(coord[0], 0.0)) + lane_1_z = sorted(coord[2] for coord in graph.coordinates.values() if np.isclose(coord[0], 1.0)) + + assert output_coordinates == {0: (0.0, 0.0, 2.0), 1: (1.0, 0.0, 2.0)} + assert lane_0_z == [0.0, 2.0] + assert lane_1_z == [0.0, 1.0, 2.0] + + +def test_stim_text_to_pattern_relocates_idle_input_without_adding_a_wire_node() -> None: + result = stim_text_to_pattern( + """ + QUBIT_COORDS(0, 0) 0 + QUBIT_COORDS(1, 0) 1 + H 0 + """ + ) + graph = result.pattern.pauli_frame.graphstate + idle_input = next(node for node, qubit in graph.input_node_indices.items() if qubit == 1) + idle_output = next(node for node, qubit in graph.output_node_indices.items() if qubit == 1) + + assert graph.number_of_nodes() == 3 + assert idle_input == idle_output + assert graph.coordinates[idle_input] == (1.0, 0.0, 1.0) + assert graph.neighbors(idle_input) == set() @pytest.mark.parametrize( @@ -154,33 +191,105 @@ def test_stim_text_to_pattern_rejects_anticommuting_mpp_in_one_tick_block() -> N stim_text_to_pattern("MPP X0\nMPP Z0") -@pytest.mark.parametrize("y_foliation", [YFoliation.TYPE_I, YFoliation.TYPE_II]) -def test_stim_text_to_pattern_serializes_cyclic_commuting_mpp_flow(y_foliation: YFoliation) -> None: +@pytest.mark.parametrize( + ("y_foliation", "expected_node_count"), + [(YFoliation.TYPE_I, 27), (YFoliation.TYPE_II, 28)], +) +def test_stim_text_to_pattern_builds_commuting_mpp_block_at_common_z( + y_foliation: YFoliation, + expected_node_count: int, +) -> None: result = stim_text_to_pattern( """ + QUBIT_COORDS(0, 0) 0 + QUBIT_COORDS(1, 0) 1 + QUBIT_COORDS(2, 0) 2 + QUBIT_COORDS(3, 0) 3 + QUBIT_COORDS(4, 0) 4 + QUBIT_COORDS(5, 0) 5 + QUBIT_COORDS(6, 0) 6 MPP X0*X1*X4*X5 + DETECTOR rec[-1] MPP Z0*Z1*Z2*Z3 + DETECTOR rec[-1] MPP Y0*X2*Z4*Z6 + DETECTOR rec[-1] MPP Z4*Z5 + DETECTOR rec[-1] MPP X1*X3 + DETECTOR rec[-1] MPP Z2*X6 + DETECTOR rec[-1] """, y_foliation=y_foliation, ) + graph = result.pattern.pauli_frame.graphstate + z_coordinates = {coordinate[2] for coordinate in graph.coordinates.values()} assert len(result.mpp_extractions) == 1 assert len(result.mpp_extractions[0].supports) == 6 + assert graph.number_of_nodes() == expected_node_count + assert np.isclose(min(z_coordinates), 0.0) + assert np.isclose(max(z_coordinates), 2.0) assert set(result.pattern.input_node_indices.values()) == set(range(7)) assert set(result.pattern.output_node_indices.values()) == set(range(7)) + mixed_check_ancilla = next(iter(result.pattern.pauli_frame.parity_check_group[2])) + mixed_loop_ancilla = next(iter(result.pattern.pauli_frame.parity_check_group[5])) + assert graph.has_edge(mixed_check_ancilla, mixed_loop_ancilla) + + +def test_stim_text_to_pattern_advances_z_once_per_mpp_tick_block() -> None: + result = stim_text_to_pattern( + """ + QUBIT_COORDS(0, 0) 0 + QUBIT_COORDS(1, 0) 1 + MPP X0 + MPP X1 + TICK + MPP X0*X1 + """ + ) + graph = result.pattern.pauli_frame.graphstate + z_coordinates = {coordinate[2] for coordinate in graph.coordinates.values()} + + assert len(result.mpp_extractions) == 2 + assert np.isclose(max(z_coordinates), 4.0) -def test_stim_text_to_pattern_uses_automatic_zflow_for_mpp_graph() -> None: - result = stim_text_to_pattern("MPP X0*Z1") +def test_stim_text_to_pattern_relocates_idle_input_to_mpp_output_layer() -> None: + result = stim_text_to_pattern( + """ + QUBIT_COORDS(0, 0) 0 + QUBIT_COORDS(1, 0) 1 + MPP X0 + """ + ) graph = result.pattern.pauli_frame.graphstate + idle_input = next(node for node, qubit in graph.input_node_indices.items() if qubit == 1) + idle_output = next(node for node, qubit in graph.output_node_indices.items() if qubit == 1) + active_output = next(node for node, qubit in graph.output_node_indices.items() if qubit == 0) + + assert idle_input == idle_output + assert graph.coordinates[idle_input] == (1.0, 0.0, 2.0) + assert graph.coordinates[active_output] == (0.0, 0.0, 2.0) + assert graph.neighbors(idle_input) == set() + + +def test_stim_text_to_pattern_derives_complete_zflow_from_xflow() -> None: + result = stim_text_to_pattern("H 0\nTICK\nMPP X0") + frame = result.pattern.pauli_frame + + assert frame.zflow == {node: odd_neighbors(targets, frame.graphstate) for node, targets in frame.xflow.items()} + - assert result.pattern.pauli_frame.xflow - for node, correction_nodes in result.pattern.pauli_frame.xflow.items(): - assert result.pattern.pauli_frame.zflow[node] == odd_neighbors(correction_nodes, graph) +def test_stim_text_to_pattern_excludes_mpp_ancilla_from_xflow() -> None: + result = stim_text_to_pattern("MPP X0\nDETECTOR rec[-1]") + frame = result.pattern.pauli_frame + + ancilla_nodes = frame.parity_check_group[0] + assert len(ancilla_nodes) == 1 + assert set(frame.xflow) == set(frame.graphstate.meas_bases) - ancilla_nodes + assert all(targets.isdisjoint(ancilla_nodes) for targets in frame.xflow.values()) def test_stim_text_to_pattern_appends_output_after_type_i_mpp_measurements() -> None: @@ -251,7 +360,7 @@ def test_stim_text_to_pattern_imports_single_qubit_pauli_measurements( def test_stim_text_to_pattern_assigns_single_measurement_to_existing_wire_node() -> None: - result = stim_text_to_pattern("H 10\nTICK\nMX 10") + result = stim_text_to_pattern("QUBIT_COORDS(1, 2) 10\nH 10\nTICK\nMX 10") output_node = next(node for node, qubit in result.pattern.output_node_indices.items() if qubit == 0) output_measurements = [ command for command in result.pattern.commands if isinstance(command, M) and command.node == output_node @@ -262,6 +371,67 @@ def test_stim_text_to_pattern_assigns_single_measurement_to_existing_wire_node() assert isinstance(output_measurements[0].meas_basis, AxisMeasBasis) assert output_measurements[0].meas_basis.axis == Axis.X assert output_measurements[0].meas_basis.sign == Sign.PLUS + assert result.pattern.pauli_frame.graphstate.coordinates[output_node] == (1.0, 2.0, 1.0) + + +def test_stim_text_to_pattern_preserves_mpp_lane_coordinate_for_terminal_measurement() -> None: + result = stim_text_to_pattern( + """ + QUBIT_COORDS(1, 2) 0 + MPP X0 + TICK + MX 0 + """ + ) + graph = result.pattern.pauli_frame.graphstate + output_node = next(node for node, qubit in result.pattern.output_node_indices.items() if qubit == 0) + output_basis = graph.meas_bases[output_node] + + assert graph.number_of_nodes() == 4 + assert graph.coordinates[output_node] == (1.0, 2.0, 2.0) + assert isinstance(output_basis, AxisMeasBasis) + assert output_basis.axis == Axis.X + + +def test_stim_text_to_pattern_places_gate_after_mpp_at_next_z_layer() -> None: + result = stim_text_to_pattern( + """ + QUBIT_COORDS(1, 2) 0 + MPP X0 + TICK + H 0 + TICK + MX 0 + """ + ) + graph = result.pattern.pauli_frame.graphstate + output_node = next(node for node, qubit in graph.output_node_indices.items() if qubit == 0) + + assert graph.coordinates[output_node] == (1.0, 2.0, 3.0) + assert all(len(coord) == 3 for coord in graph.coordinates.values()) + assert {coord[2] for coord in graph.coordinates.values()} == {0.0, 1.0, 2.0, 3.0} + + +def test_stim_text_to_pattern_composes_gate_output_with_mpp_input_at_same_z() -> None: + result = stim_text_to_pattern( + """ + QUBIT_COORDS(0, 0) 0 + QUBIT_COORDS(1, 0) 1 + H 0 + TICK + MPP X0*Z1 + """ + ) + graph = result.pattern.pauli_frame.graphstate + input_coordinates = {qubit: graph.coordinates[node] for node, qubit in graph.input_node_indices.items()} + output_coordinates = {qubit: graph.coordinates[node] for node, qubit in graph.output_node_indices.items()} + + assert input_coordinates == {0: (0.0, 0.0, 0.0), 1: (1.0, 0.0, 1.0)} + assert output_coordinates == {0: (0.0, 0.0, 3.0), 1: (1.0, 0.0, 3.0)} + assert {(coord[0], coord[2]) for coord in graph.coordinates.values()} >= { + (0.0, 1.0), + (1.0, 1.0), + } @pytest.mark.parametrize( @@ -410,6 +580,48 @@ def test_stim_text_to_pattern_preserves_deterministic_mpp_detectors(text: str) - compiled.detector_error_model() +@pytest.mark.parametrize("y_foliation", [YFoliation.TYPE_I, YFoliation.TYPE_II]) +def test_stim_text_to_pattern_preserves_detectors_for_twisted_stabilizer_orders( + y_foliation: YFoliation, +) -> None: + pattern = stim_text_to_pattern( + """ + MPP X0*Z1 + MPP Z0*X1 + TICK + MPP X0*Z1 + DETECTOR rec[-1] rec[-3] + MPP Z0*X1 + DETECTOR rec[-1] rec[-3] + """, + y_foliation=y_foliation, + ).pattern + compiled = stim.Circuit(stim_compile(pattern, emit_qubit_coords=False)) + + assert compiled.detector_error_model().num_detectors == 2 + + +@pytest.mark.parametrize( + ("text", "expected_ancilla_axis"), + [ + ("RY 0\nTICK\nMPP Y0\nDETECTOR rec[-1]", Axis.Y), + ("RY 0 1\nTICK\nMPP Y0*Y1\nDETECTOR rec[-1]", Axis.X), + ], +) +def test_stim_text_to_pattern_preserves_deterministic_type_i_y_mpp_detector( + text: str, + expected_ancilla_axis: Axis, +) -> None: + pattern = stim_text_to_pattern(text, y_foliation=YFoliation.TYPE_I).pattern + ancilla_node = next(iter(pattern.pauli_frame.parity_check_group[0])) + ancilla_basis = pattern.pauli_frame.graphstate.meas_bases[ancilla_node] + compiled = stim.Circuit(stim_compile(pattern, emit_qubit_coords=False)) + + assert isinstance(ancilla_basis, AxisMeasBasis) + assert ancilla_basis.axis == expected_ancilla_axis + compiled.detector_error_model() + + def test_stim_text_to_pattern_composes_mpp_output_into_next_mpp_input() -> None: result = stim_text_to_pattern("MPP X0\nTICK\nMPP X0") graph = result.pattern.pauli_frame.graphstate @@ -438,12 +650,66 @@ def test_stim_text_to_pattern_rejects_mixed_measurement_and_unitary_block(measur stim_text_to_pattern(f"H 0\n{measurement}\n") -@pytest.mark.parametrize("instruction", ["R 0", "RX 0", "RY 0", "MR 0", "MRX 0", "MRY 0"]) -def test_stim_text_to_pattern_defers_reset_instructions(instruction: str) -> None: +@pytest.mark.parametrize( + ("instruction", "expected_axis", "compiled_instruction"), + [ + ("R 0", Axis.Z, "R 0"), + ("RZ 0", Axis.Z, "R 0"), + ("RX 0", Axis.X, "RX 0"), + ("RY 0", Axis.Y, "RY 0"), + ], +) +def test_stim_text_to_pattern_imports_initial_reset( + instruction: str, + expected_axis: Axis, + compiled_instruction: str, +) -> None: + result = stim_text_to_pattern(instruction) + input_node = next(node for node, q_index in result.pattern.input_node_indices.items() if q_index == 0) + + assert result.pattern.input_initialization_axes[input_node] == expected_axis + assert compiled_instruction in stim_compile(result.pattern, emit_qubit_coords=False).splitlines() + + +def test_stim_text_to_pattern_uses_last_leading_reset() -> None: + result = stim_text_to_pattern("R 0\nRY 0\nH 0") + input_node = next(node for node, q_index in result.pattern.input_node_indices.items() if q_index == 0) + + assert result.pattern.input_initialization_axes[input_node] == Axis.Y + + +def test_stim_text_to_pattern_allows_initial_reset_after_other_qubit_operation() -> None: + result = stim_text_to_pattern("H 0\nR 1") + input_node = next(node for node, q_index in result.pattern.input_node_indices.items() if q_index == 1) + + assert result.pattern.input_initialization_axes[input_node] == Axis.Z + + +@pytest.mark.parametrize("instruction", ["R", "RX", "RY"]) +def test_stim_text_to_pattern_rejects_reset_after_quantum_operation(instruction: str) -> None: + with pytest.raises(ValueError, match="only initial resets are supported"): + stim_text_to_pattern(f"H 0\n{instruction} 0") + + +@pytest.mark.parametrize("instruction", ["MR 0", "MRX 0", "MRY 0"]) +def test_stim_text_to_pattern_defers_measurement_reset_instructions(instruction: str) -> None: with pytest.raises(ValueError, match="Unsupported Stim instruction"): stim_text_to_pattern(instruction) +def test_stim_text_to_pattern_rejects_duplicate_qubit_coordinates() -> None: + with pytest.raises(ValueError, match="distinct XY projections"): + stim_text_to_pattern("QUBIT_COORDS(0, 0) 0\nQUBIT_COORDS(0, 0) 1\nCZ 0 1") + + +def test_stim_text_to_pattern_rejects_coordinates_sharing_an_xy_projection() -> None: + with pytest.raises(ValueError, match="distinct XY projections"): + stim_text_to_pattern( + "QUBIT_COORDS(0, 0, 1) 0\nQUBIT_COORDS(0, 0, 2) 1\nCZ 0 1", + coord_dims=3, + ) + + def test_stim_text_to_pattern_rejects_qubit_reuse_after_single_measurement() -> None: with pytest.raises(ValueError, match="terminate those qubit lifetimes"): stim_text_to_pattern("M 0\nTICK\nMPP X0")