Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 11 additions & 4 deletions tensorflow_quantum/core/serialize/serializer.py
Original file line number Diff line number Diff line change
Expand Up @@ -808,17 +808,22 @@ def serialize_circuit(circuit_inp):
and `cirq.LineQubit` instances during serialization of circuits.

Note: once serialized terminal measurements are removed.
Note: the input circuit is not modified.

Args:
circuit_inp: A `cirq.Circuit`.

Returns:
A `tfq.proto.Program` proto.
"""
circuit = copy.deepcopy(circuit_inp)
if not isinstance(circuit, cirq.Circuit):
if not isinstance(circuit_inp, cirq.Circuit):
raise TypeError("serialize requires cirq.Circuit objects."
" Given: " + str(type(circuit)))
" Given: " + str(type(circuit_inp)))

# A shallow copy is enough. The rewrites below replace whole moments
# rather than mutating them in place, and the one operation that does get
# mutated (control demotion, further down) is copied before tagging.
circuit = circuit_inp.copy()

# This code is intentionally written to avoid using cirq functions
# as this get analyzed by tensorflow-autograph.
Expand Down Expand Up @@ -870,7 +875,9 @@ def serialize_circuit(circuit_inp):
]
new_ops = dict()
for op in controlled_ops:
tfq_compatible = op.sub_operation
# Copy before tagging: `sub_operation` is a live reference into
# the caller's circuit, and these attributes would leak onto it.
tfq_compatible = copy.copy(op.sub_operation)
tfq_compatible._tfq_control_qubits = op.controls
tfq_compatible._tfq_control_values = op.control_values
new_ops[op.qubits] = tfq_compatible
Expand Down
25 changes: 25 additions & 0 deletions tensorflow_quantum/core/serialize/serializer_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -892,6 +892,31 @@ def test_terminal_measurement_support(self):
with self.assertRaisesRegex(ValueError, expected_regex="non-terminal"):
serializer.serialize_circuit(invalid_circuit)

def test_serialize_does_not_tag_input_circuit(self):
"""Serialization must not leave control tags on the input circuit."""
q0 = cirq.GridQubit(0, 0)
q1 = cirq.GridQubit(0, 1)
q2 = cirq.GridQubit(0, 2)
circuit = cirq.Circuit([
cirq.H(q0),
cirq.ControlledOperation([q0],
cirq.ISWAP(q1, q2)**0.5)
])
circuit_before_call = copy.deepcopy(circuit)

serializer.serialize_circuit(circuit)

# Note that `==` cannot catch this. The demotion tags are plain Python
# attributes and take no part in circuit or operation equality, so they
# have to be checked for directly.
for moment in circuit:
for op in moment:
sub_op = getattr(op, 'sub_operation', op)
self.assertFalse(hasattr(sub_op, '_tfq_control_qubits'))
self.assertFalse(hasattr(sub_op, '_tfq_control_values'))

self.assertEqual(circuit, circuit_before_call)

def test_serialize_deserialize_identity(self):
"""Confirm that identity gates can be serialized and deserialized."""
q0 = cirq.GridQubit(0, 0)
Expand Down
Loading