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}")
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();
