Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
110 changes: 109 additions & 1 deletion perceval/components/linear_circuit.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,20 +31,23 @@

import copy
import random
from warnings import warn

from abc import ABC, abstractmethod
from collections.abc import Callable
from collections.abc import Callable, Iterable
from typing import Generator

import numpy as np
import sympy as sp
import scipy.optimize as so

from exqalibur import permanent_cx, estimate_permanent_cx
from perceval.components.abstract_component import AParametrizedComponent
from perceval.utils import Parameter, Matrix, MatrixN, matrix_double, global_params, InterferometerShape
from perceval.utils.logging import get_logger, channel
from perceval.utils.algorithms import decomposition, Match
from perceval.utils.algorithms.solve import solve
from perceval.utils.states import FockState


class ACircuit(AParametrizedComponent, ABC):
Expand Down Expand Up @@ -93,6 +96,111 @@ def compute_unitary(self,
return matrix_double(u)
return u

def compute_fock_matrix(self,
basis: Iterable[FockState],
perm_method: str,
n_iter = None):
"""Compute the matrix representation of the circuit in the Fock space
for a particular basis.

:param basis: Basis of states over which to calculate the transition
matrix.
:param perm_method: Permanent calculation method. Currently supported
exact methods: "ryser", "glynn". Currently supported approximate
methods: "gurvits".
:param n_iter: Number of iterations for convergence using approximate
permanent estimation method
"""
assert isinstance(perm_method, str), (
"Permanent computation method must be a string."
)
perm_method = perm_method.lower()

assert self.defined, (
'All parameters must be defined to compute numeric Fock matrix.'
)

if self.requires_polarization:
raise NotImplementedError(
"`compute_fock_matrix` cannot currently be called on "
"polarized circuits."
)

basis = tuple(basis)
prodnfact_dict = {}
mat_idx_dict = {}

for state in basis:
assert isinstance(state, FockState), (
"`basis` members must be of type `FockState`."
)
assert state.m == self.m, (
f"State {state} in basis does not match the correct number of "
f"modes. ({state.m} != {self.m})"
)

# Precompute the state's product factorial
prodnfact_dict[state] = state.prodnfact()

# Pre-compute the matrix columns/row indices for each state
rows = []
for mode, photon_count in enumerate(state):
rows.extend([mode] * photon_count)

mat_idx_dict[state] = rows

if perm_method in {"ryser", "glynn"}:
if n_iter is not None:
warn("`n_iter` argument is ignored for exact method: " +
perm_method, stacklevel=2
)
perm_estimator = lambda x: permanent_cx(
x,
ptype=perm_method
)
elif perm_method == "gurvits":
assert n_iter is not None, (
"`n_iter` must be specified for approximate permanent "
"estimation method: " + perm_method
)
perm_estimator = lambda x: estimate_permanent_cx(
x,
ptype=perm_method,
n_iter=n_iter
)[0]
else:
raise NotImplementedError(
perm_method + " permanent calculation method not supported."
)

U = self.compute_unitary(use_symbolic=False)
dim = len(basis)
U_fock = np.zeros((dim, dim), dtype=complex)

# Construct the matrix U_{t, s}
for i, t in enumerate(basis):
rows = mat_idx_dict[t]
t_prodnfact = prodnfact_dict[t]

for j, s in enumerate(basis):

# No transition amplitude between Fock states with differing n
if t.n != s.n:
continue

if t.n == 0 and s.n == 0:
perm = 1
else:
cols = mat_idx_dict[s]
U_ts = U[np.ix_(rows, cols)]
perm = perm_estimator(U_ts)

denom = (t_prodnfact * prodnfact_dict[s]) ** 0.5

U_fock[i, j] = perm / denom

return U_fock

@property
def requires_polarization(self):
"""Does the circuit require polarization?
Expand Down
22 changes: 21 additions & 1 deletion tests/components/test_circuit.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,7 +35,7 @@
from perceval import Circuit, P, BasicState, pdisplay, Matrix, BackendFactory, Processor
from perceval.rendering.pdisplay import pdisplay_circuit, pdisplay_matrix
from perceval.rendering.format import Format
from perceval.utils import InterferometerShape
from perceval.utils import InterferometerShape, FockState
import perceval.algorithm as algo
import perceval.components.unitary_components as comp
from .._test_utils import strip_line_12, assert_circuits_eq
Expand Down Expand Up @@ -468,3 +468,23 @@ def test_phase_noise():
circuit = comp.PS(0.775)
circuit.apply_phase_noise(0, 0.25)
assert pytest.approx(float(circuit.param("phi"))) == 0.75


def test_compute_fock_matrix():
circ = comp.Unitary.random(5) // comp.BS() // comp.PS(0.4)
backend = BackendFactory.get_backend("SLOS")

basis = [
FockState([3, 0, 0, 0, 0]),
FockState([0, 3, 0, 0, 0]),
]

fock_matrix = circ.compute_fock_matrix(basis=basis, perm_method="ryser")

backend = BackendFactory.get_backend("SLOS")
backend.set_circuit(circ)
backend.set_input_state(basis[0])
evolved_state = backend.evolve()

for output_idx, output_state in enumerate(basis):
assert fock_matrix[output_idx, 0] == pytest.approx(evolved_state[output_state])
Loading