Skip to main content

hugr_passes/
untuple.rs

1//! Pass for removing redundant tuple pack->unpack operations.
2
3use std::collections::VecDeque;
4
5use hugr_core::builder::{DFGBuilder, Dataflow, DataflowHugr};
6use hugr_core::extension::prelude::{MakeTuple, UnpackTuple};
7use hugr_core::hugr::SimpleReplacementError;
8use hugr_core::hugr::hugrmut::HugrMut;
9use hugr_core::hugr::views::SiblingSubgraph;
10use hugr_core::hugr::views::sibling_subgraph::TopoConvexChecker;
11use hugr_core::ops::{OpTrait, OpType};
12use hugr_core::types::Type;
13use hugr_core::{HugrView, Node, PortIndex, SimpleReplacement};
14use itertools::Itertools;
15
16use crate::composable::WithScope;
17use crate::{ComposablePass, PassScope};
18
19/// A pass that removes unnecessary `MakeTuple` operations immediately followed
20/// by `UnpackTuple`s.
21///
22/// If the tuple output is consumed by other operations, only the `UnpackTuple`s
23/// are removed and their outputs are connected to the original values
24/// accordingly.
25///
26/// Currently only unpack operations in the same region as the `MakeTuple` are
27/// removed. This may be extended in the future.
28///
29/// Removes `MakeTuple` operations that are not consumed by any other
30/// operations.
31///
32/// Ignores pack/unpack nodes with order edges.
33// TODO: Supporting those requires updating the `SiblingSubgraph` implementation. See <https://github.com/CQCL/hugr/issues/1974>.
34#[derive(Debug, Clone, Default)]
35#[deprecated(
36    note = "`hugr-passes` is deprecated. Use tket::passes instead",
37    since = "0.26.2"
38)]
39pub struct UntuplePass {
40    scope: PassScope,
41}
42
43#[derive(Debug, derive_more::Display, derive_more::Error, derive_more::From)]
44#[non_exhaustive]
45/// Errors produced by [`UntuplePass`].
46#[deprecated(
47    note = "`hugr-passes` is deprecated. Use tket::passes instead",
48    since = "0.26.2"
49)]
50pub enum UntupleError {
51    /// Rewriting the circuit failed.
52    RewriteError(SimpleReplacementError),
53}
54
55/// Result type for the untuple pass.
56#[derive(Debug, Clone, Copy, Default, PartialEq)]
57#[deprecated(
58    note = "`hugr-passes` is deprecated. Use tket::passes instead",
59    since = "0.26.2"
60)]
61pub struct UntupleResult {
62    /// Number of `MakeTuple` rewrites applied.
63    pub rewrites_applied: usize,
64}
65
66impl UntuplePass {
67    /// Find tuple pack operations followed by tuple unpack operations
68    /// and generate rewrites to remove them.
69    ///
70    /// The returned rewrites are guaranteed to be independent of each other.
71    ///
72    /// Returns an iterator over the rewrites.
73    pub fn all_rewrites<H: HugrView<Node = Node>>(
74        &self,
75        hugr: &H,
76    ) -> Vec<SimpleReplacement<H::Node>> {
77        let Some(parent) = self.scope.root(hugr) else {
78            return vec![];
79        };
80        find_rewrites(hugr, parent, self.scope.recursive())
81    }
82}
83
84fn find_rewrites<H: HugrView>(
85    hugr: &H,
86    parent: H::Node,
87    recursive: bool,
88) -> Vec<SimpleReplacement<H::Node>> {
89    let mut res = Vec::new();
90    let mut children_queue = VecDeque::new();
91    children_queue.push_back(parent);
92
93    while let Some(parent) = children_queue.pop_front() {
94        // Required to create SimpleReplacements.
95        // Reset for each parent as `TopoConvexChecker` is tied to a specific parent node.
96        let mut convex_checker: Option<TopoConvexChecker<H>> = None;
97
98        for node in hugr.children(parent) {
99            let op = hugr.get_optype(node);
100            if let Some(rw) = make_rewrite(hugr, &mut convex_checker, node, op) {
101                res.push(rw);
102            }
103            if recursive && op.is_container() {
104                children_queue.push_back(node);
105            }
106        }
107    }
108    res
109}
110
111impl<H: HugrMut<Node = Node>> ComposablePass<H> for UntuplePass {
112    type Error = UntupleError;
113    type Result = UntupleResult;
114
115    fn run(&self, hugr: &mut H) -> Result<Self::Result, Self::Error> {
116        let rewrites = self.all_rewrites(hugr);
117        let rewrites_applied = rewrites.len();
118        // The rewrites are independent, so we can always apply them all.
119        for rewrite in rewrites {
120            hugr.apply_patch(rewrite)?;
121        }
122        Ok(UntupleResult { rewrites_applied })
123    }
124}
125
126impl WithScope for UntuplePass {
127    fn with_scope(mut self, scope: impl Into<PassScope>) -> Self {
128        self.scope = scope.into();
129        self
130    }
131}
132
133/// Returns true if the given optype is a `MakeTuple` operation.
134///
135/// Boilerplate required due to <https://github.com/CQCL/hugr/issues/1496>
136fn is_make_tuple(optype: &OpType) -> bool {
137    optype.cast::<MakeTuple>().is_some()
138}
139
140/// Returns true if the given optype is an `UnpackTuple` operation.
141///
142/// Boilerplate required due to <https://github.com/CQCL/hugr/issues/1496>
143fn is_unpack_tuple(optype: &OpType) -> bool {
144    optype.cast::<UnpackTuple>().is_some()
145}
146
147/// If this is a `MakeTuple` operation followed by some number of `UnpackTuple` operations
148/// on the same region, return a rewrite to remove them.
149///
150/// Otherwise, return None.
151fn make_rewrite<'h, T: HugrView>(
152    hugr: &'h T,
153    convex_checker: &mut Option<TopoConvexChecker<'h, T>>,
154    node: T::Node,
155    op: &OpType,
156) -> Option<SimpleReplacement<T::Node>> {
157    // Only process MakeTuple operations
158    if !is_make_tuple(op) {
159        return None;
160    }
161
162    let has_order_edges = |node: T::Node| -> bool {
163        let op = hugr.get_optype(node);
164        let has_input_order = op
165            .other_input_port()
166            .and_then(|p| hugr.linked_outputs(node, p).next())
167            .is_some();
168        let has_output_order = op
169            .other_output_port()
170            .and_then(|p| hugr.linked_inputs(node, p).next())
171            .is_some();
172        has_input_order || has_output_order
173    };
174
175    // If the node has order edges, ignore it.
176    if has_order_edges(node) {
177        return None;
178    }
179
180    let tuple_types = op.dataflow_signature().unwrap().input_types().to_vec();
181    let node_parent = hugr.get_parent(node);
182
183    // See if it is followed by a tuple unpack
184    let links = hugr
185        .linked_inputs(node, 0)
186        .map(|(neigh, _)| neigh)
187        .collect_vec();
188
189    let unpack_nodes = links
190        .iter()
191        .filter(|&&neigh| hugr.get_parent(neigh) == node_parent)
192        .filter(|&&neigh| is_unpack_tuple(hugr.get_optype(neigh)))
193        .filter(|&&neigh| !has_order_edges(neigh))
194        .copied()
195        .collect_vec();
196
197    // If there are no unpacks but the tuple is being used, there's nothing to do.
198    if unpack_nodes.is_empty() && !links.is_empty() {
199        return None;
200    }
201
202    // Remove all unpack operations, and remove the pack operation if all neighbours are unpacks.
203    let num_other_outputs = links.len() - unpack_nodes.len();
204    Some(remove_pack_unpack(
205        hugr,
206        convex_checker,
207        &tuple_types,
208        node,
209        unpack_nodes,
210        num_other_outputs,
211    ))
212}
213
214/// Returns a rewrite to remove a tuple pack operation that's followed by unpack operations,
215/// and `other_tuple_links` other operations.
216fn remove_pack_unpack<'h, T: HugrView>(
217    hugr: &'h T,
218    convex_checker: &mut Option<TopoConvexChecker<'h, T>>,
219    tuple_types: &[Type],
220    pack_node: T::Node,
221    unpack_nodes: Vec<T::Node>,
222    num_other_outputs: usize,
223) -> SimpleReplacement<T::Node> {
224    let parent = hugr.get_parent(pack_node).expect("pack_node has no parent");
225    let checker = convex_checker.get_or_insert_with(|| TopoConvexChecker::new(hugr, parent));
226
227    let mut nodes = unpack_nodes.clone();
228    nodes.push(pack_node);
229    let subcirc = SiblingSubgraph::try_from_nodes_with_checker(nodes, hugr, checker).unwrap();
230    let subcirc_signature = subcirc.signature(hugr);
231
232    let mut replacement = DFGBuilder::new(subcirc_signature).unwrap();
233
234    // Wire the inputs directly to the unpack outputs
235    // We need to list the **connected** output ports from the unpack nodes.
236    // SiblingSubgraph ignores disconnected outputs, so we need these when building the replacement.
237    let mut replacement_outputs =
238        Vec::with_capacity(unpack_nodes.len() * tuple_types.len() + num_other_outputs);
239    let replacement_inputs = replacement.input_wires().collect_vec();
240    for unpack_node in unpack_nodes {
241        for out_port in hugr.node_outputs(unpack_node) {
242            if hugr.is_linked(unpack_node, out_port) {
243                let input = replacement_inputs[out_port.index()];
244                replacement_outputs.push(input);
245            }
246        }
247    }
248
249    // If needed, re-add the tuple pack node and connect its output to the tuple outputs.
250    if num_other_outputs > 0 {
251        let op = MakeTuple::new(tuple_types.to_vec().into());
252        let [tuple] = replacement
253            .add_dataflow_op(op, replacement.input_wires())
254            .unwrap()
255            .outputs_arr();
256        replacement_outputs.extend(std::iter::repeat_n(tuple, num_other_outputs));
257    }
258
259    // These should never fail, as we are defining the replacement ourselves.
260    let replacement = replacement
261        .finish_hugr_with_outputs(replacement_outputs)
262        .unwrap_or_else(|e| {
263            panic!("Failed to create replacement for removing tuple pack/unpack operations. {e}")
264        });
265    subcirc
266        .create_simple_replacement(hugr, replacement)
267        .unwrap_or_else(|e| {
268            panic!("Failed to create rewrite for removing tuple pack/unpack operations. {e}")
269        })
270}
271
272#[cfg(test)]
273mod test {
274    use super::*;
275    use crate::composable::WithScope;
276    use hugr_core::Hugr;
277    use hugr_core::builder::FunctionBuilder;
278    use hugr_core::extension::prelude::{UnpackTuple, bool_t, qb_t};
279    use hugr_core::ops::handle::NodeHandle;
280    use hugr_core::std_extensions::arithmetic::float_types::float64_type;
281    use hugr_core::types::Signature;
282    use rstest::{fixture, rstest};
283
284    /// A simple pack operation with unused output.
285    ///
286    /// These can be removed entirely.
287    #[fixture]
288    fn unused_pack() -> Hugr {
289        let mut h = DFGBuilder::new(Signature::new(vec![bool_t(), bool_t()], vec![])).unwrap();
290        let mut inps = h.input_wires();
291        let b1 = inps.next().unwrap();
292        let b2 = inps.next().unwrap();
293
294        let _tuple = h.make_tuple([b1, b2]).unwrap();
295
296        h.finish_hugr_with_outputs([]).unwrap()
297    }
298
299    /// A simple pack operation followed by an unpack operation.
300    ///
301    /// These can be removed entirely.
302    #[fixture]
303    fn simple_pack_unpack() -> Hugr {
304        let mut h = DFGBuilder::new(Signature::new_endo([qb_t(), bool_t()])).unwrap();
305        let mut inps = h.input_wires();
306        let qb1 = inps.next().unwrap();
307        let b2 = inps.next().unwrap();
308
309        let tuple = h.make_tuple([qb1, b2]).unwrap();
310
311        let op = UnpackTuple::new(vec![qb_t(), bool_t()].into());
312        let [qb1, b2] = h.add_dataflow_op(op, [tuple]).unwrap().outputs_arr();
313
314        h.finish_hugr_with_outputs([qb1, b2]).unwrap()
315    }
316
317    /// A simple pack/unpack pair with order edges between them.
318    ///
319    /// In the future we should be able to preserve some order edges, but for now
320    /// we just remove everything.
321    #[fixture]
322    fn ordered_pack_unpack() -> Hugr {
323        let mut h = DFGBuilder::new(Signature::new_endo(vec![qb_t(), bool_t()])).unwrap();
324        let mut inps = h.input_wires();
325        let qb1 = inps.next().unwrap();
326        let b2 = inps.next().unwrap();
327
328        let tuple = h.make_tuple([qb1, b2]).unwrap();
329        h.set_order(&h.input(), &tuple.node());
330
331        let op = UnpackTuple::new(vec![qb_t(), bool_t()].into());
332        let untuple = h.add_dataflow_op(op, [tuple]).unwrap();
333        let [qb1, b2] = untuple.outputs_arr();
334        h.set_order(&tuple.node(), &untuple.node());
335
336        h.set_order(&untuple.node(), &h.output());
337        h.finish_hugr_with_outputs([qb1, b2]).unwrap()
338    }
339
340    /// A simple pack/unpack pair with an order from the pack node to a downstream node.
341    ///
342    /// The order edge should be preserved, so we move it to the predecessor of the pack node.
343    #[fixture]
344    fn outgoing_ordered_pack_unpack() -> Hugr {
345        let mut h = DFGBuilder::new(Signature::new_endo(vec![qb_t(), bool_t()])).unwrap();
346        let mut inps = h.input_wires();
347        let qb1 = inps.next().unwrap();
348        let b2 = inps.next().unwrap();
349
350        let tuple = h.make_tuple([qb1, b2]).unwrap();
351
352        let op = UnpackTuple::new(vec![qb_t(), bool_t()].into());
353        let untuple = h.add_dataflow_op(op, [tuple]).unwrap();
354        let [qb1, b2] = untuple.outputs_arr();
355
356        h.set_order(&tuple.node(), &h.output());
357        h.finish_hugr_with_outputs([qb1, b2]).unwrap()
358    }
359
360    /// A simple pack/unpack pair with an order from a downstream node to the pack node.
361    ///
362    /// The order edge should be preserved, so we move it to the successor of the unpack node.
363    #[fixture]
364    fn incoming_ordered_pack_unpack() -> Hugr {
365        let mut h = DFGBuilder::new(Signature::new_endo(vec![qb_t(), bool_t()])).unwrap();
366        let mut inps = h.input_wires();
367        let qb1 = inps.next().unwrap();
368        let b2 = inps.next().unwrap();
369
370        let tuple = h.make_tuple([qb1, b2]).unwrap();
371
372        let op = UnpackTuple::new(vec![qb_t(), bool_t()].into());
373        let untuple = h.add_dataflow_op(op, [tuple]).unwrap();
374        let [qb1, b2] = untuple.outputs_arr();
375
376        h.set_order(&h.input(), &untuple.node());
377        h.finish_hugr_with_outputs([qb1, b2]).unwrap()
378    }
379
380    /// A pack operation followed by three unpack operations from the same tuple.
381    ///
382    /// These can be removed entirely.
383    #[fixture]
384    fn multi_unpack() -> Hugr {
385        let mut h = DFGBuilder::new(Signature::new(
386            vec![bool_t(), bool_t()],
387            vec![bool_t(), bool_t(), bool_t(), bool_t()],
388        ))
389        .unwrap();
390        let mut inps = h.input_wires();
391        let b1 = inps.next().unwrap();
392        let b2 = inps.next().unwrap();
393
394        let tuple = h.make_tuple([b1, b2]).unwrap();
395
396        let op = UnpackTuple::new(vec![bool_t(), bool_t()].into());
397        let [b1, b2] = h.add_dataflow_op(op, [tuple]).unwrap().outputs_arr();
398
399        let op = UnpackTuple::new(vec![bool_t(), bool_t()].into());
400        let [b3, b4] = h.add_dataflow_op(op, [tuple]).unwrap().outputs_arr();
401
402        // The last one's outputs are disconnected.
403        // TODO: Adding this causes the test to fail due to a `NonCovex` error.
404        //let op = UnpackTuple::new(vec![bool_t(), bool_t()].into());
405        //let _ = h.add_dataflow_op(op, [tuple]).unwrap();
406
407        h.finish_hugr_with_outputs([b1, b2, b3, b4]).unwrap()
408    }
409
410    /// A pack operation followed by an unpack operation, where the tuple is also returned.
411    ///
412    /// The unpack operation can be removed, but the pack operation cannot.
413    #[fixture]
414    fn partial_unpack() -> Hugr {
415        let mut h = DFGBuilder::new(Signature::new(
416            vec![bool_t(), bool_t()],
417            vec![
418                bool_t(),
419                bool_t(),
420                Type::new_tuple(vec![bool_t(), bool_t()]),
421            ],
422        ))
423        .unwrap();
424        let mut inps = h.input_wires();
425        let b1 = inps.next().unwrap();
426        let b2 = inps.next().unwrap();
427
428        let tuple = h.make_tuple([b1, b2]).unwrap();
429
430        let op = UnpackTuple::new(vec![bool_t(), bool_t()].into());
431        let [b1, b2] = h.add_dataflow_op(op, [tuple]).unwrap().outputs_arr();
432
433        h.finish_hugr_with_outputs([b1, b2, tuple]).unwrap()
434    }
435
436    /// A pack operation followed by an unpack that discards its first output.
437    ///
438    /// The unpack operation can be removed, but the pack operation cannot.
439    ///
440    /// This is a minimal error case for <https://github.com/Quantinuum/tket2/issues/1347>.
441    #[fixture]
442    fn unpack_discard_first() -> Hugr {
443        let mut h = FunctionBuilder::new(
444            "test",
445            Signature::new(vec![bool_t(), float64_type()], vec![float64_type()]),
446        )
447        .unwrap();
448        let [b, f] = h.input_wires_arr();
449
450        let tuple = h.make_tuple([b, f]).unwrap();
451
452        let op = UnpackTuple::new(vec![bool_t(), float64_type()].into());
453        let [_b, f] = h.add_dataflow_op(op, [tuple]).unwrap().outputs_arr();
454
455        h.finish_hugr_with_outputs([f]).unwrap()
456    }
457
458    #[rstest]
459    #[case::unused(unused_pack(), 1, 2)]
460    #[case::simple(simple_pack_unpack(), 1, 2)]
461    #[case::multi(multi_unpack(), 1, 2)]
462    #[case::partial(partial_unpack(), 1, 3)]
463    #[case::unpack_discard_first(unpack_discard_first(), 1, 2)]
464    // Nodes with order edges are ignored.
465    #[case::ordered(ordered_pack_unpack(), 0, 4)]
466    #[case::outgoing_ordered(outgoing_ordered_pack_unpack(), 0, 4)]
467    #[case::incoming_ordered(incoming_ordered_pack_unpack(), 0, 4)]
468    fn test_pack_unpack(
469        #[case] mut hugr: Hugr,
470        #[case] expected_rewrites: usize,
471        #[case] remaining_nodes: usize,
472    ) {
473        let parent = hugr.entrypoint();
474        let pass = UntuplePass::default().with_scope(PassScope::EntrypointFlat);
475        let res = pass.run(&mut hugr).unwrap_or_else(|e| panic!("{e}"));
476        assert_eq!(res.rewrites_applied, expected_rewrites);
477        assert_eq!(hugr.children(parent).count(), remaining_nodes);
478    }
479}