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)
|