use std::collections::HashMap;
use crate::operations::Operation;
use crate::registers::Registers;
use crate::registers::{BitOutputRegister, ComplexOutputRegister, FloatOutputRegister};
use crate::Circuit;
use crate::{
measurements::{Measure, MeasureExpectationValues},
RoqoqoBackendError,
};
#[cfg(feature = "async")]
use async_trait::async_trait;
pub type RegisterResult = Result<Registers, RoqoqoBackendError>;
pub trait EvaluatingBackend: Sized {
fn run_circuit(&self, circuit: &Circuit) -> RegisterResult {
self.run_circuit_iterator(circuit.iter())
}
fn run_circuit_iterator<'a>(
&self,
circuit: impl Iterator<Item = &'a Operation>,
) -> RegisterResult;
fn run_measurement_registers<T>(&self, measurement: &T) -> RegisterResult
where
T: Measure,
{
let mut bit_registers: HashMap<String, BitOutputRegister> = HashMap::new();
let mut float_registers: HashMap<String, FloatOutputRegister> = HashMap::new();
let mut complex_registers: HashMap<String, ComplexOutputRegister> = HashMap::new();
for circuit in measurement.circuits() {
let (tmp_bit_reg, tmp_float_reg, tmp_complex_reg) = match measurement.constant_circuit()
{
Some(x) => self.run_circuit_iterator(x.iter().chain(circuit.iter()))?,
None => self.run_circuit_iterator(circuit.iter())?,
};
for (key, mut val) in tmp_bit_reg.into_iter() {
if let Some(x) = bit_registers.get_mut(&key) {
x.append(&mut val);
} else {
let _ = bit_registers.insert(key, val);
}
}
for (key, mut val) in tmp_float_reg.into_iter() {
if let Some(x) = float_registers.get_mut(&key) {
x.append(&mut val);
} else {
let _ = float_registers.insert(key, val);
}
}
for (key, mut val) in tmp_complex_reg.into_iter() {
if let Some(x) = complex_registers.get_mut(&key) {
x.append(&mut val);
} else {
let _ = complex_registers.insert(key, val);
}
}
}
Ok((bit_registers, float_registers, complex_registers))
}
fn run_measurement<T>(
&self,
measurement: &T,
) -> Result<Option<HashMap<String, f64>>, RoqoqoBackendError>
where
T: MeasureExpectationValues,
{
let (bit_registers, float_registers, complex_registers) =
self.run_measurement_registers(measurement)?;
Ok(measurement.evaluate(bit_registers, float_registers, complex_registers)?)
}
}
#[cfg(feature = "async")]
#[async_trait]
pub trait AsyncEvaluatingBackend: Sized {
async fn async_run_circuit(&self, circuit: &Circuit) -> RegisterResult {
self.async_run_circuit_iterator(circuit.iter()).await
}
async fn async_run_circuit_iterator<'a>(
&self,
circuit: impl Iterator<Item = &'a Operation> + std::marker::Send,
) -> RegisterResult;
async fn async_run_measurement_registers<T>(&self, measurement: &T) -> RegisterResult
where
T: Measure,
T: std::marker::Sync,
{
let mut bit_registers: HashMap<String, BitOutputRegister> = HashMap::new();
let mut float_registers: HashMap<String, FloatOutputRegister> = HashMap::new();
let mut complex_registers: HashMap<String, ComplexOutputRegister> = HashMap::new();
let mut circuit_futures = Vec::new();
for circuit in measurement.circuits() {
let circuit_future = match measurement.constant_circuit() {
Some(x) => self.async_run_circuit_iterator(x.iter().chain(circuit.iter())),
None => self.async_run_circuit_iterator(circuit.iter()),
};
circuit_futures.push(circuit_future)
}
let circuit_results = futures::future::try_join_all(circuit_futures).await?;
for (tmp_bit_reg, tmp_float_reg, tmp_complex_reg) in circuit_results {
for (key, mut val) in tmp_bit_reg.into_iter() {
if let Some(x) = bit_registers.get_mut(&key) {
x.append(&mut val);
} else {
let _ = bit_registers.insert(key, val);
}
}
for (key, mut val) in tmp_float_reg.into_iter() {
if let Some(x) = float_registers.get_mut(&key) {
x.append(&mut val);
} else {
let _ = float_registers.insert(key, val);
}
}
for (key, mut val) in tmp_complex_reg.into_iter() {
if let Some(x) = complex_registers.get_mut(&key) {
x.append(&mut val);
} else {
let _ = complex_registers.insert(key, val);
}
}
}
Ok((bit_registers, float_registers, complex_registers))
}
async fn async_run_measurement<T>(
&self,
measurement: &T,
) -> Result<Option<HashMap<String, f64>>, RoqoqoBackendError>
where
T: MeasureExpectationValues,
T: std::marker::Sync,
{
let (bit_registers, float_registers, complex_registers) =
self.async_run_measurement_registers(measurement).await?;
Ok(measurement.evaluate(bit_registers, float_registers, complex_registers)?)
}
}