1use 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#[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#[deprecated(
47 note = "`hugr-passes` is deprecated. Use tket::passes instead",
48 since = "0.26.2"
49)]
50pub enum UntupleError {
51 RewriteError(SimpleReplacementError),
53}
54
55#[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 pub rewrites_applied: usize,
64}
65
66impl UntuplePass {
67 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 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 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
133fn is_make_tuple(optype: &OpType) -> bool {
137 optype.cast::<MakeTuple>().is_some()
138}
139
140fn is_unpack_tuple(optype: &OpType) -> bool {
144 optype.cast::<UnpackTuple>().is_some()
145}
146
147fn 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 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 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 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 unpack_nodes.is_empty() && !links.is_empty() {
199 return None;
200 }
201
202 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
214fn 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 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 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 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 #[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 #[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 #[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 #[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 #[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 #[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 h.finish_hugr_with_outputs([b1, b2, b3, b4]).unwrap()
408 }
409
410 #[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 #[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 #[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}