File size: 6,634 Bytes
4173e5f
 
 
 
 
b0b2232
4173e5f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b0b2232
 
 
 
4173e5f
 
 
 
b0b2232
 
 
 
 
 
 
 
 
 
 
 
 
4173e5f
 
 
 
 
 
 
 
 
 
 
 
 
b0b2232
 
4173e5f
 
 
 
 
 
 
 
 
 
 
b0b2232
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from typing import Callable, Optional, Tuple

from .parser import QASMCircuit
from .gates import GATE_IDS
from .compiler import QuantumTranspiler
from .registry import HAS_JAX, NoiseModel, NoiseSpec

if HAS_JAX:
    import jax
    import jax.numpy as jnp
    from .compiler import _compile_and_run_circuit_jit
else:
    jnp = None


#: gates that receive a value from theta β€” must match the n_params count
#: below exactly, or theta's allocation order desyncs from the template's
#: injection order. Kept identical to dashboard_core's own list (same
#: engine, one source of truth).
_PARAMETRIC_GATES = ('rx', 'ry', 'rz', 'u1', 'p', 'cp', 'crz')
_TWO_QUBIT_GATES = ('cx', 'cy', 'cz', 'cp', 'crz', 'swap')


def _require_jax():
    if not HAS_JAX:
        raise ImportError(
            "circuit_to_energy_fn requires JAX. "
            "Install it with: pip install dense-evolution[jax]")


def _build_template(circuit: QASMCircuit, n_qubits: int) -> "jnp.ndarray":
    """Builds the (n_ops, 4) float64 [g_id, q1, q2, sentinel] template that
    the energy function injects theta into (-1.0 in the param slot for
    gates whose value comes from theta, patched in via jnp.where inside a
    jax.lax.scan, never a Python float() β€” that would sever the JAX trace).

    Structural pass only: build (name, *qubits) tuples (no param values β€”
    QuantumTranspiler.transpile only inspects gate name/qubit-count, for
    ccx/swap decomposition), transpile once, then look up g_id per gate and
    mark parametric slots with the sentinel. ccx/toffoli decomposes into
    non-parametric gates only, so this never desyncs theta's order.

    circuit.ops qubits are always plain ints here (QASMCircuit is the
    package's own interchange type, produced by QASMParser.parse and by
    the Qiskit/PennyLane interop bridge alike) β€” no defensive unwrapping
    of framework-specific qubit/wire objects needed.
    """
    tuples = []
    for op in circuit.ops:
        name = str(op['name']).lower().strip()
        qubits = [int(q) for q in op.get('qubits', [])]
        if not qubits or any(q >= n_qubits for q in qubits):
            continue
        tuples.append((name, *qubits))

    target = QuantumTranspiler.transpile(tuples)

    rows = []
    for cmd in target:
        name = cmd[0].lower()
        if name not in GATE_IDS:
            continue
        g_id = float(GATE_IDS[name])
        qubits = cmd[1:]
        sentinel = -1.0 if name in _PARAMETRIC_GATES else 0.0
        if name in _TWO_QUBIT_GATES and len(qubits) >= 2:
            rows.append([g_id, float(qubits[0]), float(qubits[1]), sentinel])
        elif qubits:
            rows.append([g_id, float(qubits[0]), 0.0, sentinel])

    if not rows:
        return jnp.empty((0, 4), dtype=jnp.float64)
    return jnp.array(rows, dtype=jnp.float64)


def circuit_to_energy_fn(
    circuit: QASMCircuit, n_qubits: int
) -> Tuple[Callable, int]:
    """
    Convert a QASMCircuit into a JAX-differentiable energy function.

    circuit : QASMCircuit β€” from QASMParser.parse(qasm), or from the
              Qiskit/PennyLane interop bridge (from_qiskit/from_pennylane).

    Returns (energy_fn, n_params):
      energy_fn(theta, h_matrix, stato_zero=None, noise=None) ->
      (energy, statevector) is a pure JAX function, differentiable w.r.t.
      theta via jax.grad / jax.value_and_grad(energy_fn, argnums=0,
      has_aux=True). stato_zero defaults to |0...0> if not given.
      n_params is the number of parametric gates in the circuit, in the
      same order theta is injected β€” build theta as an array of that
      length.

      noise, when given, is a registry.NoiseSpec (a JAX PyTree) applied
      to the statevector right after the circuit and before the energy
      expectation value is computed β€” natively inside the same traced
      computation as theta, not as an external step the caller has to
      splice in around energy_fn themselves. Because NoiseSpec carries
      its own jax_key as a pytree leaf, the whole thing stays
      jit/grad/vmap-composable with no OS-entropy fallback and no
      external key-management workaround:

          noise = NoiseSpec(model='depolarizing', p=0.05,
                             jax_key=jax.random.PRNGKey(0))
          energy, sv = energy_fn(theta, h_matrix, noise=noise)

    This is the same engine dashboard_core.py's real VQE gradient uses
    internally (verified against finite differences, ~1e-11 agreement) β€”
    exposed here as public API so it's reachable without reading
    dashboard_core.py, and so circuits imported via from_qiskit/
    from_pennylane (which are NOT differentiable on their own β€” see
    run_pennylane_circuit's docstring) have a real way to become
    differentiable instead of just a documented dead end.
    """
    _require_jax()
    template = _build_template(circuit, n_qubits)
    n_params = sum(1 for op in circuit.ops
                   if str(op['name']).lower().strip() in _PARAMETRIC_GATES)

    def energy_fn(theta, h_matrix, stato_zero: Optional["jnp.ndarray"] = None,
                  noise: Optional["NoiseSpec"] = None):
        if stato_zero is None:
            stato_zero = jnp.zeros(2 ** n_qubits, dtype=jnp.complex128).at[0].set(1.0)

        if n_params == 0:
            # No parametric gates -> no sentinel (-1.0) rows in template, so
            # patch_and_apply below would never take its is_param branch.
            # Skip the scan entirely rather than index into an empty theta
            # array during tracing (n_params is a static Python int, fixed
            # at circuit_to_energy_fn() call time, so this branch is
            # resolved before any tracing happens β€” not a jax.lax.cond).
            sv = _compile_and_run_circuit_jit(stato_zero, template)
        else:
            def patch_and_apply(carry, op):
                idx = carry
                is_param = op[3] == -1.0
                final_p = jnp.where(is_param, theta[idx], op[3])
                next_idx = jnp.where(is_param, idx + jnp.int32(1), idx)
                return next_idx, jnp.array([op[0], op[1], op[2], final_p], dtype=jnp.float64)

            _, patched_ops = jax.lax.scan(patch_and_apply, jnp.int32(0), template)
            sv = _compile_and_run_circuit_jit(stato_zero, patched_ops)

        if noise is not None:
            sv = NoiseModel.apply_to_sv(
                sv, n_qubits, model=noise.model, p=noise.p,
                jax_key=noise.jax_key, qubits=list(noise.qubits) if noise.qubits else None,
            )

        energy = jnp.real(jnp.vdot(sv, h_matrix @ sv))
        return energy, sv

    return energy_fn, n_params