File size: 11,728 Bytes
4173e5f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4cdc0c8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4173e5f
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
"""
Tests for dense_evolution.autodiff.circuit_to_energy_fn — the differentiable
VQE engine extracted from dashboard_core.py (_build_vqe_template /
_vqe_energy_fn) so it's reachable as public API, independent of Streamlit/
pandas/dashboard, and composable with the Qiskit/PennyLane interop bridge.
"""
import numpy as np
import pytest

import dense_evolution as de
from dense_evolution import autodiff

jax = pytest.importorskip("jax")
import jax.numpy as jnp


VQE_QASM = (
    'OPENQASM 2.0; include "qelib1.inc"; qreg q[2]; creg c[2]; '
    'ry(0.5) q[0]; rx(0.5) q[1]; cx q[0],q[1]; rz(0.2) q[1]; cx q[0],q[1]; '
    'ry(0.5) q[0]; rx(0.5) q[1]; measure q -> c;'
)


def _random_hamiltonian(n_qubits, seed=7):
    rng = np.random.default_rng(seed)
    values = np.sort(rng.uniform(-2.5, 2.5, 2 ** n_qubits))
    return jnp.diag(jnp.array(values, dtype=jnp.float64))


class TestCircuitToEnergyFn:

    def test_n_params_matches_parametric_gate_count(self):
        circ = de.QASMParser().parse(VQE_QASM)
        _, n_params = de.circuit_to_energy_fn(circ, circ.n_qubits)
        # ry, rx, rz, ry, rx = 5 parametric gates in VQE_QASM
        assert n_params == 5

    def test_gradient_matches_finite_difference(self):
        circ = de.QASMParser().parse(VQE_QASM)
        energy_fn, n_params = de.circuit_to_energy_fn(circ, circ.n_qubits)
        h_matrix = _random_hamiltonian(circ.n_qubits)

        rng = np.random.default_rng(3)
        theta0 = rng.uniform(-np.pi, np.pi, n_params)

        energy_and_grad = jax.jit(jax.value_and_grad(energy_fn, argnums=0, has_aux=True))
        (_, _), grad = energy_and_grad(jnp.asarray(theta0), h_matrix)

        eps = 1e-6
        fd_grad = np.zeros(n_params)
        for i in range(n_params):
            tp, tm = theta0.copy(), theta0.copy()
            tp[i] += eps
            tm[i] -= eps
            (ep, _), _ = energy_and_grad(jnp.asarray(tp), h_matrix)
            (em, _), _ = energy_and_grad(jnp.asarray(tm), h_matrix)
            fd_grad[i] = float((ep - em) / (2 * eps))

        np.testing.assert_allclose(np.asarray(grad), fd_grad, atol=1e-6)

    def test_default_stato_zero_is_ground_state(self):
        circ = de.QASMParser().parse('OPENQASM 2.0; include "qelib1.inc"; qreg q[2]; creg c[2]; measure q -> c;')
        energy_fn, n_params = de.circuit_to_energy_fn(circ, circ.n_qubits)
        assert n_params == 0
        h_matrix = jnp.diag(jnp.array([1.0, 2.0, 3.0, 4.0], dtype=jnp.float64))
        energy, sv = energy_fn(jnp.array([]), h_matrix)
        # empty circuit -> stays |00>, energy = h_matrix[0,0]
        assert abs(float(energy) - 1.0) < 1e-9
        assert abs(abs(complex(sv[0])) - 1.0) < 1e-9

    def test_optimization_loop_standalone_energy_descends(self):
        # Same correctness bar as test_dashboard_core.py's
        # test_vqe_energy_trends_downward_over_epochs, but built directly on
        # circuit_to_energy_fn with a hand-rolled Adam loop — no
        # run_vqe_telemetry/dashboard_core involved at all. Demonstrates the
        # public API is self-sufficient for real VQE outside the dashboard.
        circ = de.QASMParser().parse(VQE_QASM)
        energy_fn, n_params = de.circuit_to_energy_fn(circ, circ.n_qubits)
        h_matrix = _random_hamiltonian(circ.n_qubits, seed=11)

        rng = np.random.default_rng(11)
        theta = rng.uniform(-np.pi, np.pi, n_params)
        m, v = np.zeros(n_params), np.zeros(n_params)
        lr, beta1, beta2 = 0.1, 0.9, 0.999

        energy_and_grad = jax.jit(jax.value_and_grad(energy_fn, argnums=0, has_aux=True))
        energies = []
        for epoch in range(1, 41):
            (energy, _), grad = energy_and_grad(jnp.asarray(theta), h_matrix)
            energies.append(float(energy))
            grad = np.asarray(grad)
            m = beta1 * m + (1 - beta1) * grad
            v = beta2 * v + (1 - beta2) * (grad ** 2)
            m_hat = m / (1 - beta1 ** epoch)
            v_hat = v / (1 - beta2 ** epoch)
            theta -= (lr / (np.sqrt(v_hat) + 1e-8)) * m_hat

        assert np.mean(energies[-5:]) < np.mean(energies[:5])

    def test_from_pennylane_circuit_is_now_differentiable(self):
        # This is the gap found while auditing the interop bridge last
        # turn: jax.grad through run_pennylane_circuit silently returns 0.0
        # because from_pennylane bakes theta into a plain float. Composed
        # through circuit_to_energy_fn instead (which works on the
        # QASMCircuit from_pennylane returns, before that float baking
        # matters), a real circuit must now show a non-zero gradient.
        pennylane = pytest.importorskip("pennylane")
        import pennylane as qml

        dev = qml.device('default.qubit', wires=2)

        @qml.qnode(dev)
        def circuit(theta):
            qml.RY(theta[0], wires=0)
            qml.RX(theta[1], wires=1)
            qml.CNOT(wires=[0, 1])
            return qml.probs(wires=[0, 1])

        theta0 = np.array([0.5, 0.5])
        circ = de.from_pennylane(circuit, theta0)
        energy_fn, n_params = de.circuit_to_energy_fn(circ, circ.n_qubits)
        assert n_params == 2

        h_matrix = _random_hamiltonian(circ.n_qubits, seed=5)

        def energy(theta):
            e, _ = energy_fn(theta, h_matrix)
            return e

        grad = jax.grad(energy)(jnp.asarray(theta0))
        assert float(jnp.linalg.norm(grad)) > 1e-6


