use super::*;
use crate::Result;
use crate::operation::{OpCode, OpKind, Operation};
use crate::wire::WireList;
use qudit_core::array::Tensor;
use qudit_core::{ComplexScalar, HybridSystem, ParamIndices, ParamInfo};
use qudit_expr::FUNCTION;
use qudit_expr::index::IndexDirection;
use qudit_tensor::{QuditCircuitTensorNetworkBuilder, QuditTensor, QuditTensorNetwork};
impl QuditCircuit {
pub fn kraus_ops<C: ComplexScalar>(&self, args: &[C::R]) -> Tensor<C, 3> {
let network = self.to_tensor_network();
let code = qudit_tensor::compile_network(network);
let mut tnvm =
qudit_tensor::TNVM::<C, FUNCTION>::new(&code, Some(&self.params.const_map()));
let result = tnvm.evaluate::<FUNCTION>(args);
result.get_fn_result2().unpack_tensor3d().to_owned()
}
pub fn to_tensor_network(&self) -> QuditTensorNetwork {
self.as_tensor_network_builder().build()
}
pub fn as_tensor_network_builder(&self) -> QuditCircuitTensorNetworkBuilder {
let mut network = QuditCircuitTensorNetworkBuilder::new(
self.qudit_radices(),
Some(self.operations.expressions()),
);
for inst in self.iter() {
network = self
.add_instruction_to_builder(network, inst.op_code(), inst.wires(), inst.params())
.expect("TODO");
}
network
}
fn add_instruction_to_builder(
&self,
mut network: QuditCircuitTensorNetworkBuilder,
op_code: OpCode,
wires: WireList,
params: ParamIndices,
) -> Result<QuditCircuitTensorNetworkBuilder> {
if op_code.kind() == OpKind::Expression {
let param_indices = self.params.convert_ids_to_indices(params);
let constant = param_indices
.iter()
.map(|i| self.params[i].is_assigned())
.collect();
let param_info = ParamInfo::new(param_indices, constant);
let indices = self.operations.indices(op_code);
let input_index_map = if indices
.iter()
.any(|idx| idx.direction() == IndexDirection::Input && idx.index_size() > 1)
{
wires.qudits().collect()
} else {
vec![]
};
let output_index_map = if indices
.iter()
.any(|idx| idx.direction() == IndexDirection::Output && idx.index_size() > 1)
{
wires.qudits().collect()
} else {
vec![]
};
let batch_index_map: Vec<String> = wires.dits().map(|id| id.to_string()).collect();
let tensor = QuditTensor::new(indices, op_code.id(), param_info);
network = network.prepend(tensor, input_index_map, output_index_map, batch_index_map);
Ok(network)
} else {
let op = self
.operations
.get(op_code)
.ok_or(crate::Error::MissingOperation(op_code))?;
match op {
Operation::Expression(_) => unreachable!("Already handled expressions."),
Operation::Subcircuit(sub) => {
for sub_inst in sub.iter() {
let mapped_wires = sub_inst.wires().map_through(&wires);
let mapped_params = sub_inst.params().map_through(¶ms);
network = self.add_instruction_to_builder(
network,
sub_inst.op_code(),
mapped_wires,
mapped_params,
)?;
}
Ok(network)
}
Operation::Directive(_) => Ok(network),
}
}
}
}