From a8da2cccd52e0675024ba0a9252e472fe140dbe4 Mon Sep 17 00:00:00 2001 From: ddddddanni Date: Wed, 8 Jul 2026 11:38:19 -0700 Subject: [PATCH 01/12] init version --- .../cirq/contrib/paulistring/__init__.py | 1 + ...ing_measurement_with_readout_mitigation.py | 50 +++++++- ...easurement_with_readout_mitigation_test.py | 116 +++++++++--------- 3 files changed, 111 insertions(+), 56 deletions(-) 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..206b93c69c2 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 @@ -1155,6 +1155,7 @@ def generate_trex_and_readout_circuits( num_twirls: int, num_readout_circuits: int, rng: np.random.Generator, + insert_strategy: circuits.InsertStrategy = circuits.InsertStrategy.INLINE, ) -> tuple[list[circuits.Circuit], TRexMetadata]: """Generates a list of circuits for TREX benchmarking and readout calibration. @@ -1169,13 +1170,60 @@ def generate_trex_and_readout_circuits( 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 TRexMetadata object containing the random choices needed for post-processing. """ - raise NotImplementedError("T-REX error mitigation is not yet implemented.") + all_generated_circuits: list[circuits.Circuit] = [] + metadata_list: list[TRexMetadata] = [] + + circuit = cirq.drop_terminal_measurements(circuit_to_pauli.circuit.unfreeze()) + + for pauli_group in circuit_to_pauli.pauli_strings: + + qubit_pauli_dict = {} + for pauli_str in pauli_group: + for qubit, pauli in pauli_str.items(): + qubit_pauli_dict[qubit] = pauli + + joint_basis_pauli = ops.PauliString(qubit_pauli_dict) + num_qubits = len(qubit_pauli_dict) + + twirl_choices = _generate_random_boolean_choices(num_twirls, num_qubits, rng) + # overall_flip = twirl_choices.sum(axis=1) % 2 + pauli_circuits = _build_trex_twirled_pauli_circuits( + circuit, joint_basis_pauli, twirl_choices, insert_strategy + ) + all_generated_circuits.extend(pauli_circuits) + + readout_choices = _generate_random_boolean_choices(num_readout_circuits, num_qubits, rng) + # overall_readout_parity = readout_choices.sum(axis=1) % 2 + + readout_circuits = [ + cirq.Circuit.from_moments( + cirq.Moment( + cirq.X.on_each( + q for q, flip in zip(joint_basis_pauli, readout_choices_i) if flip + ) + ), + cirq.Moment(cirq.M(*joint_basis_pauli.qubits, key='result')), + ) + for readout_choices_i in readout_choices + ] + all_generated_circuits.extend(readout_circuits) + + metadata_list.append( + TRexMetadata( + pauli_str=joint_basis_pauli, + twirl_choices=twirl_choices, + readout_choices=readout_choices, + ) + ) + return all_generated_circuits, metadata_list 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 34595def610..0a797760ee8 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 @@ -939,71 +939,77 @@ 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) - - # 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)] - ] - ) - - circuits = _build_trex_twirled_pauli_circuits(base_circuit, basis_ps, twirl_choices) - - assert len(circuits) == 3 +def test_generate_trex_and_readout_circuits() -> 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))) - q0_no_flip = cirq.Ry(rads=-np.pi / 2)(q0) - q0_flip = cirq.Ry(rads=np.pi / 2)(q0) + pauli_group_1 = [cirq.PauliString(cirq.Z(q0) * cirq.Z(q1))] - q1_no_flip = cirq.Rx(rads=np.pi / 2)(q1) - q1_flip = cirq.Rx(rads=-np.pi / 2)(q1) + pauli_group_2 = [cirq.PauliString(cirq.X(q0)), cirq.PauliString(cirq.X(q1))] - q2_no_flip = cirq.I(q2) - q2_flip = cirq.X(q2) + pauli_strings = [pauli_group_1, pauli_group_2] - # 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 + params = CircuitToPauliStringsParameters(circuit=base_circuit, pauli_strings=pauli_strings) - # 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 + num_twirls = 3 + num_readout_circuits = 2 + rng = np.random.default_rng(seed=42) - # 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 + all_circuits, metadata_list = generate_trex_and_readout_circuits( + circuit_to_pauli=params, + num_twirls=num_twirls, + num_readout_circuits=num_readout_circuits, + rng=rng, + ) - # Verify that every generated circuit ends with the correct joint measurement - for circuit in circuits: - meas_op = circuit.moments[-1].operations[0] + # For each group, we should get `num_twirls` + `num_readout_circuits` circuits. + # Total circuits = (3 twirls + 2 readouts) * 2 groups = 10 circuits + assert len(all_circuits) == 10 + assert len(metadata_list) == 2 + + # Verify Group1 Metadata and Circuits + meta1 = metadata_list[0] + 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) + + # Indices 0, 1, 2 belong to the twirled circuits of group 1 + for i in range(3): + # Twirled circuits should contain the base operations + basis changes + measurement + assert len(all_circuits[i]) >= len(base_circuit) + meas_op = all_circuits[i].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' + # Indices 3, 4 belong to the readout circuits of group 1 + for i in range(3, 5): + readout_circuit = all_circuits[i] + # Readout circuits have exactly 2 moments: Optional X gates, then Measurement + assert len(readout_circuit.moments) == 2 -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) + # Check first moment (State preparation via X gates) + for op in readout_circuit.moments[0].operations: + assert op.gate == cirq.X - # 2D boolean arrays of shape (num_readout_circuits, num_qubits) - twirl_choices = np.array([[True, False], [False, True], [True, True]]) - - readout_choices = np.array([[False, False], [True, True], [False, True]]) - - metadata = TRexMetadata( - pauli_str=pauli_str, twirl_choices=twirl_choices, readout_choices=readout_choices - ) + # Check second moment (Measurement) + meas_op = readout_circuit.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) + # Verify Group 2 Metadata and Circuits + meta2 = metadata_list[1] + # 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) + + # Indices 8, 9 belong to the readout circuits of group 2 + for i in range(8, 10): + readout_circuit = all_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} From ec54b854dafa98737f63653841b1d2531afc17ac Mon Sep 17 00:00:00 2001 From: ddddddanni Date: Wed, 8 Jul 2026 11:59:12 -0700 Subject: [PATCH 02/12] fix test --- ...string_measurement_with_readout_mitigation.py | 16 ++++++++-------- ...g_measurement_with_readout_mitigation_test.py | 11 ++++++----- 2 files changed, 14 insertions(+), 13 deletions(-) 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 206b93c69c2..e21a276f286 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, work, transformers from cirq.experiments.readout_confusion_matrix import TensoredConfusionMatrices if TYPE_CHECKING: @@ -1156,7 +1156,7 @@ def generate_trex_and_readout_circuits( num_readout_circuits: int, rng: np.random.Generator, insert_strategy: circuits.InsertStrategy = circuits.InsertStrategy.INLINE, -) -> tuple[list[circuits.Circuit], TRexMetadata]: +) -> tuple[list[circuits.Circuit], list[TRexMetadata]]: """Generates a list of circuits for TREX benchmarking and readout calibration. This function generates `num_twirls` circuits by applying random Pauli twirls @@ -1181,7 +1181,7 @@ def generate_trex_and_readout_circuits( all_generated_circuits: list[circuits.Circuit] = [] metadata_list: list[TRexMetadata] = [] - circuit = cirq.drop_terminal_measurements(circuit_to_pauli.circuit.unfreeze()) + circuit = transformers.drop_terminal_measurements(circuit_to_pauli.circuit.unfreeze()) for pauli_group in circuit_to_pauli.pauli_strings: @@ -1190,7 +1190,7 @@ def generate_trex_and_readout_circuits( for qubit, pauli in pauli_str.items(): qubit_pauli_dict[qubit] = pauli - joint_basis_pauli = ops.PauliString(qubit_pauli_dict) + joint_basis_pauli: cirq.PauliString = ops.PauliString(qubit_pauli_dict) num_qubits = len(qubit_pauli_dict) twirl_choices = _generate_random_boolean_choices(num_twirls, num_qubits, rng) @@ -1204,13 +1204,13 @@ def generate_trex_and_readout_circuits( # overall_readout_parity = readout_choices.sum(axis=1) % 2 readout_circuits = [ - cirq.Circuit.from_moments( - cirq.Moment( - cirq.X.on_each( + circuits.Circuit.from_moments( + circuits.Moment( + ops.X.on_each( q for q, flip in zip(joint_basis_pauli, readout_choices_i) if flip ) ), - cirq.Moment(cirq.M(*joint_basis_pauli.qubits, key='result')), + circuits.Moment(ops.M(*joint_basis_pauli.qubits, key='result')), ) for readout_choices_i in readout_choices ] 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 0a797760ee8..617a139883d 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, PostFilteringSymmetryCalibrationResult as PostFilteringSymmetryCalibrationResult, - TRexMetadata, + generate_trex_and_readout_circuits, ) from cirq.experiments import SingleQubitReadoutCalibrationResult from cirq.experiments.single_qubit_readout_calibration_test import NoisySingleQubitReadoutSampler @@ -944,9 +943,11 @@ def test_generate_trex_and_readout_circuits() -> None: q0, q1 = cirq.LineQubit.range(2) base_circuit = cirq.FrozenCircuit(cirq.Circuit(cirq.H(q0), cirq.CNOT(q0, q1))) - pauli_group_1 = [cirq.PauliString(cirq.Z(q0) * cirq.Z(q1))] - - pauli_group_2 = [cirq.PauliString(cirq.X(q0)), cirq.PauliString(cirq.X(q1))] + pauli_group_1: list[cirq.PauliString] = [cirq.PauliString(cirq.Z(q0) * cirq.Z(q1))] + pauli_group_2: list[cirq.PauliString] = [ + cirq.PauliString(cirq.X(q0)), + cirq.PauliString(cirq.X(q1)), + ] pauli_strings = [pauli_group_1, pauli_group_2] From 63f17ec78eb39d0402dafc31bb7f94bd2dfa9732 Mon Sep 17 00:00:00 2001 From: ddddddanni Date: Wed, 8 Jul 2026 12:03:10 -0700 Subject: [PATCH 03/12] fix format --- .../pauli_string_measurement_with_readout_mitigation.py | 2 +- .../pauli_string_measurement_with_readout_mitigation_test.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) 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 e21a276f286..d16f4a49c68 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, transformers +from cirq import circuits, ops, study, transformers, work from cirq.experiments.readout_confusion_matrix import TensoredConfusionMatrices if TYPE_CHECKING: 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 617a139883d..6d88f69721a 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,8 +24,8 @@ import cirq from cirq.contrib.paulistring import CircuitToPauliStringsParameters, measure_pauli_strings from cirq.contrib.paulistring.pauli_string_measurement_with_readout_mitigation import ( - PostFilteringSymmetryCalibrationResult as PostFilteringSymmetryCalibrationResult, generate_trex_and_readout_circuits, + PostFilteringSymmetryCalibrationResult as PostFilteringSymmetryCalibrationResult, ) from cirq.experiments import SingleQubitReadoutCalibrationResult from cirq.experiments.single_qubit_readout_calibration_test import NoisySingleQubitReadoutSampler From 1e57ea6731e075e55c7356bf75f7d6ddac9eae73 Mon Sep 17 00:00:00 2001 From: ddddddanni Date: Wed, 8 Jul 2026 12:12:03 -0700 Subject: [PATCH 04/12] fix lint again --- .../pauli_string_measurement_with_readout_mitigation.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) 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 d16f4a49c68..6435377484a 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 @@ -1187,8 +1187,7 @@ def generate_trex_and_readout_circuits( qubit_pauli_dict = {} for pauli_str in pauli_group: - for qubit, pauli in pauli_str.items(): - qubit_pauli_dict[qubit] = pauli + qubit_pauli_dict.update(pauli_str) joint_basis_pauli: cirq.PauliString = ops.PauliString(qubit_pauli_dict) num_qubits = len(qubit_pauli_dict) From 2f5370662f47f0ab0438f5a7b15b8bd837ce804f Mon Sep 17 00:00:00 2001 From: ddddddanni Date: Sun, 19 Jul 2026 20:54:46 -0700 Subject: [PATCH 05/12] fix type check --- .../pauli_string_measurement_with_readout_mitigation.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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 6435377484a..059ad6ab58d 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 @@ -1185,7 +1185,7 @@ def generate_trex_and_readout_circuits( for pauli_group in circuit_to_pauli.pauli_strings: - qubit_pauli_dict = {} + qubit_pauli_dict: dict[ops.Qid, ops.Pauli] = {} for pauli_str in pauli_group: qubit_pauli_dict.update(pauli_str) From f193b347b2e2514ec5a670cd878bc35c818420f2 Mon Sep 17 00:00:00 2001 From: ddddddanni Date: Sun, 26 Jul 2026 15:56:56 -0700 Subject: [PATCH 06/12] Change based on comments --- ...ing_measurement_with_readout_mitigation.py | 31 +++++++------ ...easurement_with_readout_mitigation_test.py | 45 +++++++++++++------ 2 files changed, 46 insertions(+), 30 deletions(-) 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 059ad6ab58d..49ff5e0a9d4 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 @@ -1156,7 +1156,7 @@ def generate_trex_and_readout_circuits( num_readout_circuits: int, rng: np.random.Generator, insert_strategy: circuits.InsertStrategy = circuits.InsertStrategy.INLINE, -) -> tuple[list[circuits.Circuit], list[TRexMetadata]]: +) -> 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 @@ -1174,12 +1174,12 @@ def generate_trex_and_readout_circuits( 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. """ - all_generated_circuits: list[circuits.Circuit] = [] - metadata_list: list[TRexMetadata] = [] + results: list[tuple[list[circuits.Circuit], TRexMetadata]] = [] circuit = transformers.drop_terminal_measurements(circuit_to_pauli.circuit.unfreeze()) @@ -1193,14 +1193,11 @@ def generate_trex_and_readout_circuits( num_qubits = len(qubit_pauli_dict) twirl_choices = _generate_random_boolean_choices(num_twirls, num_qubits, rng) - # overall_flip = twirl_choices.sum(axis=1) % 2 pauli_circuits = _build_trex_twirled_pauli_circuits( circuit, joint_basis_pauli, twirl_choices, insert_strategy ) - all_generated_circuits.extend(pauli_circuits) readout_choices = _generate_random_boolean_choices(num_readout_circuits, num_qubits, rng) - # overall_readout_parity = readout_choices.sum(axis=1) % 2 readout_circuits = [ circuits.Circuit.from_moments( @@ -1213,16 +1210,18 @@ def generate_trex_and_readout_circuits( ) for readout_choices_i in readout_choices ] - all_generated_circuits.extend(readout_circuits) - metadata_list.append( - TRexMetadata( - pauli_str=joint_basis_pauli, - twirl_choices=twirl_choices, - readout_choices=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, ) - return all_generated_circuits, metadata_list + + 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 6d88f69721a..314d1081b82 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 @@ -957,20 +957,23 @@ def test_generate_trex_and_readout_circuits() -> None: num_readout_circuits = 2 rng = np.random.default_rng(seed=42) - all_circuits, metadata_list = generate_trex_and_readout_circuits( + results = generate_trex_and_readout_circuits( circuit_to_pauli=params, num_twirls=num_twirls, num_readout_circuits=num_readout_circuits, rng=rng, ) - # For each group, we should get `num_twirls` + `num_readout_circuits` circuits. - # Total circuits = (3 twirls + 2 readouts) * 2 groups = 10 circuits - assert len(all_circuits) == 10 - assert len(metadata_list) == 2 + # We have 2 Pauli groups, so we should get 2 tuples of (circuits, metadata). + assert len(results) == 2 - # Verify Group1 Metadata and Circuits - meta1 = metadata_list[0] + # Verify Group 1 Metadata and Circuits + group1_circuits, meta1 = results[0] + + # We should get (3 twirls + 2 readouts) = 5 circuits + assert len(group1_circuits) == 5 + + # 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) @@ -978,15 +981,15 @@ def test_generate_trex_and_readout_circuits() -> None: # Indices 0, 1, 2 belong to the twirled circuits of group 1 for i in range(3): # Twirled circuits should contain the base operations + basis changes + measurement - assert len(all_circuits[i]) >= len(base_circuit) - meas_op = all_circuits[i].moments[-1].operations[0] + assert len(group1_circuits[i]) >= len(base_circuit) + 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' # Indices 3, 4 belong to the readout circuits of group 1 for i in range(3, 5): - readout_circuit = all_circuits[i] + readout_circuit = group1_circuits[i] # Readout circuits have exactly 2 moments: Optional X gates, then Measurement assert len(readout_circuit.moments) == 2 @@ -1001,15 +1004,29 @@ def test_generate_trex_and_readout_circuits() -> None: assert meas_op.gate.key == 'result' # Verify Group 2 Metadata and Circuits - meta2 = metadata_list[1] + group2_circuits, meta2 = results[1] + + # We should get (3 twirls + 2 readouts) = 5 circuits + assert len(group2_circuits) == 5 + + # 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) - # Indices 8, 9 belong to the readout circuits of group 2 - for i in range(8, 10): - readout_circuit = all_circuits[i] + # Indices 0, 1, 2 belong to the twirled circuits of group 2 + for i in range(3): + # Twirled circuits should contain the base operations + basis changes + measurement + assert len(group2_circuits[i]) >= len(base_circuit) + 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' + + # 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) From 9cc57f5ae13847ac2e1072c52909b9e82c49ca02 Mon Sep 17 00:00:00 2001 From: ddddddanni Date: Sun, 26 Jul 2026 16:01:50 -0700 Subject: [PATCH 07/12] Change doc string --- .../pauli_string_measurement_with_readout_mitigation.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) 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 49ff5e0a9d4..afde953f161 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 @@ -1166,7 +1166,8 @@ 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. From d9da249ac05afa809c49a72c3acfd29a5b8d4e05 Mon Sep 17 00:00:00 2001 From: ddddddanni Date: Tue, 28 Jul 2026 21:20:10 -0700 Subject: [PATCH 08/12] Covers insert strategy in test --- ...easurement_with_readout_mitigation_test.py | 23 ++++++++++++++----- 1 file changed, 17 insertions(+), 6 deletions(-) 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 314d1081b82..cdfa30e8853 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 @@ -938,10 +938,16 @@ def test_sampler_receives_correct_circuits(use_sweep: bool) -> None: assert measured == expected_qubits -def test_generate_trex_and_readout_circuits() -> None: +@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))) + base_circuit = cirq.FrozenCircuit(cirq.Circuit(cirq.H(q0), cirq.CNOT(q0, q1), cirq.H(q1))) pauli_group_1: list[cirq.PauliString] = [cirq.PauliString(cirq.Z(q0) * cirq.Z(q1))] pauli_group_2: list[cirq.PauliString] = [ @@ -962,6 +968,7 @@ def test_generate_trex_and_readout_circuits() -> None: num_twirls=num_twirls, num_readout_circuits=num_readout_circuits, rng=rng, + insert_strategy=insert_strategy, ) # We have 2 Pauli groups, so we should get 2 tuples of (circuits, metadata). @@ -980,8 +987,10 @@ def test_generate_trex_and_readout_circuits() -> None: # Indices 0, 1, 2 belong to the twirled circuits of group 1 for i in range(3): - # Twirled circuits should contain the base operations + basis changes + measurement - assert len(group1_circuits[i]) >= len(base_circuit) + # 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 + meas_op = group1_circuits[i].moments[-1].operations[0] assert isinstance(meas_op.gate, cirq.MeasurementGate) assert set(meas_op.qubits) == {q0, q1} @@ -1017,8 +1026,10 @@ def test_generate_trex_and_readout_circuits() -> None: # Indices 0, 1, 2 belong to the twirled circuits of group 2 for i in range(3): - # Twirled circuits should contain the base operations + basis changes + measurement - assert len(group2_circuits[i]) >= len(base_circuit) + # 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 + meas_op = group2_circuits[i].moments[-1].operations[0] assert isinstance(meas_op.gate, cirq.MeasurementGate) assert set(meas_op.qubits) == {q0, q1} From 3ebc40a99c218dc7ab998ff0913c447eb550e198 Mon Sep 17 00:00:00 2001 From: ddddddanni Date: Sun, 2 Aug 2026 14:47:34 -0700 Subject: [PATCH 09/12] Fix based on comments --- .../pauli_string_measurement_with_readout_mitigation.py | 8 ++++++-- ...uli_string_measurement_with_readout_mitigation_test.py | 7 +++++-- 2 files changed, 11 insertions(+), 4 deletions(-) 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 afde953f161..1f477cadd39 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 @@ -1186,9 +1186,13 @@ def generate_trex_and_readout_circuits( for pauli_group in circuit_to_pauli.pauli_strings: - qubit_pauli_dict: dict[ops.Qid, ops.Pauli] = {} + qubit_pauli_dict_unsorted: dict[ops.Qid, ops.Pauli] = {} for pauli_str in pauli_group: - qubit_pauli_dict.update(pauli_str) + 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: cirq.PauliString = ops.PauliString(qubit_pauli_dict) num_qubits = len(qubit_pauli_dict) 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 cdfa30e8853..1a87b93a05d 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 @@ -950,9 +950,10 @@ def test_generate_trex_and_readout_circuits( base_circuit = cirq.FrozenCircuit(cirq.Circuit(cirq.H(q0), cirq.CNOT(q0, q1), cirq.H(q1))) 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(q0)), cirq.PauliString(cirq.X(q1)), + cirq.PauliString(cirq.X(q0)), ] pauli_strings = [pauli_group_1, pauli_group_2] @@ -1023,12 +1024,14 @@ def test_generate_trex_and_readout_circuits( 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) # 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(group1_circuits[i].moments[2].operations) == expected_ops_in_moment_3 + assert len(group2_circuits[i].moments[2].operations) == expected_ops_in_moment_3 meas_op = group2_circuits[i].moments[-1].operations[0] assert isinstance(meas_op.gate, cirq.MeasurementGate) From 37cc0bdb8e7cce1ec63f1e1fd0088d6772988162 Mon Sep 17 00:00:00 2001 From: ddddddanni Date: Tue, 18 Aug 2026 17:43:13 -0700 Subject: [PATCH 10/12] Fix comment --- .../pauli_string_measurement_with_readout_mitigation.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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 1f477cadd39..a66fa24998b 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 @@ -1194,7 +1194,7 @@ def generate_trex_and_readout_circuits( q: qubit_pauli_dict_unsorted[q] for q in sorted(qubit_pauli_dict_unsorted) } - joint_basis_pauli: cirq.PauliString = ops.PauliString(qubit_pauli_dict) + 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) From 32a923806bb56221bbd66aac6efc98287dd131f5 Mon Sep 17 00:00:00 2001 From: ddddddanni Date: Thu, 20 Aug 2026 18:49:05 -0700 Subject: [PATCH 11/12] Add comment explaining empty moments in readout circuits --- .../pauli_string_measurement_with_readout_mitigation.py | 3 +++ 1 file changed, 3 insertions(+) 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 a66fa24998b..c0cf061ba5b 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 @@ -1206,6 +1206,9 @@ def generate_trex_and_readout_circuits( 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 From 4f1354e0a60297da5dd01d7a44e26fc4652fd260 Mon Sep 17 00:00:00 2001 From: ddddddanni Date: Thu, 20 Aug 2026 18:51:31 -0700 Subject: [PATCH 12/12] Fix lint --- .../pauli_string_measurement_with_readout_mitigation.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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 c0cf061ba5b..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 @@ -1207,7 +1207,7 @@ def generate_trex_and_readout_circuits( 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 + # 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(