diff --git a/cirq-core/cirq/sim/simulator.py b/cirq-core/cirq/sim/simulator.py index 3afd69cfd89..4d407a8b075 100644 --- a/cirq-core/cirq/sim/simulator.py +++ b/cirq-core/cirq/sim/simulator.py @@ -85,8 +85,25 @@ def run_sweep_iter( for param_resolver in study.to_resolvers(params): records = {} if repetitions == 0: + record_shapes: dict[str, tuple[int, tuple[int, ...]]] = {} for _, op, _ in program.findall_operations_with_gate_type(ops.MeasurementGate): - records[protocols.measurement_key_name(op)] = np.empty([0, 1, 1]) + key = protocols.measurement_key_name(op) + qid_shape = protocols.qid_shape(op) + if key in record_shapes: + num_instances, expected_qid_shape = record_shapes[key] + if qid_shape != expected_qid_shape: + raise ValueError( + 'Different qid shapes for repeated measurement: ' + f'key={key!r}, prev_qid_shape={expected_qid_shape}, ' + f'qid_shape={qid_shape}' + ) + record_shapes[key] = (num_instances + 1, qid_shape) + else: + record_shapes[key] = (1, qid_shape) + records = { + key: np.empty((0, num_instances, len(qid_shape))) + for key, (num_instances, qid_shape) in record_shapes.items() + } else: records = self._run( circuit=program, param_resolver=param_resolver, repetitions=repetitions diff --git a/cirq-core/cirq/sim/simulator_test.py b/cirq-core/cirq/sim/simulator_test.py index eadb8f32f81..0883c1316a0 100644 --- a/cirq-core/cirq/sim/simulator_test.py +++ b/cirq-core/cirq/sim/simulator_test.py @@ -115,6 +115,25 @@ def test_run_simulator_run() -> None: ) +def test_run_zero_repetitions_preserves_measurement_record_shape() -> None: + q0, q1 = cirq.LineQubit.range(2) + circuit = cirq.Circuit(cirq.measure(q0, q1, key='m'), cirq.measure(q0, q1, key='m')) + + result = cirq.Simulator().run(circuit, repetitions=0) + + assert result.repetitions == 0 + assert result.records['m'].shape == (0, 2, 2) + + +def test_run_zero_repetitions_rejects_repeated_key_with_mismatched_qid_shape() -> None: + q0 = cirq.LineQid.for_qid_shape((2,))[0] + q1 = cirq.LineQid.for_qid_shape((3,))[0] + circuit = cirq.Circuit(cirq.measure(q0, key='m'), cirq.measure(q1, key='m')) + + with pytest.raises(ValueError, match='Different qid shapes for repeated measurement'): + cirq.Simulator().run(circuit, repetitions=0) + + def test_run_simulator_sweeps() -> None: expected_records = {'a': np.array([[[1]]])} simulator = FakeSimulatesSamples(expected_records)