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
38 changes: 38 additions & 0 deletions cirq-core/cirq/ops/eigen_gate_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,8 @@
from cirq import value
from cirq.testing import assert_has_consistent_trace_distance_bound

_NUMPY_SCALAR_TYPES = (np.float32, np.float64, np.double, np.int32, np.int64, np.short)


class CExpZinGate(cirq.EigenGate, cirq.testing.TwoQubitGate):
"""Two-qubit gate for the following matrix:
Expand Down Expand Up @@ -345,6 +347,16 @@ def test_is_parameterized() -> None:
assert not cirq.is_parameterized(CExpZinGate(1))
assert not cirq.is_parameterized(CExpZinGate(3))
assert cirq.is_parameterized(CExpZinGate(sympy.Symbol('a')))
for dtype in _NUMPY_SCALAR_TYPES:
assert not cirq.is_parameterized(CExpZinGate(dtype(1)))


@pytest.mark.parametrize('dtype', _NUMPY_SCALAR_TYPES)
def test_is_parameterized_numpy(dtype) -> None:
gate = CExpZinGate(dtype(1))
assert not cirq.is_parameterized(gate)
assert isinstance(gate.exponent, np.number)
assert type(gate.exponent) is dtype


@pytest.mark.parametrize('resolve_fn', [cirq.resolve_parameters, cirq.resolve_parameters_once])
Expand All @@ -358,6 +370,32 @@ def test_resolve_parameters(resolve_fn) -> None:
with pytest.raises(ValueError, match='Complex exponent'):
resolve_fn(CExpZinGate(sympy.Symbol('a')), cirq.ParamResolver({'a': 0.5j}))

for dtype in _NUMPY_SCALAR_TYPES:
assert resolve_fn(
CExpZinGate(sympy.Symbol('a')), cirq.ParamResolver({'a': dtype(1)})
) == CExpZinGate(1)
for dtype in (np.float32, np.float64, np.double):
assert resolve_fn(
CExpZinGate(sympy.Symbol('a')), cirq.ParamResolver({'a': dtype(0.5)})
) == CExpZinGate(0.5)


@pytest.mark.parametrize('dtype', _NUMPY_SCALAR_TYPES)
@pytest.mark.parametrize('resolve_fn', [cirq.resolve_parameters, cirq.resolve_parameters_once])
def test_resolve_parameters_numpy(resolve_fn, dtype) -> None:
resolved = resolve_fn(CExpZinGate(sympy.Symbol('a')), cirq.ParamResolver({'a': dtype(1)}))
assert resolved == CExpZinGate(1)
identity = resolve_fn(CExpZinGate(dtype(1)), cirq.ParamResolver({}))
assert identity == CExpZinGate(dtype(1))
assert type(identity.exponent) is dtype


@pytest.mark.parametrize('dtype', (np.float32, np.float64, np.double))
@pytest.mark.parametrize('resolve_fn', [cirq.resolve_parameters, cirq.resolve_parameters_once])
def test_resolve_parameters_numpy_half(resolve_fn, dtype) -> None:
resolved = resolve_fn(CExpZinGate(sympy.Symbol('a')), cirq.ParamResolver({'a': dtype(0.5)}))
assert resolved == CExpZinGate(0.5)


class WeightedZPowGate(cirq.EigenGate, cirq.testing.SingleQubitGate):
def __init__(self, weight, **kwargs):
Expand Down
234 changes: 234 additions & 0 deletions cirq-core/cirq/protocols/resolve_parameters_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,11 +14,51 @@

from __future__ import annotations

import numpy as np
import pytest
import sympy

import cirq

_NUMPY_SCALAR_TYPES = (np.float32, np.float64, np.double, np.int32, np.int64, np.short)
_NUMPY_FLOAT_TYPES = (np.float32, np.float64, np.double)
_EIGEN_GATES = (
cirq.XPowGate,
cirq.YPowGate,
cirq.ZPowGate,
cirq.HPowGate,
cirq.CZPowGate,
cirq.CXPowGate,
cirq.SwapPowGate,
cirq.ISwapPowGate,
cirq.ZZPowGate,
cirq.CCZPowGate,
cirq.CCXPowGate,
)
_OTHER_GATES = (
'rx',
'ry',
'rz',
'FSimGate',
'PhasedXZGate',
'GlobalPhaseGate',
'WaitGate',
'ControlledXPow',
)


def _gate_with_symbol(kind: str, a: sympy.Symbol) -> cirq.Gate:
return {
'rx': cirq.rx(a),
'ry': cirq.ry(a),
'rz': cirq.rz(a),
'FSimGate': cirq.FSimGate(a, a),
'PhasedXZGate': cirq.PhasedXZGate(x_exponent=a, z_exponent=a, axis_phase_exponent=a),
'GlobalPhaseGate': cirq.GlobalPhaseGate(a),
'WaitGate': cirq.WaitGate(cirq.Duration(nanos=a)),
'ControlledXPow': cirq.ControlledGate(cirq.XPowGate(exponent=a)),
}[kind]


@pytest.mark.parametrize('resolve_fn', [cirq.resolve_parameters, cirq.resolve_parameters_once])
def test_resolve_parameters(resolve_fn) -> None:
Expand Down Expand Up @@ -74,6 +114,24 @@ def _resolve_parameters_(self, resolver: cirq.ParamResolver, recursive: bool):
assert resolve_fn(1.1, resolver) == 1.1
assert resolve_fn(1j, resolver) == 1j

for dtype in _NUMPY_SCALAR_TYPES:
val = dtype(1)
zero = dtype(0)
assert resolve_fn(val, resolver) == val
assert type(resolve_fn(val, resolver)) is dtype
assert resolve_fn((val, val), resolver) == (val, val)
assert resolve_fn([val, val], resolver) == [val, val]
assert resolve_fn(a, {a: val}) == val
assert type(resolve_fn(a, {a: val})) is dtype
assert resolve_fn((a, b, c), {a: val, b: val, c: val}) == (val, val, val)
assert resolve_fn([a, b, c], {a: val, b: val, c: val}) == [val, val, val]
resolved_switch = resolve_fn(SimpleParameterSwitch('a'), {a: val})
assert resolved_switch.parameter == val
assert type(resolved_switch.parameter) is dtype
assert resolve_fn(SimpleParameterSwitch(zero), r).parameter == zero
assert not cirq.is_parameterized(SimpleParameterSwitch(zero))
assert cirq.is_parameterized(SimpleParameterSwitch(val))


def test_is_parameterized() -> None:
a, b = tuple(sympy.Symbol(l) for l in 'ab')
Expand All @@ -89,6 +147,13 @@ def test_is_parameterized() -> None:
assert not cirq.is_parameterized(1)
assert not cirq.is_parameterized(1.1)
assert not cirq.is_parameterized(1j)
for dtype in _NUMPY_SCALAR_TYPES:
val = dtype(1)
assert not cirq.is_parameterized(val)
assert not cirq.is_parameterized((val, val))
assert not cirq.is_parameterized([val, x])
assert cirq.is_parameterized([a, val])
assert cirq.is_parameterized((a, val))


def test_parameter_names() -> None:
Expand All @@ -103,6 +168,11 @@ def test_parameter_names() -> None:
assert cirq.parameter_names(1) == set()
assert cirq.parameter_names(1.1) == set()
assert cirq.parameter_names(1j) == set()
for dtype in _NUMPY_SCALAR_TYPES:
val = dtype(1)
assert cirq.parameter_names(val) == set()
assert cirq.parameter_names((val, val)) == set()
assert cirq.parameter_names([a, val]) == {'a'}


@pytest.mark.parametrize('resolve_fn', [cirq.resolve_parameters, cirq.resolve_parameters_once])
Expand All @@ -129,7 +199,171 @@ def test_recursive_resolve() -> None:
assert cirq.resolve_parameters_once([a, b], {a: b, b: c}) == [b, c]
assert cirq.resolve_parameters_once(a, {}) == a

for dtype in _NUMPY_SCALAR_TYPES:
resolver = cirq.ParamResolver({a: b + 3, b: c + 2, c: dtype(1)})
assert cirq.resolve_parameters(a, resolver) == 6
assert cirq.resolve_parameters(b, resolver) == 3
assert cirq.resolve_parameters(c, resolver) == 1
assert cirq.resolve_parameters_once(c, resolver) == dtype(1)
assert type(cirq.resolve_parameters_once(c, resolver)) is dtype
chained = cirq.ParamResolver({a: b, b: dtype(1)})
assert cirq.resolve_parameters(a, chained) == 1
assert type(cirq.resolve_parameters(a, chained)) is dtype
assert cirq.resolve_parameters_once(a, chained) == b
assert cirq.resolve_parameters_once(b, chained) == dtype(1)

resolver = cirq.ParamResolver({a: b, b: a})
assert cirq.resolve_parameters_once(a, resolver) == b
with pytest.raises(RecursionError):
_ = cirq.resolve_parameters(a, resolver)


@pytest.mark.parametrize('dtype', _NUMPY_SCALAR_TYPES)
@pytest.mark.parametrize('resolve_fn', [cirq.resolve_parameters, cirq.resolve_parameters_once])
def test_resolve_parameters_numpy_scalars(resolve_fn, dtype) -> None:
val = dtype(1)
assert not cirq.is_parameterized(val)
assert cirq.parameter_names(val) == set()
assert resolve_fn(val, {'a': 0}) is val
assert resolve_fn((val, val), {}) == (val, val)


@pytest.mark.parametrize('gate_cls', _EIGEN_GATES)
@pytest.mark.parametrize('dtype', _NUMPY_SCALAR_TYPES)
@pytest.mark.parametrize('resolve_fn', [cirq.resolve_parameters, cirq.resolve_parameters_once])
def test_resolve_numpy_values_on_gates(resolve_fn, gate_cls, dtype) -> None:
a = sympy.Symbol('a')
gate = gate_cls(exponent=a)
resolved = resolve_fn(gate, {a: dtype(1)})
assert not cirq.is_parameterized(resolved)
assert resolved.exponent == 1.0
assert type(resolved.exponent) is float
assert resolved == gate_cls(exponent=1.0)


@pytest.mark.parametrize('kind', _OTHER_GATES)
@pytest.mark.parametrize('dtype', _NUMPY_SCALAR_TYPES)
@pytest.mark.parametrize('resolve_fn', [cirq.resolve_parameters, cirq.resolve_parameters_once])
def test_resolve_numpy_values_on_other_gates(resolve_fn, kind, dtype) -> None:
a = sympy.Symbol('a')
gate = _gate_with_symbol(kind, a)
resolved = resolve_fn(gate, {a: dtype(1)})
expected = resolve_fn(gate, {a: 1})
assert not cirq.is_parameterized(resolved)
assert resolved == expected


@pytest.mark.parametrize('kind', ('rx', 'ry', 'rz', 'FSimGate', 'PhasedXZGate', 'WaitGate'))
@pytest.mark.parametrize('dtype', _NUMPY_FLOAT_TYPES)
@pytest.mark.parametrize('resolve_fn', [cirq.resolve_parameters, cirq.resolve_parameters_once])
def test_resolve_numpy_half_on_other_gates(resolve_fn, kind, dtype) -> None:
a = sympy.Symbol('a')
gate = _gate_with_symbol(kind, a)
resolved = resolve_fn(gate, {a: dtype(0.5)})
expected = resolve_fn(gate, {a: 0.5})
assert not cirq.is_parameterized(resolved)
assert resolved == expected


@pytest.mark.parametrize('dtype', _NUMPY_SCALAR_TYPES)
def test_numpy_exponent_is_not_parameterized(dtype) -> None:
gate = cirq.XPowGate(exponent=dtype(1))
assert not cirq.is_parameterized(gate)
assert cirq.parameter_names(gate) == set()
assert cirq.resolve_parameters(gate, {'a': 0.5}) is gate


def test_numpy_double_exponent_float_isinstance_back_compat() -> None:
# https://github.com/quantumlib/Cirq/issues/5758#issuecomment-3608357176
gate = cirq.XPowGate(exponent=np.double(0.5))
is_parameterized = not isinstance(gate.exponent, float)
assert is_parameterized is False
assert not cirq.is_parameterized(gate)
assert gate.exponent == 0.5
assert type(gate.exponent) is np.float64


@pytest.mark.parametrize('dtype', (np.float32, np.int32, np.int64, np.short))
def test_numpy_nonfloat64_exponent_cirq_is_parameterized(dtype) -> None:
gate = cirq.XPowGate(exponent=dtype(1))
assert not cirq.is_parameterized(gate)
assert isinstance(gate.exponent, np.number)
assert type(gate.exponent) is dtype


@pytest.mark.parametrize('dtype', _NUMPY_SCALAR_TYPES)
@pytest.mark.parametrize('resolve_fn', [cirq.resolve_parameters, cirq.resolve_parameters_once])
def test_resolve_numpy_values_on_operations(resolve_fn, dtype) -> None:
q0, q1, q2 = cirq.LineQubit.range(3)
a = sympy.Symbol('a')
val = dtype(1)
ops = [
cirq.XPowGate(exponent=a).on(q0),
cirq.CZPowGate(exponent=a).on(q0, q1),
cirq.CCZPowGate(exponent=a).on(q0, q1, q2),
cirq.rx(a).on(q0),
cirq.FSimGate(a, a).on(q0, q1),
cirq.WaitGate(cirq.Duration(nanos=a)).on(q0),
cirq.ControlledGate(cirq.XPowGate(exponent=a)).on(q0, q1),
cirq.PhasedXZGate(x_exponent=a, z_exponent=a, axis_phase_exponent=a).on(q0),
]
for op in ops:
resolved = resolve_fn(op, {'a': val})
expected = resolve_fn(op, {'a': 1})
assert not cirq.is_parameterized(resolved)
assert resolved == expected


@pytest.mark.parametrize('dtype', (np.float32, np.float64, np.double))
@pytest.mark.parametrize('resolve_fn', [cirq.resolve_parameters, cirq.resolve_parameters_once])
def test_resolve_numpy_half_exponent_on_operations(resolve_fn, dtype) -> None:
q = cirq.LineQubit(0)
a = sympy.Symbol('a')
op = cirq.XPowGate(exponent=a).on(q)
half = resolve_fn(op, {'a': dtype(0.5)})
assert half == cirq.XPowGate(exponent=0.5).on(q)


@pytest.mark.parametrize('dtype', _NUMPY_SCALAR_TYPES)
@pytest.mark.parametrize('resolve_fn', [cirq.resolve_parameters, cirq.resolve_parameters_once])
def test_resolve_numpy_values_on_phased_x(resolve_fn, dtype) -> None:
a = sympy.Symbol('a')
gate = cirq.PhasedXPowGate(phase_exponent=a, exponent=a)
resolved = resolve_fn(gate, {a: dtype(1)})
assert not cirq.is_parameterized(resolved)
assert resolved.exponent == 1.0
assert resolved.phase_exponent == 1


@pytest.mark.parametrize('dtype', (np.float32, np.float64, np.double))
def test_phased_x_numpy_phase_exponent_canonicalize(dtype) -> None:
gate = cirq.PhasedXPowGate(phase_exponent=dtype(1.5))
assert gate.phase_exponent == -0.5
assert type(gate.phase_exponent) is dtype
assert isinstance(gate.phase_exponent, np.number)


@pytest.mark.parametrize('dtype', (np.int32, np.int64, np.short))
def test_phased_x_numpy_int_phase_exponent_canonicalize(dtype) -> None:
gate = cirq.PhasedXPowGate(phase_exponent=dtype(3))
assert gate.phase_exponent == 1
assert type(gate.phase_exponent) is dtype


@pytest.mark.parametrize('dtype', _NUMPY_SCALAR_TYPES)
def test_resolve_numpy_values_on_circuit(dtype) -> None:
q0, q1 = cirq.LineQubit.range(2)
a = sympy.Symbol('a')
circuit = cirq.Circuit(
cirq.XPowGate(exponent=a).on(q0),
cirq.HPowGate(exponent=a).on(q0),
cirq.rx(a).on(q0),
cirq.CZPowGate(exponent=a).on(q0, q1),
cirq.FSimGate(a, a).on(q0, q1),
cirq.WaitGate(cirq.Duration(nanos=a)).on(q0),
cirq.PhasedXZGate(x_exponent=a, z_exponent=0, axis_phase_exponent=0).on(q0),
)
resolved = cirq.resolve_parameters(circuit, {a: dtype(1)})
expected = cirq.resolve_parameters(circuit, {a: 1})
assert not cirq.is_parameterized(resolved)
assert resolved == expected
30 changes: 30 additions & 0 deletions cirq-core/cirq/study/resolver_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,7 +40,10 @@ def test_symbol() -> None:
1,
np.int32(45),
np.float64(6.3),
np.double(6.3),
np.int32(2),
np.int64(7),
np.short(3),
np.complex64(1j),
np.complex128(2j),
1j,
Expand Down Expand Up @@ -137,12 +140,39 @@ def test_value_of_calculations() -> None:
assert r.value_of(sympy.Symbol('b') / 0.1 - sympy.Symbol('a')) == 0.5


@pytest.mark.parametrize(
'a_val,b_val',
[
(np.float32(0.5), np.float32(0.1)),
(np.float64(0.5), np.float64(0.1)),
(np.double(0.5), np.double(0.1)),
(np.int32(2), np.int32(4)),
(np.int64(2), np.int64(4)),
(np.short(2), np.short(4)),
],
)
def test_value_of_calculations_numpy(a_val, b_val) -> None:
r = cirq.ParamResolver({'a': a_val, 'b': b_val})
a = sympy.Symbol('a')
b = sympy.Symbol('b')
assert r.value_of(a + b) == a_val + b_val
assert r.value_of(b - a) == b_val - a_val
assert r.value_of(a * b) == a_val * b_val


def test_resolve_integer_division() -> None:
r = cirq.ParamResolver({'a': 1, 'b': 2})
resolved = r.value_of(sympy.Symbol('a') / sympy.Symbol('b'))
assert resolved == 0.5


@pytest.mark.parametrize('dtype', (np.float32, np.float64, np.double, np.int32, np.int64, np.short))
def test_resolve_integer_division_numpy(dtype) -> None:
r = cirq.ParamResolver({'a': dtype(1), 'b': dtype(2)})
resolved = r.value_of(sympy.Symbol('a') / sympy.Symbol('b'))
assert resolved == pytest.approx(0.5)


def test_resolve_symbol_division() -> None:
B = sympy.Symbol('B')
r = cirq.ParamResolver({'a': 1, 'b': B})
Expand Down
Loading
Loading