use hugr_core::builder::{
DFGBuilder, Dataflow, DataflowHugr, DataflowSubContainer, HugrBuilder, SubContainer, endo_sig,
inout_sig,
};
use hugr_core::extension::prelude::{UnwrapBuilder, option_type};
use hugr_core::hugr::linking::{NameLinkingPolicy, OnMultiDefn};
use hugr_core::ops::constant::{CustomConst, OpaqueValue};
use hugr_core::ops::{OpTrait, OpType, Tag, Value};
use hugr_core::std_extensions::arithmetic::conversions::ConvertOpDef;
use hugr_core::std_extensions::arithmetic::int_ops::IntOpDef;
use hugr_core::std_extensions::arithmetic::int_types::{ConstInt, INT_TYPES};
use hugr_core::std_extensions::collections::array::{
Array, ArrayClone, ArrayDiscard, ArrayKind, ArrayOpBuilder, GenericArrayOpDef,
GenericArrayRepeat, GenericArrayScan, GenericArrayValue, array_type,
};
use hugr_core::std_extensions::collections::borrow_array::{
BArrayClone, BArrayDiscard, BArrayOpBuilder, BorrowArray, borrow_array_type,
};
use hugr_core::std_extensions::collections::list::ListValue;
use hugr_core::types::{SumType, Transformable, Type, TypeArg};
use hugr_core::{Visibility, type_row};
use itertools::Itertools;
use crate::mangle_name;
use super::{
CallbackHandler, LinearizeError, Linearizer, NodeTemplate, ReplaceTypes, ReplaceTypesError,
};
#[deprecated(
note = "`hugr-passes` is deprecated. Use tket::passes instead",
since = "0.26.2"
)]
pub fn list_const(
val: &OpaqueValue,
repl: &ReplaceTypes,
) -> Result<Option<Value>, ReplaceTypesError> {
let Some(lv) = val.value().downcast_ref::<ListValue>() else {
return Ok(None);
};
let mut elem_t = lv.get_element_type().clone();
if !elem_t.transform(repl)? {
return Ok(None);
}
let mut vals: Vec<Value> = lv.get_contents().to_vec();
for v in &mut vals {
repl.change_value(v)?;
}
Ok(Some(ListValue::new(elem_t, vals).into()))
}
#[deprecated(
note = "`hugr-passes` is deprecated. Use tket::passes instead",
since = "0.26.2"
)]
pub fn generic_array_const<AK: ArrayKind>(
val: &OpaqueValue,
repl: &ReplaceTypes,
) -> Result<Option<Value>, ReplaceTypesError>
where
GenericArrayValue<AK>: CustomConst,
{
let Some(av) = val.value().downcast_ref::<GenericArrayValue<AK>>() else {
return Ok(None);
};
let mut elem_t = av.get_element_type().clone();
if !elem_t.transform(repl)? {
return Ok(None);
}
let mut vals: Vec<Value> = av.get_contents().to_vec();
for v in &mut vals {
repl.change_value(v)?;
}
Ok(Some(GenericArrayValue::<AK>::new(elem_t, vals).into()))
}
#[deprecated(
note = "`hugr-passes` is deprecated. Use tket::passes instead",
since = "0.26.2"
)]
pub fn array_const(
val: &OpaqueValue,
repl: &ReplaceTypes,
) -> Result<Option<Value>, ReplaceTypesError> {
generic_array_const::<Array>(val, repl)
}
pub(super) const DISCARD_TO_UNIT_PREFIX: &str = "__discard_unit";
pub(super) const COPY_SCAN_PREFIX: &str = "__copy_scan";
pub(super) const UNWRAP_PREFIX: &str = "__unwrap";
pub(super) const MAKE_NONE_PREFIX: &str = "__mk_none";
#[deprecated(
note = "`hugr-passes` is deprecated. Use tket::passes instead",
since = "0.26.2"
)]
pub fn linearize_generic_array<AK: ArrayKind>(
args: &[TypeArg],
num_outports: usize,
lin: &CallbackHandler,
) -> Result<NodeTemplate, LinearizeError> {
let [TypeArg::BoundedNat(n), TypeArg::Runtime(ty)] = args else {
panic!("Illegal TypeArgs to array: {args:?}")
};
if num_outports == 0 {
let array_scan = GenericArrayScan::<AK>::new(ty.clone(), Type::UNIT, vec![], *n);
let in_type = AK::ty(*n, ty.clone());
return Ok(NodeTemplate::LinkedHugr(
Box::new({
let mut dfb = DFGBuilder::new(inout_sig([in_type], [])).unwrap();
let map_fn = {
let mut mb = dfb.module_root_builder();
let mut fb = mb
.define_function_vis(
mangle_name(DISCARD_TO_UNIT_PREFIX, &[ty.clone().into()]),
inout_sig([ty.clone()], [Type::UNIT]),
Visibility::Public,
)
.unwrap();
let [to_discard] = fb.input_wires_arr();
let disc = lin.copy_discard_op(ty, 0)?;
disc.add(&mut fb, [to_discard]).map_err(|e| {
LinearizeError::NestedTemplateError(Box::new(ty.clone()), Box::new(e))
})?;
let ret = fb.add_load_value(Value::unary_unit_sum());
fb.finish_with_outputs([ret]).unwrap()
};
let [in_array] = dfb.input_wires_arr();
let map_fn = dfb.load_func(map_fn.handle(), &[]).unwrap();
let unit_arr = dfb
.add_dataflow_op(array_scan, [in_array, map_fn])
.unwrap()
.out_wire(0);
AK::build_discard(&mut dfb, Type::UNIT, *n, unit_arr).unwrap();
dfb.finish_hugr_with_outputs([]).unwrap()
}),
NameLinkingPolicy::default().on_multiple_defn(OnMultiDefn::UseSource),
));
}
let num_new = num_outports - 1;
let array_ty = AK::ty(*n, ty.clone());
let mut dfb = DFGBuilder::new(inout_sig(
[array_ty.clone()],
vec![array_ty.clone(); num_outports],
))
.unwrap();
let option_sty = option_type([ty.clone()]);
let option_ty = Type::from(option_sty.clone());
let arrays_of_none = {
let fn_none = {
let mut mb = dfb.module_root_builder();
let mut fb = mb
.define_function_vis(
mangle_name(MAKE_NONE_PREFIX, &[ty.clone().into()]),
inout_sig(vec![], [option_ty.clone()]),
Visibility::Public,
)
.unwrap();
let none = fb
.add_dataflow_op(Tag::new(0, vec![type_row![], [ty.clone()].into()]), [])
.unwrap();
fb.finish_with_outputs(none.outputs()).unwrap()
};
let repeats = vec![GenericArrayRepeat::<AK>::new(option_ty.clone(), *n); num_new];
let fn_none = dfb.load_func(fn_none.handle(), &[]).unwrap();
repeats
.into_iter()
.map(|rpt| {
let [arr] = dfb.add_dataflow_op(rpt, [fn_none]).unwrap().outputs_arr();
arr
})
.collect::<Vec<_>>()
};
let i64_t = INT_TYPES[6].clone();
let option_array = AK::ty(*n, option_ty.clone());
let copy_elem = {
let mut io = vec![ty.clone(), i64_t.clone()];
io.extend(vec![option_array.clone(); num_new]);
let mut mb = dfb.module_root_builder();
let mut fb = mb
.define_function_vis(
mangle_name(
COPY_SCAN_PREFIX,
&[(*n).into(), ty.clone().into(), (num_new as u64).into()],
),
endo_sig(io),
Visibility::Public,
)
.unwrap();
let mut inputs = fb.input_wires();
let elem = inputs.next().unwrap();
let idx = inputs.next().unwrap();
let opt_arrays = inputs.collect::<Vec<_>>();
let [idx_usz] = fb
.add_dataflow_op(ConvertOpDef::itousize.without_log_width(), [idx])
.unwrap()
.outputs_arr();
let mut copies = lin
.copy_discard_op(ty, num_outports)?
.add(&mut fb, [elem])
.map_err(|e| LinearizeError::NestedTemplateError(Box::new(ty.clone()), Box::new(e)))?
.outputs();
let copy0 = copies.next().unwrap();
let set_op = OpType::from(GenericArrayOpDef::<AK>::set.to_concrete(option_ty.clone(), *n));
let either_st = set_op.dataflow_signature().unwrap().output[0]
.as_sum()
.unwrap()
.clone();
let opt_arrays = opt_arrays
.into_iter()
.zip_eq(copies)
.map(|(opt_array, copy1)| {
let [tag] = fb
.add_dataflow_op(Tag::new(1, vec![type_row![], [ty.clone()].into()]), [copy1])
.unwrap()
.outputs_arr();
let [set_result] = fb
.add_dataflow_op(set_op.clone(), [opt_array, idx_usz, tag])
.unwrap()
.outputs_arr();
let [none, opt_array] = fb
.build_unwrap_sum(1, either_st.clone(), set_result)
.unwrap();
let [] = fb
.build_unwrap_sum(0, SumType::new_option([ty.clone()]), none)
.unwrap();
opt_array
})
.collect::<Vec<_>>();
let cst1 = fb.add_load_value(ConstInt::new_u(6, 1).unwrap());
let [new_idx] = fb
.add_dataflow_op(IntOpDef::iadd.with_log_width(6), [idx, cst1])
.unwrap()
.outputs_arr();
fb.finish_with_outputs([copy0, new_idx].into_iter().chain(opt_arrays))
.unwrap()
};
let [in_array] = dfb.input_wires_arr();
let scan1 = GenericArrayScan::<AK>::new(
ty.clone(),
ty.clone(),
std::iter::once(i64_t)
.chain(vec![option_array; num_new])
.collect(),
*n,
);
let copy_elem = dfb.load_func(copy_elem.handle(), &[]).unwrap();
let cst0 = dfb.add_load_value(ConstInt::new_u(6, 0).unwrap());
let mut outs = dfb
.add_dataflow_op(
scan1,
[in_array, copy_elem, cst0]
.into_iter()
.chain(arrays_of_none),
)
.unwrap()
.outputs();
let out_array1 = outs.next().unwrap();
let _idx_out = outs.next().unwrap();
let opt_arrays = outs;
let unwrap_elem = {
let mut mb = dfb.module_root_builder();
let mut fb = mb
.define_function_vis(
mangle_name(UNWRAP_PREFIX, &[ty.clone().into()]),
inout_sig([option_ty.clone()], [ty.clone()]),
Visibility::Public,
)
.unwrap();
let [opt] = fb.input_wires_arr();
let [val] = fb.build_unwrap_sum(1, option_sty, opt).unwrap();
fb.finish_with_outputs([val]).unwrap()
};
let unwrap_scan = GenericArrayScan::<AK>::new(option_ty, ty.clone(), vec![], *n);
let unwrap_elem = dfb.load_func(unwrap_elem.handle(), &[]).unwrap();
let out_arrays = std::iter::once(out_array1)
.chain(opt_arrays.map(|opt_array| {
let [out_array] = dfb
.add_dataflow_op(unwrap_scan.clone(), [opt_array, unwrap_elem])
.unwrap()
.outputs_arr();
out_array
}))
.collect::<Vec<_>>();
Ok(NodeTemplate::LinkedHugr(
Box::new(dfb.finish_hugr_with_outputs(out_arrays).unwrap()),
NameLinkingPolicy::default().on_multiple_defn(OnMultiDefn::UseSource),
))
}
#[deprecated(
note = "`hugr-passes` is deprecated. Use tket::passes instead",
since = "0.26.2"
)]
pub fn copy_discard_array(
args: &[TypeArg],
num_outports: usize,
lin: &CallbackHandler,
) -> Result<NodeTemplate, LinearizeError> {
let [TypeArg::BoundedNat(n), TypeArg::Runtime(ty)] = args else {
panic!("Illegal TypeArgs to array: {args:?}")
};
if ty.copyable() {
if num_outports == 0 {
Ok(NodeTemplate::SingleOp(
ArrayDiscard::new(ty.clone(), *n).unwrap().into(),
))
} else if num_outports == 2 {
Ok(NodeTemplate::SingleOp(
ArrayClone::new(ty.clone(), *n).unwrap().into(),
))
} else {
let array_ty = array_type(*n, ty.clone());
Ok(NodeTemplate::CompoundOp(Box::new({
let mut dfb =
DFGBuilder::new(inout_sig([array_ty.clone()], vec![array_ty; *n as usize]))
.unwrap();
let [mut arr] = dfb.input_wires_arr();
let mut outs = vec![];
for _ in 0..(num_outports - 1) {
let (arr1, arr2) = dfb.add_array_clone(ty.clone(), *n, arr).unwrap();
arr = arr1;
outs.push(arr2);
}
outs.push(arr);
dfb.finish_hugr_with_outputs(outs).unwrap()
})))
}
} else {
linearize_generic_array::<Array>(args, num_outports, lin)
}
}
#[deprecated(
note = "`hugr-passes` is deprecated. Use tket::passes instead",
since = "0.26.2"
)]
pub fn copy_discard_borrow_array(
args: &[TypeArg],
num_outports: usize,
lin: &CallbackHandler,
) -> Result<NodeTemplate, LinearizeError> {
let [TypeArg::BoundedNat(n), TypeArg::Runtime(ty)] = args else {
panic!("Illegal TypeArgs to borrow array: {args:?}")
};
if ty.copyable() {
if num_outports == 0 {
Ok(NodeTemplate::SingleOp(
BArrayDiscard::new(ty.clone(), *n).unwrap().into(),
))
} else if num_outports == 2 {
Ok(NodeTemplate::SingleOp(
BArrayClone::new(ty.clone(), *n).unwrap().into(),
))
} else {
let array_ty = borrow_array_type(*n, ty.clone());
Ok(NodeTemplate::CompoundOp(Box::new({
let mut dfb =
DFGBuilder::new(inout_sig([array_ty.clone()], vec![array_ty; *n as usize]))
.unwrap();
let [mut arr] = dfb.input_wires_arr();
let mut outs = vec![];
for _ in 0..(num_outports - 1) {
let (arr1, arr2) = dfb.add_borrow_array_clone(ty.clone(), *n, arr).unwrap();
arr = arr1;
outs.push(arr2);
}
outs.push(arr);
dfb.finish_hugr_with_outputs(outs).unwrap()
})))
}
} else if num_outports == 0 {
let elem_discard = lin.copy_discard_op(ty, 0)?;
let array_ty = || borrow_array_type(*n, ty.clone());
let i64_t = || INT_TYPES[6].clone();
let mut dfb = DFGBuilder::new(inout_sig([array_ty()], [])).unwrap();
let [in_array] = dfb.input_wires_arr();
let zero = dfb.add_load_value(ConstInt::new_u(6, 0).unwrap());
let one = dfb.add_load_value(ConstInt::new_u(6, 1).unwrap());
let len = dfb.add_load_value(ConstInt::new_u(6, *n).unwrap());
let mut tl = dfb
.tail_loop_builder([(i64_t(), zero), (array_ty(), in_array)], [], type_row![])
.unwrap();
let [idx, arr] = tl.input_wires_arr();
let [in_range] = tl
.add_dataflow_op(IntOpDef::ilt_u.with_log_width(6), [idx, len])
.unwrap()
.outputs_arr();
let loop_variants = vec![[i64_t(), array_ty()].into(), type_row![]];
let mut cond = tl
.conditional_builder(
(vec![type_row![]; 2], in_range),
[(array_ty(), arr)],
[Type::new_sum(loop_variants.clone())].into(),
)
.unwrap();
{
let mut out_range = cond.case_builder(0).unwrap();
let [arr] = out_range.input_wires_arr();
let () = out_range
.add_discard_all_borrowed(ty.clone(), *n, arr)
.unwrap();
let res = out_range
.add_dataflow_op(Tag::new(1, loop_variants.clone()), [])
.unwrap();
out_range.finish_with_outputs(res.outputs()).unwrap();
}
{
let mut in_range = cond.case_builder(1).unwrap();
let [arr] = in_range.input_wires_arr();
let [idx_u] = in_range
.add_dataflow_op(ConvertOpDef::itousize.without_log_width(), [idx])
.unwrap()
.outputs_arr();
let (arr, is_borrowed) = in_range
.add_is_borrowed(ty.clone(), *n, arr, idx_u)
.unwrap();
let mut cond2 = in_range
.conditional_builder(
(vec![type_row![]; 2], is_borrowed),
[(array_ty(), arr)],
[array_ty()].into(),
)
.unwrap();
{
let borrowed_case = cond2.case_builder(1).unwrap();
let [arr] = borrowed_case.input_wires_arr();
borrowed_case.finish_with_outputs([arr]).unwrap();
}
{
let mut not_borrowed_case = cond2.case_builder(0).unwrap();
let [arr] = not_borrowed_case.input_wires_arr();
let (arr, elem) = not_borrowed_case
.add_borrow_array_borrow(ty.clone(), *n, arr, idx_u)
.unwrap();
elem_discard.add(&mut not_borrowed_case, [elem]).unwrap();
not_borrowed_case.finish_with_outputs([arr]).unwrap();
}
let [arr_out] = cond2.finish_sub_container().unwrap().outputs_arr();
let [idx_out] = in_range
.add_dataflow_op(IntOpDef::iadd.with_log_width(6), [idx, one])
.unwrap()
.outputs_arr();
let res = in_range
.add_dataflow_op(Tag::new(0, loop_variants), [idx_out, arr_out])
.unwrap();
in_range.finish_with_outputs(res.outputs()).unwrap();
}
let [loop_pred] = cond.finish_sub_container().unwrap().outputs_arr();
let [] = tl.finish_with_outputs(loop_pred, []).unwrap().outputs_arr();
let h = dfb.finish_hugr_with_outputs([]).unwrap();
Ok(NodeTemplate::CompoundOp(Box::new(h)))
} else {
linearize_generic_array::<BorrowArray>(args, num_outports, lin)
}
}
#[cfg(test)]
mod test {
use hugr_core::builder::{DFGBuilder, Dataflow, DataflowHugr};
use hugr_core::{
extension::prelude::usize_t, std_extensions::collections::borrow_array::borrow_array_type,
types::Signature,
};
use crate::replace_types::{DelegatingLinearizer, Linearizer};
#[test]
fn test_borrow_array_discard() {
let arr_ty = borrow_array_type(5, borrow_array_type(7, usize_t()));
let dl = DelegatingLinearizer::default();
let mut dfb = DFGBuilder::new(Signature::new([arr_ty.clone()], [])).unwrap();
let nt = dl.copy_discard_op(&arr_ty, 0).unwrap();
let ins = dfb.input_wires();
nt.add(&mut dfb, ins).unwrap();
dfb.finish_hugr_with_outputs([]).unwrap();
}
}