class TestEnergyFnNoiseSpec:
    """circuit_to_energy_fn's `noise=` argument (registry.NoiseSpec, a
    JAX PyTree): applies NoiseModel.apply_to_sv natively inside the same
    traced computation as theta, instead of as an external Python-side
    step the caller has to splice in around energy_fn. jax_key is a
    pytree leaf (not an aux_data/static field), so it flows through
    jit/grad/vmap the way any other JAX array does -- no external
    key-management workaround, no OS-entropy fallback."""

    def test_default_none_matches_pre_noise_behavior(self):
        circ = de.QASMParser().parse(VQE_QASM)
        energy_fn, n_params = de.circuit_to_energy_fn(circ, circ.n_qubits)
        h_matrix = _random_hamiltonian(circ.n_qubits)
        theta = jnp.asarray(np.random.default_rng(1).uniform(-np.pi, np.pi, n_params))
        e_no_arg, _ = energy_fn(theta, h_matrix)
        e_explicit_none, _ = energy_fn(theta, h_matrix, None, None)
        assert float(e_no_arg) == pytest.approx(float(e_explicit_none))

    def test_same_key_is_reproducible(self):
        circ = de.QASMParser().parse(VQE_QASM)
        energy_fn, n_params = de.circuit_to_energy_fn(circ, circ.n_qubits)
        h_matrix = _random_hamiltonian(circ.n_qubits)
        theta = jnp.asarray(np.random.default_rng(2).uniform(-np.pi, np.pi, n_params))
        noise = de.NoiseSpec(model='depolarizing', p=0.15, jax_key=jax.random.PRNGKey(42))
        e1, _ = energy_fn(theta, h_matrix, noise=noise)
        e2, _ = energy_fn(theta, h_matrix, noise=noise)
        assert float(e1) == pytest.approx(float(e2))

    def test_noisy_energy_differs_from_ideal(self):
        circ = de.QASMParser().parse(VQE_QASM)
        energy_fn, n_params = de.circuit_to_energy_fn(circ, circ.n_qubits)
        h_matrix = _random_hamiltonian(circ.n_qubits)
        theta = jnp.asarray(np.random.default_rng(3).uniform(-np.pi, np.pi, n_params))
        e_ideal, _ = energy_fn(theta, h_matrix)
        noise = de.NoiseSpec(model='depolarizing', p=0.3, jax_key=jax.random.PRNGKey(7))
        e_noisy, _ = energy_fn(theta, h_matrix, noise=noise)
        assert float(e_ideal) != pytest.approx(float(e_noisy))

    def test_zero_probability_reduces_to_ideal_even_when_traced(self):
        # p=0.0 passed as a *traced* pytree leaf (under jax.jit) used to
        # crash with TracerBoolConversionError -- apply_to_sv's early
        # `if p <= 0.0: return sv` shortcut couldn't take a Python bool
        # of a tracer. Fixed to try/except around that optimization only;
        # every channel's math already reduces to a no-op at p=0
        # (`fire = r < p` is always False), so falling through instead of
        # short-circuiting is still correct.
        circ = de.QASMParser().parse(VQE_QASM)
        energy_fn, n_params = de.circuit_to_energy_fn(circ, circ.n_qubits)
        h_matrix = _random_hamiltonian(circ.n_qubits)
        theta = jnp.asarray(np.random.default_rng(4).uniform(-np.pi, np.pi, n_params))
        e_ideal, _ = energy_fn(theta, h_matrix)

        @jax.jit
        def run(p):
            noise = de.NoiseSpec(model='depolarizing', p=p, jax_key=jax.random.PRNGKey(1))
            e, _ = energy_fn(theta, h_matrix, None, noise)
            return e

        e_zero = run(0.0)
        assert float(e_zero) == pytest.approx(float(e_ideal))

    def test_composable_with_jax_jit(self):
        circ = de.QASMParser().parse(VQE_QASM)
        energy_fn, n_params = de.circuit_to_energy_fn(circ, circ.n_qubits)
        h_matrix = _random_hamiltonian(circ.n_qubits)
        theta = jnp.asarray(np.random.default_rng(5).uniform(-np.pi, np.pi, n_params))
        noise = de.NoiseSpec(model='depolarizing', p=0.1, jax_key=jax.random.PRNGKey(3))

        e_eager, _ = energy_fn(theta, h_matrix, None, noise)
        e_jit, _ = jax.jit(energy_fn)(theta, h_matrix, None, noise)
        assert float(e_eager) == pytest.approx(float(e_jit))

    def test_composable_with_jax_grad_through_noise(self):
        # gradient w.r.t. theta must flow through the noisy pipeline,
        # not just the ideal one
        circ = de.QASMParser().parse(VQE_QASM)
        energy_fn, n_params = de.circuit_to_energy_fn(circ, circ.n_qubits)
        h_matrix = _random_hamiltonian(circ.n_qubits)
        theta = jnp.asarray(np.random.default_rng(6).uniform(-np.pi, np.pi, n_params))
        noise = de.NoiseSpec(model='depolarizing', p=0.1, jax_key=jax.random.PRNGKey(9))

        def energy(t):
            e, _ = energy_fn(t, h_matrix, None, noise)
            return e

        grad = jax.grad(energy)(theta)
        assert float(jnp.linalg.norm(grad)) > 0.0

    def test_composable_with_jax_vmap_over_keys(self):
        # a batch of independent, reproducible noise realizations,
        # natively -- no external Python loop managing keys
        circ = de.QASMParser().parse(VQE_QASM)
        energy_fn, n_params = de.circuit_to_energy_fn(circ, circ.n_qubits)
        h_matrix = _random_hamiltonian(circ.n_qubits)
        theta = jnp.asarray(np.random.default_rng(8).uniform(-np.pi, np.pi, n_params))
        keys = jax.random.split(jax.random.PRNGKey(0), 6)

        def run(k):
            noise = de.NoiseSpec(model='depolarizing', p=0.15, jax_key=k)
            e, _ = energy_fn(theta, h_matrix, None, noise)
            return e

        batch = jax.vmap(run)(keys)
        assert batch.shape == (6,)
        # independent keys must not all collapse to the same value
        assert float(jnp.std(batch)) > 0.0

    def test_qubits_subset_still_works(self):
        circ = de.QASMParser().parse(VQE_QASM)
        energy_fn, n_params = de.circuit_to_energy_fn(circ, circ.n_qubits)
        h_matrix = _random_hamiltonian(circ.n_qubits)
        theta = jnp.asarray(np.random.default_rng(9).uniform(-np.pi, np.pi, n_params))
        noise = de.NoiseSpec(model='bitflip', p=0.5, jax_key=jax.random.PRNGKey(2), qubits=[0])
        e, sv = energy_fn(theta, h_matrix, None, noise)
        assert sv.shape == (2 ** circ.n_qubits,)


class TestImportSafety:

    def test_root_import_never_fails(self):
        assert hasattr(de, 'circuit_to_energy_fn')

    def test_missing_jax_raises_clear_importerror(self, monkeypatch):
        monkeypatch.setattr(autodiff, 'HAS_JAX', False)
        with pytest.raises(ImportError, match='JAX'):
            autodiff.circuit_to_energy_fn(None, 1)