Skip to content

Commit ca5f19c

Browse files
authored
feat: Add a borrow_array type replacement pass (#975)
Draft because performance benchmarks with Guppy using this are needed first to see if this is suitable. There are probably a lot of ways that some of the more repetitive code in all the `dest` functions could be generalised but for now I prioritised getting something working. Closes [#2397](Quantinuum/hugr#2397) (but should add a new issue for doing this properly in LLVM at some point)
1 parent 64412df commit ca5f19c

4 files changed

Lines changed: 1312 additions & 10 deletions

File tree

‎tket-qsystem/src/lib.rs‎

Lines changed: 18 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@ pub mod llvm;
99
mod lower_drops;
1010
pub mod pytket;
1111
pub mod replace_bools;
12+
pub mod replace_borrow_arrays;
1213

1314
use derive_more::{Display, Error, From};
1415
use hugr::{
@@ -24,6 +25,7 @@ use hugr::{
2425
};
2526
use lower_drops::LowerDropsPass;
2627
use replace_bools::{ReplaceBoolPass, ReplaceBoolPassError};
28+
use replace_borrow_arrays::{ReplaceBorrowArrayPass, ReplaceBorrowArrayPassError};
2729
use tket::TketOp;
2830

2931
use extension::{
@@ -44,6 +46,7 @@ pub struct QSystemPass {
4446
monomorphize: bool,
4547
force_order: bool,
4648
lazify: bool,
49+
lower_borrow_arrays: bool,
4750
}
4851

4952
impl Default for QSystemPass {
@@ -53,6 +56,7 @@ impl Default for QSystemPass {
5356
monomorphize: true,
5457
force_order: true,
5558
lazify: true,
59+
lower_borrow_arrays: true,
5660
}
5761
}
5862
}
@@ -63,6 +67,8 @@ impl Default for QSystemPass {
6367
pub enum QSystemPassError<N = Node> {
6468
/// An error from the component [ReplaceBoolPass].
6569
ReplaceBoolError(ReplaceBoolPassError<N>),
70+
/// An error from the component [ReplaceBorrowArrayPass].
71+
ReplaceBorrowArrayError(ReplaceBorrowArrayPassError<N>),
6672
/// An error from the component [force_order()] pass.
6773
ForceOrderError(HugrError),
6874
/// An error from the component [LowerTketToQSystemPass] pass.
@@ -116,12 +122,19 @@ impl QSystemPass {
116122
}
117123

118124
self.lower_tk2().run(hugr)?;
125+
// Only has an effect if there are borrow arrays - should be run before
126+
// lazification and drop lowering as no copy discard handler currently exists
127+
// for BorrowArrays.
128+
if self.lower_borrow_arrays {
129+
self.replace_borrow_arrays().run(hugr)?;
130+
}
119131
if self.lazify {
120132
self.replace_bools().run(hugr)?;
121133
}
122-
123134
// We expect any Hugr will have *either* drop ops, or ValueArrays (without drops),
124135
// so only one of these passes will do anything; the order is thus immaterial.
136+
// Drop should come after borrow array replacement so that we don't require a
137+
// copy/discard handler for borrow arrays + avoid lowering the discard function.
125138
self.lower_drops().run(hugr)?;
126139
self.linearize_arrays().run(hugr)?;
127140

@@ -190,6 +203,10 @@ impl QSystemPass {
190203
ReplaceBoolPass
191204
}
192205

206+
fn replace_borrow_arrays(&self) -> ReplaceBorrowArrayPass {
207+
ReplaceBorrowArrayPass
208+
}
209+
193210
fn constant_fold(&self) -> ConstantFoldPass {
194211
ConstantFoldPass::default()
195212
}

‎tket-qsystem/src/lower_drops.rs‎

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,10 +2,14 @@
22
use hugr::algorithms::replace_types::{NodeTemplate, ReplaceTypesError, ReplacementOptions};
33
use hugr::algorithms::{ComposablePass, ReplaceTypes};
44
use hugr::builder::{Container, DFGBuilder};
5+
use hugr::extension::prelude::bool_t;
6+
use hugr::extension::simple_op::MakeRegisteredOp;
57
use hugr::types::{Signature, Term};
68
use hugr::{hugr::hugrmut::HugrMut, Node};
79
use tket::extension::guppy::{DROP_OP_NAME, GUPPY_EXTENSION};
810

11+
use crate::extension::futures::{future_type, FutureOp, FutureOpDef};
12+
913
/// A pass that lowers "drop" ops from [GUPPY_EXTENSION]
1014
#[derive(Default, Debug, Clone)]
1115
pub struct LowerDropsPass;
@@ -18,6 +22,30 @@ impl<H: HugrMut<Node = Node>> ComposablePass<H> for LowerDropsPass {
1822

1923
fn run(&self, hugr: &mut H) -> Result<Self::Result, Self::Error> {
2024
let mut rt = ReplaceTypes::default();
25+
26+
// future(bool) is not in the default linearizer handler so we add it here.
27+
// TODO: Create ReplaceTypes with future(bool) linearized by default to avoid
28+
// code duplication with ReplaceBools pass.
29+
let dup_op = FutureOp {
30+
op: FutureOpDef::Dup,
31+
typ: bool_t(),
32+
}
33+
.to_extension_op()
34+
.unwrap();
35+
let free_op = FutureOp {
36+
op: FutureOpDef::Free,
37+
typ: bool_t(),
38+
}
39+
.to_extension_op()
40+
.unwrap();
41+
rt.linearizer()
42+
.register_simple(
43+
future_type(bool_t()).as_extension().unwrap().clone(),
44+
NodeTemplate::SingleOp(dup_op.into()),
45+
NodeTemplate::SingleOp(free_op.into()),
46+
)
47+
.unwrap();
48+
2149
rt.replace_parametrized_op_with(
2250
GUPPY_EXTENSION.get_op(DROP_OP_NAME.as_str()).unwrap(),
2351
|targs| {

‎tket-qsystem/src/replace_bools.rs‎

Lines changed: 114 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -5,27 +5,34 @@ mod static_array;
55
use derive_more::{Display, Error, From};
66
use hugr::{
77
algorithms::{
8+
ensure_no_nonlocal_edges,
89
non_local::FindNonLocalEdgesError,
9-
replace_types::{NodeTemplate, ReplaceTypesError},
10+
replace_types::{NodeTemplate, ReplaceTypesError, ReplacementOptions},
1011
ComposablePass, ReplaceTypes,
1112
},
1213
builder::{
13-
inout_sig, BuildHandle, DFGBuilder, Dataflow, DataflowHugr, DataflowSubContainer,
14-
SubContainer,
14+
inout_sig, BuildHandle, Container, DFGBuilder, Dataflow, DataflowHugr,
15+
DataflowSubContainer, SubContainer,
1516
},
1617
extension::{
1718
prelude::{bool_t, qb_t},
1819
simple_op::MakeRegisteredOp,
1920
},
2021
hugr::hugrmut::HugrMut,
21-
ops::{handle::ConditionalID, Tag, Value},
22-
std_extensions::logic::LogicOp,
22+
ops::{handle::ConditionalID, ExtensionOp, Tag, Value},
23+
std_extensions::{
24+
collections::array::{self, array_type, ARRAY_CLONE_OP_ID, ARRAY_DISCARD_OP_ID},
25+
logic::LogicOp,
26+
},
2327
types::{SumType, Type},
24-
Hugr, Node, Wire,
28+
Hugr, HugrView, Node, Wire,
2529
};
2630
use static_array::{ReplaceStaticArrayBoolPass, ReplaceStaticArrayBoolPassError};
2731
use tket::{
28-
extension::bool::{bool_type, BoolOp, ConstBool},
32+
extension::{
33+
bool::{bool_type, BoolOp, ConstBool},
34+
guppy::{DROP_OP_NAME, GUPPY_EXTENSION},
35+
},
2936
TketOp,
3037
};
3138

@@ -71,8 +78,7 @@ impl<H: HugrMut<Node = Node>> ComposablePass<H> for ReplaceBoolPass {
7178
type Result = ();
7279

7380
fn run(&self, hugr: &mut H) -> Result<(), Self::Error> {
74-
// TODO uncomment once https://github.com/CQCL/hugr/issues/1234 is complete
75-
// ensure_no_nonlocal_edges(hugr)?;
81+
ensure_no_nonlocal_edges(hugr)?;
7682
ReplaceStaticArrayBoolPass::default().run(hugr)?;
7783
let lowerer = lowerer();
7884
lowerer.run(hugr)?;
@@ -213,6 +219,20 @@ fn measure_reset_dest() -> NodeTemplate {
213219
NodeTemplate::CompoundOp(Box::new(h))
214220
}
215221

222+
fn array_clone_dest(size: u64, elem_ty: Type) -> NodeTemplate {
223+
let array_ty = array_type(size, elem_ty.clone());
224+
let mut dfb = DFGBuilder::new(inout_sig(
225+
vec![array_ty.clone()],
226+
vec![array_ty.clone(), array_ty],
227+
))
228+
.unwrap();
229+
let mut h = std::mem::take(dfb.hugr_mut());
230+
let [inp, outp] = h.get_io(h.entrypoint()).unwrap();
231+
h.connect(inp, 0, outp, 0);
232+
h.connect(inp, 0, outp, 1);
233+
NodeTemplate::CompoundOp(Box::new(h))
234+
}
235+
216236
/// The configuration used for replacing tket.bool extension types and ops.
217237
fn lowerer() -> ReplaceTypes {
218238
let mut lw = ReplaceTypes::default();
@@ -277,6 +297,44 @@ fn lowerer() -> ReplaceTypes {
277297
lw.replace_op(&qsystem_measure, measure_dest());
278298
lw.replace_op(&qsystem_measure_reset, measure_reset_dest());
279299

300+
// Replace array ops with copyable bounds with DFGs that the linearizer can act on
301+
// now that the elements are no longer copyable.
302+
lw.replace_parametrized_op_with(
303+
array::EXTENSION.get_op(ARRAY_CLONE_OP_ID.as_str()).unwrap(),
304+
move |args| {
305+
let [size, elem_ty] = args else {
306+
unreachable!()
307+
};
308+
let size = size.as_nat().unwrap();
309+
let elem_ty = elem_ty.as_runtime().unwrap();
310+
(!elem_ty.copyable()).then(|| array_clone_dest(size, elem_ty))
311+
},
312+
ReplacementOptions::default().with_linearization(true),
313+
);
314+
315+
lw.replace_parametrized_op(
316+
array::EXTENSION
317+
.get_op(ARRAY_DISCARD_OP_ID.as_str())
318+
.unwrap(),
319+
move |args| {
320+
let [size, elem_ty] = args else {
321+
unreachable!()
322+
};
323+
let size = size.as_nat().unwrap();
324+
let elem_ty = elem_ty.as_runtime().unwrap();
325+
if elem_ty.copyable() {
326+
return None;
327+
}
328+
let drop_op_def = GUPPY_EXTENSION.get_op(DROP_OP_NAME.as_str()).unwrap();
329+
let drop_op = ExtensionOp::new(
330+
drop_op_def.clone(),
331+
vec![array_type(size, elem_ty.clone()).into()],
332+
)
333+
.unwrap();
334+
Some(NodeTemplate::SingleOp(drop_op.into()))
335+
},
336+
);
337+
280338
lw
281339
}
282340

@@ -286,6 +344,8 @@ mod test {
286344

287345
use super::*;
288346
use hugr::ops::OpType;
347+
use hugr::std_extensions::collections::array::ArrayOpBuilder;
348+
use hugr::type_row;
289349
use hugr::{
290350
builder::{inout_sig, DFGBuilder, Dataflow, DataflowHugr},
291351
extension::prelude::qb_t,
@@ -437,4 +497,49 @@ mod test {
437497
let sig = h.signature(h.entrypoint()).unwrap();
438498
assert_eq!(sig.output(), &TypeRow::from(vec![qb_t(), bool_dest()]));
439499
}
500+
501+
#[test]
502+
fn test_array_clone_bool() {
503+
let elem_ty = bool_type();
504+
let size = 4;
505+
let arr_ty = array_type(size, elem_ty.clone());
506+
let mut dfb = DFGBuilder::new(inout_sig(
507+
vec![arr_ty.clone()],
508+
vec![arr_ty.clone(), arr_ty.clone()],
509+
))
510+
.unwrap();
511+
let [arr_in] = dfb.input_wires_arr();
512+
let (arr1, arr2) = dfb.add_array_clone(elem_ty, size, arr_in).unwrap();
513+
let mut h = dfb.finish_hugr_with_outputs([arr1, arr2]).unwrap();
514+
515+
h.validate().unwrap();
516+
let pass = ReplaceBoolPass;
517+
pass.run(&mut h).unwrap();
518+
h.validate().unwrap();
519+
520+
let sig = h.signature(h.entrypoint()).unwrap();
521+
let bool_dest_ty = bool_dest();
522+
let arr_dest_ty = array_type(size, bool_dest_ty);
523+
assert_eq!(sig.input(), &TypeRow::from(vec![arr_dest_ty.clone()]));
524+
assert_eq!(
525+
sig.output(),
526+
&TypeRow::from(vec![arr_dest_ty.clone(), arr_dest_ty])
527+
);
528+
}
529+
530+
#[test]
531+
fn test_array_discard_bool() {
532+
let elem_ty = bool_type();
533+
let size = 4;
534+
let arr_ty = array_type(size, elem_ty.clone());
535+
let mut dfb = DFGBuilder::new(inout_sig(vec![arr_ty.clone()], type_row![])).unwrap();
536+
let [arr_in] = dfb.input_wires_arr();
537+
dfb.add_array_discard(elem_ty, size, arr_in).unwrap();
538+
let mut h = dfb.finish_hugr_with_outputs([]).unwrap();
539+
540+
h.validate().unwrap();
541+
let pass = ReplaceBoolPass;
542+
pass.run(&mut h).unwrap();
543+
h.validate().unwrap();
544+
}
440545
}

0 commit comments

Comments
 (0)