diff --git a/cirq-core/cirq/contrib/paulistring/__init__.py b/cirq-core/cirq/contrib/paulistring/__init__.py index 678f6b7ca7d..282744d2457 100644 --- a/cirq-core/cirq/contrib/paulistring/__init__.py +++ b/cirq-core/cirq/contrib/paulistring/__init__.py @@ -47,4 +47,5 @@ measure_pauli_strings as measure_pauli_strings, CircuitToPauliStringsParameters as CircuitToPauliStringsParameters, CircuitToPauliStringsMeasurementResult as CircuitToPauliStringsMeasurementResult, + generate_trex_and_readout_circuits as generate_trex_and_readout_circuits, ) diff --git a/cirq-core/cirq/contrib/paulistring/pauli_string_measurement_with_readout_mitigation.py b/cirq-core/cirq/contrib/paulistring/pauli_string_measurement_with_readout_mitigation.py index 71d08c77ec3..9c144099dbb 100644 --- a/cirq-core/cirq/contrib/paulistring/pauli_string_measurement_with_readout_mitigation.py +++ b/cirq-core/cirq/contrib/paulistring/pauli_string_measurement_with_readout_mitigation.py @@ -26,7 +26,7 @@ import sympy import cirq.contrib.shuffle_circuits.shuffle_circuits_with_readout_benchmarking as sc_readout -from cirq import circuits, ops, study, work +from cirq import circuits, ops, study, transformers, work from cirq.experiments.readout_confusion_matrix import TensoredConfusionMatrices if TYPE_CHECKING: @@ -1155,7 +1155,8 @@ def generate_trex_and_readout_circuits( num_twirls: int, num_readout_circuits: int, rng: np.random.Generator, -) -> tuple[list[circuits.Circuit], TRexMetadata]: + insert_strategy: circuits.InsertStrategy = circuits.InsertStrategy.INLINE, +) -> list[tuple[list[circuits.Circuit], TRexMetadata]]: """Generates a list of circuits for TREX benchmarking and readout calibration. This function generates `num_twirls` circuits by applying random Pauli twirls @@ -1165,17 +1166,70 @@ def generate_trex_and_readout_circuits( Args: circuit_to_pauli: A CircuitToPauliStringsParameters object containing the original - circuit and its associated Pauli strings. + circuit and the Pauli strings that the user wishes to measure on the output + of the circuit. num_twirls: The number of twirled circuits to generate for each original circuit. num_readout_circuits: The number of readout calibration circuits to generate. rng: A NumPy random number generator for generating random Pauli twirls. + insert_strategy: The strategy for inserting measurement operations into the circuit. + Defaults to circuits.InsertStrategy.INLINE. Returns: - A tuple containing: - - A combined list of the twirled Pauli circuits followed by the readout circuits. + A list of tuples, one for each Pauli group in `circuit_to_pauli.pauli_strings`. + Each tuple contains: + - A list of the generated circuits (twirled circuits followed by readout circuits). - A TRexMetadata object containing the random choices needed for post-processing. """ - raise NotImplementedError("T-REX error mitigation is not yet implemented.") + results: list[tuple[list[circuits.Circuit], TRexMetadata]] = [] + + circuit = transformers.drop_terminal_measurements(circuit_to_pauli.circuit.unfreeze()) + + for pauli_group in circuit_to_pauli.pauli_strings: + + qubit_pauli_dict_unsorted: dict[ops.Qid, ops.Pauli] = {} + for pauli_str in pauli_group: + qubit_pauli_dict_unsorted.update(pauli_str) + + qubit_pauli_dict = { + q: qubit_pauli_dict_unsorted[q] for q in sorted(qubit_pauli_dict_unsorted) + } + + joint_basis_pauli: ops.PauliString = ops.PauliString(qubit_pauli_dict) + num_qubits = len(qubit_pauli_dict) + + twirl_choices = _generate_random_boolean_choices(num_twirls, num_qubits, rng) + pauli_circuits = _build_trex_twirled_pauli_circuits( + circuit, joint_basis_pauli, twirl_choices, insert_strategy + ) + + readout_choices = _generate_random_boolean_choices(num_readout_circuits, num_qubits, rng) + + readout_circuits = [ + circuits.Circuit.from_moments( + # Note: An empty moment may be inserted here if readout_choices_i contains + # all false values. This is intentional to ensure the depth and measurement + # timing of all readout circuits remain consistent and uniform. + circuits.Moment( + ops.X.on_each( + q for q, flip in zip(joint_basis_pauli, readout_choices_i) if flip + ) + ), + circuits.Moment(ops.M(*joint_basis_pauli.qubits, key='result')), + ) + for readout_choices_i in readout_choices + ] + + # Bundle the circuits and their corresponding metadata together + group_circuits = pauli_circuits + readout_circuits + group_metadata = TRexMetadata( + pauli_str=joint_basis_pauli, + twirl_choices=twirl_choices, + readout_choices=readout_choices, + ) + + results.append((group_circuits, group_metadata)) + + return results def measure_pauli_strings( diff --git a/cirq-core/cirq/contrib/paulistring/pauli_string_measurement_with_readout_mitigation_test.py b/cirq-core/cirq/contrib/paulistring/pauli_string_measurement_with_readout_mitigation_test.py index 72127f3902b..b509803871b 100644 --- a/cirq-core/cirq/contrib/paulistring/pauli_string_measurement_with_readout_mitigation_test.py +++ b/cirq-core/cirq/contrib/paulistring/pauli_string_measurement_with_readout_mitigation_test.py @@ -24,9 +24,8 @@ import cirq from cirq.contrib.paulistring import CircuitToPauliStringsParameters, measure_pauli_strings from cirq.contrib.paulistring.pauli_string_measurement_with_readout_mitigation import ( - _build_trex_twirled_pauli_circuits, + generate_trex_and_readout_circuits, PostFilteringSymmetryCalibrationResult as PostFilteringSymmetryCalibrationResult, - TRexMetadata, ) from cirq.experiments import SingleQubitReadoutCalibrationResult from cirq.experiments.single_qubit_readout_calibration_test import NoisySingleQubitReadoutSampler @@ -939,71 +938,110 @@ def test_sampler_receives_correct_circuits(use_sweep: bool) -> None: assert measured == expected_qubits -def test_build_trex_twirled_pauli_circuits_multiple_twirls(): - """Test generating multiple circuits from a multi-row twirl_choices array.""" - q0, q1, q2 = cirq.LineQubit.range(3) - base_circuit = cirq.Circuit(cirq.H(q0), cirq.CNOT(q0, q1), cirq.CNOT(q0, q2)) - basis_ps = cirq.X(q0) * cirq.Y(q1) * cirq.Z(q2) +@pytest.mark.parametrize( + "insert_strategy, expected_ops_in_moment_3", + [(cirq.InsertStrategy.INLINE, 1), (cirq.InsertStrategy.EARLIEST, 2)], +) +def test_generate_trex_and_readout_circuits( + insert_strategy: cirq.InsertStrategy, expected_ops_in_moment_3: int +) -> None: + """Test the generation of TRex twirled circuits and readout calibration circuits.""" + q0, q1 = cirq.LineQubit.range(2) + base_circuit = cirq.FrozenCircuit(cirq.Circuit(cirq.H(q0), cirq.CNOT(q0, q1), cirq.H(q1))) - # 3 different twirl choices - twirl_choices = np.array( - [ - [False, True, True], # q0(no flip), q1(flip), q2(flip)] - [True, False, False], # q0(flip), q1(no flip), q2(no flip)] - [False, False, False], # q0(no flip), q1(no flip), q2(no flip)] - ] + pauli_group_1: list[cirq.PauliString] = [cirq.PauliString(cirq.Z(q0) * cirq.Z(q1))] + # Qubits in group 2 are unsorted. + pauli_group_2: list[cirq.PauliString] = [ + cirq.PauliString(cirq.X(q1)), + cirq.PauliString(cirq.X(q0)), + ] + + pauli_strings = [pauli_group_1, pauli_group_2] + + params = CircuitToPauliStringsParameters(circuit=base_circuit, pauli_strings=pauli_strings) + + num_twirls = 3 + num_readout_circuits = 2 + rng = np.random.default_rng(seed=42) + + results = generate_trex_and_readout_circuits( + circuit_to_pauli=params, + num_twirls=num_twirls, + num_readout_circuits=num_readout_circuits, + rng=rng, + insert_strategy=insert_strategy, ) - circuits = _build_trex_twirled_pauli_circuits(base_circuit, basis_ps, twirl_choices) + # We have 2 Pauli groups, so we should get 2 tuples of (circuits, metadata). + assert len(results) == 2 - assert len(circuits) == 3 + # Verify Group 1 Metadata and Circuits + group1_circuits, meta1 = results[0] - q0_no_flip = cirq.Ry(rads=-np.pi / 2)(q0) - q0_flip = cirq.Ry(rads=np.pi / 2)(q0) + # We should get (3 twirls + 2 readouts) = 5 circuits + assert len(group1_circuits) == 5 - q1_no_flip = cirq.Rx(rads=np.pi / 2)(q1) - q1_flip = cirq.Rx(rads=-np.pi / 2)(q1) + # Verify metadata. + assert meta1.pauli_str == cirq.Z(q0) * cirq.Z(q1) + assert meta1.twirl_choices.shape == (num_twirls, 2) + assert meta1.readout_choices.shape == (num_readout_circuits, 2) - q2_no_flip = cirq.I(q2) - q2_flip = cirq.X(q2) + # Indices 0, 1, 2 belong to the twirled circuits of group 1 + for i in range(3): + # Check how many operations are in the 3rd moment + # EARLIEST packs 2 ops, while INLINE leaves it at 1 op + assert len(group1_circuits[i].moments[2].operations) == expected_ops_in_moment_3 - # Verify Circuit 0: row [False, True, True] - assert q0_no_flip in circuits[0].moments[-2].operations - assert q1_flip in circuits[0].moments[-2].operations - assert q2_flip in circuits[0].moments[-2].operations + meas_op = group1_circuits[i].moments[-1].operations[0] + assert isinstance(meas_op.gate, cirq.MeasurementGate) + assert set(meas_op.qubits) == {q0, q1} + assert meas_op.gate.key == 'result' - # Verify Circuit 1: row [True, False, False] - assert q0_flip in circuits[1].moments[-2].operations - assert q1_no_flip in circuits[1].moments[-2].operations - assert q2_no_flip in circuits[1].moments[-2].operations + # Indices 3, 4 belong to the readout circuits of group 1 + for i in range(3, 5): + readout_circuit = group1_circuits[i] + # Readout circuits have exactly 2 moments: Optional X gates, then Measurement + assert len(readout_circuit.moments) == 2 - # Verify Circuit 2: row [False, False, False] - assert q0_no_flip in circuits[2].moments[-2].operations - assert q1_no_flip in circuits[2].moments[-2].operations - assert q2_no_flip in circuits[2].moments[-2].operations + # Check first moment (State preparation via X gates) + for op in readout_circuit.moments[0].operations: + assert op.gate == cirq.X - # Verify that every generated circuit ends with the correct joint measurement - for circuit in circuits: - meas_op = circuit.moments[-1].operations[0] + # Check second moment (Measurement) + meas_op = readout_circuit.moments[-1].operations[0] assert isinstance(meas_op.gate, cirq.MeasurementGate) - assert meas_op.qubits == (q0, q1, q2) + assert set(meas_op.qubits) == {q0, q1} assert meas_op.gate.key == 'result' + # Verify Group 2 Metadata and Circuits + group2_circuits, meta2 = results[1] -def test_trex_metadata_instantiation() -> None: - """Test the instantiation and attributes of TRexMetadata.""" - q0, q1 = cirq.LineQubit.range(2) - pauli_str = cirq.X(q0) * cirq.Z(q1) + # We should get (3 twirls + 2 readouts) = 5 circuits + assert len(group2_circuits) == 5 - # 2D boolean arrays of shape (num_readout_circuits, num_qubits) - twirl_choices = np.array([[True, False], [False, True], [True, True]]) + # Verify metadata. + # The two X Paulis should have been combined into a joint X0*X1 string + assert meta2.pauli_str == cirq.X(q0) * cirq.X(q1) + assert meta2.twirl_choices.shape == (num_twirls, 2) + assert meta2.readout_choices.shape == (num_readout_circuits, 2) + # The qubits order are sorted. + assert meta2.pauli_str.qubits == (q0, q1) - readout_choices = np.array([[False, False], [True, True], [False, True]]) + # Indices 0, 1, 2 belong to the twirled circuits of group 2 + for i in range(3): + # Check how many operations are in the 3rd moment + # EARLIEST packs 2 ops, while INLINE leaves it at 1 op + assert len(group2_circuits[i].moments[2].operations) == expected_ops_in_moment_3 - metadata = TRexMetadata( - pauli_str=pauli_str, twirl_choices=twirl_choices, readout_choices=readout_choices - ) + meas_op = group2_circuits[i].moments[-1].operations[0] + assert isinstance(meas_op.gate, cirq.MeasurementGate) + assert set(meas_op.qubits) == {q0, q1} + assert meas_op.gate.key == 'result' - assert metadata.pauli_str == pauli_str - np.testing.assert_array_equal(metadata.twirl_choices, twirl_choices) - np.testing.assert_array_equal(metadata.readout_choices, readout_choices) + # Indices 3, 4 belong to the readout circuits of group 2 + for i in range(3, 5): + readout_circuit = group2_circuits[i] + assert len(readout_circuit.moments) == 2 + meas_op = readout_circuit.moments[-1].operations[0] + assert isinstance(meas_op.gate, cirq.MeasurementGate) + assert set(meas_op.qubits) == {q0, q1}