Skip to content

Single qubit: phase estimation

This example shows a one-qubit interference experiment.

import itertools

import equinox as eqx
import jax
import jax.numpy as jnp
import seaborn as sns
import ultraplot as uplt
from rich.pretty import pprint

from squint.interface.base import Circuit, Wire
from squint.interface.dv import DiscreteVariableState, HGate, RZGate
from squint.backends.tensornetwork.simulator import Simulator
from squint.math.information_matrices import quantum_fisher_information_matrix, classical_fisher_information_matrix
from squint.utils import partition_op
---------------------------------------------------------------------------
ModuleNotFoundError                       Traceback (most recent call last)
Cell In[2], line 12
     10 from squint.interface.base import Circuit, Wire
     11 from squint.interface.dv import DiscreteVariableState, HGate, RZGate
---> 12 from squint.backends.tensornetwork.simulator import Simulator
     13 from squint.utils import partition_op

File ~/Desktop/1 - Projects/Quantum Intelligence Lab/repos/squint/src/squint/backends/tensornetwork/simulator.py:39
     35 from ordered_set import OrderedSet
     37 __all__ = ["SimulatorQuantumAmplitudes", "SimulatorClassicalProbabilities", "Simulator"]
---> 39 from squint.circuit import Circuit
     40 from squint.ops import (
     41     AbstractErasureChannel,
     42     AbstractGate,
   (...)
     46     AbstractPureState,
     47 )
     48 from squint.interface.base import Block, wire_sort_key

ModuleNotFoundError: No module named 'squint.circuit'
wire = Wire(dim=2, idx=0)

circuit = Circuit()

#          ____      ___________      ____
# |0> --- | H | --- | Rz(\phi) | --- | H | ----
#         ----      -----------      ----

circuit.add(DiscreteVariableState(wires=(wire,), n=(0,)))
circuit.add(HGate(wires=(wire,)))
circuit.add(RZGate(wires=(wire,), phi=0.0 * jnp.pi), "phase")
circuit.add(HGate(wires=(wire,)))

pprint(circuit)
params, static = partition_op(circuit, "phase")
sim = Simulator(static=static, params=params)

def forward_probs(p):
    return jnp.abs(sim.forward(p)) ** 2

grad_probs = jax.jacfwd(forward_probs)
ket = sim.forward(params)
dket = sim.grad(params)
prob = forward_probs(params)
dprob = grad_probs(params)

print(f"Shape of ket is: {ket.shape}, with dtype {ket.dtype}")
print(f"Shape of prob is: {prob.shape}, with dtype {prob.dtype}")
Shape of ket is: (2,), with dtype complex128
Shape of prob is: (2,), with dtype float64
phis = jnp.linspace(-jnp.pi, jnp.pi, 100)
params = eqx.tree_at(lambda pytree: pytree.ops["phase"].phi, params, phis)

probs = jax.vmap(forward_probs)(params)
qfims = jax.vmap(lambda p: quantum_fisher_information_matrix(sim.forward, sim.grad, p))(params)
cfims = jax.vmap(lambda p: classical_fisher_information_matrix(forward_probs, grad_probs, p))(params)
colors = sns.color_palette("Set2", n_colors=jnp.prod(jnp.array(probs.shape[1:])))
fig, axs = uplt.subplots(nrows=2, figsize=(6, 4), sharey=False)

for i, idx in enumerate(
    itertools.product(*[list(range(ell)) for ell in probs.shape[1:]])
):
    axs[0].plot(phis, probs[:, *idx], label=f"{idx}", color=colors[i])
axs[0].legend()
axs[0].set(xlabel=r"Phase, $\varphi$", ylabel=r"Probability, $p(\mathbf{x} | \varphi)$")

axs[1].plot(phis, qfims.squeeze(), color=colors[0], label=r"$\mathcal{I}_\varphi^Q$")
axs[1].plot(phis, cfims.squeeze(), color=colors[-1], label=r"$\mathcal{I}_\varphi^C$")
axs[1].set(
    xlabel=r"Phase, $\varphi$",
    ylabel=r"Fisher Information, $\mathcal{I}_\varphi$",
    ylim=[0, 1.05 * jnp.max(qfims)],
)
axs[1].legend();

img