@@ -5,27 +5,34 @@ mod static_array;
55use derive_more:: { Display , Error , From } ;
66use 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} ;
2630use static_array:: { ReplaceStaticArrayBoolPass , ReplaceStaticArrayBoolPassError } ;
2731use 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.
217237fn 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