1use std::collections::HashMap;
12
13use nxpu_ir::graph::{ActivationFunction, ComputeGraph, EdgeId, GraphOp};
14
15use crate::Pass;
16
17#[derive(Debug)]
19pub struct OperatorFusion;
20
21impl Pass for OperatorFusion {
22 fn name(&self) -> &str {
23 "operator-fusion"
24 }
25
26 fn run(&self, module: &mut nxpu_ir::Module) -> bool {
27 let _ = module;
30 false
31 }
32}
33
34impl OperatorFusion {
35 pub fn run_on_graph(&self, graph: &mut ComputeGraph) -> bool {
38 let mut changed = false;
39
40 loop {
42 let fused_any = try_fuse_conv_bias_activation(graph)
43 | try_fuse_matmul_bias(graph)
44 | try_fuse_elementwise_activation(graph);
45
46 if fused_any {
47 changed = true;
48 } else {
49 break;
50 }
51 }
52
53 changed
54 }
55}
56
57fn is_activation(op: &GraphOp) -> Option<ActivationFunction> {
59 match op {
60 GraphOp::Relu => Some(ActivationFunction::Relu),
61 GraphOp::Sigmoid => Some(ActivationFunction::Sigmoid),
62 _ => None,
63 }
64}
65
66fn is_elementwise_binary(op: &GraphOp) -> bool {
68 matches!(
69 op,
70 GraphOp::Add | GraphOp::Sub | GraphOp::Mul | GraphOp::Div
71 )
72}
73
74fn build_edge_consumer_map(graph: &ComputeGraph) -> HashMap<EdgeId, Vec<usize>> {
76 let mut map: HashMap<EdgeId, Vec<usize>> = HashMap::new();
77 for (i, node) in graph.nodes.iter().enumerate() {
78 for &inp in &node.inputs {
79 map.entry(inp).or_default().push(i);
80 }
81 }
82 map
83}
84
85fn has_single_consumer(edge: EdgeId, consumer_map: &HashMap<EdgeId, Vec<usize>>) -> bool {
87 match consumer_map.get(&edge) {
88 Some(consumers) => consumers.len() == 1,
89 None => false,
90 }
91}
92
93fn remove_node_and_rewire(
98 graph: &mut ComputeGraph,
99 remove_idx: usize,
100 fused_idx: usize,
101 intermediate_edge: EdgeId,
102) {
103 let new_outputs = graph.nodes[remove_idx].outputs.clone();
105 graph.nodes[fused_idx].outputs = new_outputs;
106
107 graph.inputs.retain(|e| *e != intermediate_edge);
109 graph.outputs.retain(|e| *e != intermediate_edge);
110
111 graph.nodes.remove(remove_idx);
113
114 graph.edges.remove(&intermediate_edge);
116}
117
118fn try_fuse_conv_bias_activation(graph: &mut ComputeGraph) -> bool {
123 let edge_consumer = build_edge_consumer_map(graph);
124
125 for conv_idx in 0..graph.nodes.len() {
127 if graph.nodes[conv_idx].op != GraphOp::Conv2d {
128 continue;
129 }
130
131 if graph.nodes[conv_idx].outputs.len() != 1 {
133 continue;
134 }
135 let conv_out_edge = graph.nodes[conv_idx].outputs[0];
136
137 if !has_single_consumer(conv_out_edge, &edge_consumer) {
139 continue;
140 }
141
142 let consumer_indices = &edge_consumer[&conv_out_edge];
143 let next_idx = consumer_indices[0];
144
145 match &graph.nodes[next_idx].op {
147 GraphOp::Add => {
148 let add_idx = next_idx;
150 if graph.nodes[add_idx].outputs.len() != 1 {
151 continue;
152 }
153 let add_out_edge = graph.nodes[add_idx].outputs[0];
154
155 let bias_edge = graph.nodes[add_idx]
157 .inputs
158 .iter()
159 .find(|&&e| e != conv_out_edge)
160 .copied();
161
162 let Some(bias_edge) = bias_edge else {
163 continue;
164 };
165
166 let mut fused_inputs = graph.nodes[conv_idx].inputs.clone();
168 fused_inputs.push(bias_edge);
169
170 let activation = if has_single_consumer(add_out_edge, &edge_consumer) {
172 let act_consumers = &edge_consumer[&add_out_edge];
173 let act_idx = act_consumers[0];
174 is_activation(&graph.nodes[act_idx].op)
175 } else {
176 None
177 };
178
179 let act = activation.unwrap_or(ActivationFunction::None);
180
181 graph.nodes[conv_idx].op = GraphOp::FusedConv2d { activation: act };
183 graph.nodes[conv_idx].inputs = fused_inputs;
184 graph.nodes[conv_idx].name = format!("{}_fused", graph.nodes[conv_idx].name);
185
186 remove_node_and_rewire(graph, add_idx, conv_idx, conv_out_edge);
187
188 if activation.is_some() {
190 remove_activation_after_fused_conv(graph, add_out_edge, conv_idx);
191 }
192
193 return true;
194 }
195 op if is_activation(op).is_some() => {
196 let act = is_activation(op).unwrap();
198 let act_idx = next_idx;
199
200 graph.nodes[conv_idx].op = GraphOp::FusedConv2d { activation: act };
201 graph.nodes[conv_idx].name = format!("{}_fused", graph.nodes[conv_idx].name);
202
203 remove_node_and_rewire(graph, act_idx, conv_idx, conv_out_edge);
204
205 return true;
206 }
207 _ => {}
208 }
209 }
210
211 false
212}
213
214fn remove_activation_after_fused_conv(
216 graph: &mut ComputeGraph,
217 add_out_edge: EdgeId,
218 conv_idx: usize,
219) {
220 let edge_consumer_new = build_edge_consumer_map(graph);
221 let act_consumers = match edge_consumer_new.get(&add_out_edge) {
222 Some(c) if c.len() == 1 => c,
223 _ => return,
224 };
225
226 let act_idx = act_consumers[0];
227 let fused_idx = graph
228 .nodes
229 .iter()
230 .position(|n| {
231 matches!(n.op, GraphOp::FusedConv2d { .. }) && n.outputs.contains(&add_out_edge)
232 })
233 .unwrap_or(conv_idx.min(graph.nodes.len() - 1));
234
235 remove_node_and_rewire(graph, act_idx, fused_idx, add_out_edge);
236}
237
238fn try_fuse_matmul_bias(graph: &mut ComputeGraph) -> bool {
243 let edge_consumer = build_edge_consumer_map(graph);
244
245 for mm_idx in 0..graph.nodes.len() {
246 if graph.nodes[mm_idx].op != GraphOp::MatMul {
247 continue;
248 }
249
250 if graph.nodes[mm_idx].outputs.len() != 1 {
251 continue;
252 }
253 let mm_out_edge = graph.nodes[mm_idx].outputs[0];
254
255 if !has_single_consumer(mm_out_edge, &edge_consumer) {
256 continue;
257 }
258
259 let consumer_indices = &edge_consumer[&mm_out_edge];
260 let add_idx = consumer_indices[0];
261
262 if graph.nodes[add_idx].op != GraphOp::Add {
263 continue;
264 }
265
266 let bias_edge = graph.nodes[add_idx]
268 .inputs
269 .iter()
270 .find(|&&e| e != mm_out_edge)
271 .copied();
272
273 let Some(bias_edge) = bias_edge else {
274 continue;
275 };
276
277 let mut gemm_inputs = graph.nodes[mm_idx].inputs.clone();
279 gemm_inputs.push(bias_edge);
280
281 graph.nodes[mm_idx].op = GraphOp::Gemm { alpha: 1, beta: 1 };
282 graph.nodes[mm_idx].inputs = gemm_inputs;
283 graph.nodes[mm_idx].name = format!("{}_fused", graph.nodes[mm_idx].name);
284
285 remove_node_and_rewire(graph, add_idx, mm_idx, mm_out_edge);
286
287 return true;
288 }
289
290 false
291}
292
293fn try_fuse_elementwise_activation(graph: &mut ComputeGraph) -> bool {
298 let edge_consumer = build_edge_consumer_map(graph);
299
300 for ew_idx in 0..graph.nodes.len() {
301 if !is_elementwise_binary(&graph.nodes[ew_idx].op) {
302 continue;
303 }
304
305 if graph.nodes[ew_idx].outputs.len() != 1 {
306 continue;
307 }
308 let ew_out_edge = graph.nodes[ew_idx].outputs[0];
309
310 if !has_single_consumer(ew_out_edge, &edge_consumer) {
311 continue;
312 }
313
314 let consumer_indices = &edge_consumer[&ew_out_edge];
315 let act_idx = consumer_indices[0];
316
317 let act = match is_activation(&graph.nodes[act_idx].op) {
318 Some(a) => a,
319 None => continue,
320 };
321
322 let base_op = graph.nodes[ew_idx].op.clone();
324 graph.nodes[ew_idx].op = GraphOp::FusedElementWise {
325 base_op: Box::new(base_op),
326 activation: act,
327 };
328 graph.nodes[ew_idx].name = format!("{}_fused", graph.nodes[ew_idx].name);
329
330 remove_node_and_rewire(graph, act_idx, ew_idx, ew_out_edge);
331
332 return true;
333 }
334
335 false
336}
337
338#[cfg(test)]
339mod tests {
340 use super::*;
341 use nxpu_ir::graph::{ComputeGraph, GraphOp, TensorInfo};
342 use nxpu_ir::{Dimension, Scalar, TensorShape};
343
344 fn make_tensor(name: &str, shape: &[i64]) -> TensorInfo {
345 TensorInfo {
346 name: name.into(),
347 scalar: Scalar::F32,
348 shape: TensorShape {
349 dims: shape
350 .iter()
351 .map(|&d| {
352 if d < 0 {
353 Dimension::Dynamic(None)
354 } else {
355 Dimension::Fixed(d as u32)
356 }
357 })
358 .collect(),
359 },
360 }
361 }
362
363 #[test]
364 fn fuse_conv_relu() {
365 let mut graph = ComputeGraph::new();
366
367 let input = graph.add_edge(make_tensor("input", &[-1, 3, 224, 224]));
368 let weight = graph.add_edge(make_tensor("weight", &[64, 3, 3, 3]));
369 let conv_out = graph.add_edge(make_tensor("conv_out", &[-1, 64, 222, 222]));
370 let relu_out = graph.add_edge(make_tensor("relu_out", &[-1, 64, 222, 222]));
371
372 graph.inputs = vec![input, weight];
373 graph.outputs = vec![relu_out];
374
375 graph
376 .add_node(GraphOp::Conv2d, vec![input, weight], vec![conv_out], "conv")
377 .unwrap();
378 graph
379 .add_node(GraphOp::Relu, vec![conv_out], vec![relu_out], "relu")
380 .unwrap();
381
382 assert_eq!(graph.node_count(), 2);
383
384 let fusion = OperatorFusion;
385 let changed = fusion.run_on_graph(&mut graph);
386
387 assert!(changed);
388 assert_eq!(graph.node_count(), 1);
389 assert_eq!(
390 graph.nodes[0].op,
391 GraphOp::FusedConv2d {
392 activation: ActivationFunction::Relu
393 }
394 );
395 assert_eq!(graph.nodes[0].outputs, vec![relu_out]);
397 }
398
399 #[test]
400 fn fuse_conv_bias_relu() {
401 let mut graph = ComputeGraph::new();
402
403 let input = graph.add_edge(make_tensor("input", &[-1, 3, 224, 224]));
404 let weight = graph.add_edge(make_tensor("weight", &[64, 3, 3, 3]));
405 let bias = graph.add_edge(make_tensor("bias", &[64]));
406 let conv_out = graph.add_edge(make_tensor("conv_out", &[-1, 64, 222, 222]));
407 let add_out = graph.add_edge(make_tensor("add_out", &[-1, 64, 222, 222]));
408 let relu_out = graph.add_edge(make_tensor("relu_out", &[-1, 64, 222, 222]));
409
410 graph.inputs = vec![input, weight, bias];
411 graph.outputs = vec![relu_out];
412
413 graph
414 .add_node(GraphOp::Conv2d, vec![input, weight], vec![conv_out], "conv")
415 .unwrap();
416 graph
417 .add_node(
418 GraphOp::Add,
419 vec![conv_out, bias],
420 vec![add_out],
421 "bias_add",
422 )
423 .unwrap();
424 graph
425 .add_node(GraphOp::Relu, vec![add_out], vec![relu_out], "relu")
426 .unwrap();
427
428 assert_eq!(graph.node_count(), 3);
429
430 let fusion = OperatorFusion;
431 let changed = fusion.run_on_graph(&mut graph);
432
433 assert!(changed);
434 assert_eq!(graph.node_count(), 1);
435 assert_eq!(
436 graph.nodes[0].op,
437 GraphOp::FusedConv2d {
438 activation: ActivationFunction::Relu
439 }
440 );
441 assert_eq!(graph.nodes[0].inputs, vec![input, weight, bias]);
443 assert_eq!(graph.nodes[0].outputs, vec![relu_out]);
444 }
445
446 #[test]
447 fn fuse_matmul_add_to_gemm() {
448 let mut graph = ComputeGraph::new();
449
450 let a = graph.add_edge(make_tensor("A", &[-1, 768]));
451 let b = graph.add_edge(make_tensor("B", &[768, 768]));
452 let mm_out = graph.add_edge(make_tensor("mm_out", &[-1, 768]));
453 let bias = graph.add_edge(make_tensor("bias", &[768]));
454 let add_out = graph.add_edge(make_tensor("add_out", &[-1, 768]));
455
456 graph.inputs = vec![a, b, bias];
457 graph.outputs = vec![add_out];
458
459 graph
460 .add_node(GraphOp::MatMul, vec![a, b], vec![mm_out], "matmul")
461 .unwrap();
462 graph
463 .add_node(GraphOp::Add, vec![mm_out, bias], vec![add_out], "add")
464 .unwrap();
465
466 assert_eq!(graph.node_count(), 2);
467
468 let fusion = OperatorFusion;
469 let changed = fusion.run_on_graph(&mut graph);
470
471 assert!(changed);
472 assert_eq!(graph.node_count(), 1);
473 assert_eq!(graph.nodes[0].op, GraphOp::Gemm { alpha: 1, beta: 1 });
474 assert_eq!(graph.nodes[0].inputs, vec![a, b, bias]);
475 assert_eq!(graph.nodes[0].outputs, vec![add_out]);
476 }
477
478 #[test]
479 fn fuse_add_relu() {
480 let mut graph = ComputeGraph::new();
481
482 let a = graph.add_edge(make_tensor("a", &[-1, 256]));
483 let b = graph.add_edge(make_tensor("b", &[-1, 256]));
484 let add_out = graph.add_edge(make_tensor("add_out", &[-1, 256]));
485 let relu_out = graph.add_edge(make_tensor("relu_out", &[-1, 256]));
486
487 graph.inputs = vec![a, b];
488 graph.outputs = vec![relu_out];
489
490 graph
491 .add_node(GraphOp::Add, vec![a, b], vec![add_out], "add")
492 .unwrap();
493 graph
494 .add_node(GraphOp::Relu, vec![add_out], vec![relu_out], "relu")
495 .unwrap();
496
497 assert_eq!(graph.node_count(), 2);
498
499 let fusion = OperatorFusion;
500 let changed = fusion.run_on_graph(&mut graph);
501
502 assert!(changed);
503 assert_eq!(graph.node_count(), 1);
504 assert_eq!(
505 graph.nodes[0].op,
506 GraphOp::FusedElementWise {
507 base_op: Box::new(GraphOp::Add),
508 activation: ActivationFunction::Relu,
509 }
510 );
511 assert_eq!(graph.nodes[0].inputs, vec![a, b]);
512 assert_eq!(graph.nodes[0].outputs, vec![relu_out]);
513 }
514
515 #[test]
516 fn no_fusion_unfusible_pattern() {
517 let mut graph = ComputeGraph::new();
519
520 let a = graph.add_edge(make_tensor("A", &[-1, 768]));
521 let b = graph.add_edge(make_tensor("B", &[768, 768]));
522 let mm_out = graph.add_edge(make_tensor("mm_out", &[-1, 768]));
523 let reshape_out = graph.add_edge(make_tensor("reshape_out", &[-1, 768]));
524
525 graph.inputs = vec![a, b];
526 graph.outputs = vec![reshape_out];
527
528 graph
529 .add_node(GraphOp::MatMul, vec![a, b], vec![mm_out], "matmul")
530 .unwrap();
531 graph
532 .add_node(GraphOp::Reshape, vec![mm_out], vec![reshape_out], "reshape")
533 .unwrap();
534
535 let fusion = OperatorFusion;
536 let changed = fusion.run_on_graph(&mut graph);
537
538 assert!(!changed);
539 assert_eq!(graph.node_count(), 2);
540 }
541
542 #[test]
543 fn no_fusion_multi_consumer() {
544 let mut graph = ComputeGraph::new();
546
547 let a = graph.add_edge(make_tensor("A", &[-1, 768]));
548 let b = graph.add_edge(make_tensor("B", &[768, 768]));
549 let mm_out = graph.add_edge(make_tensor("mm_out", &[-1, 768]));
550 let bias = graph.add_edge(make_tensor("bias", &[768]));
551 let add_out = graph.add_edge(make_tensor("add_out", &[-1, 768]));
552 let relu_out = graph.add_edge(make_tensor("relu_out", &[-1, 768]));
553
554 graph.inputs = vec![a, b, bias];
555 graph.outputs = vec![add_out, relu_out];
556
557 graph
558 .add_node(GraphOp::MatMul, vec![a, b], vec![mm_out], "matmul")
559 .unwrap();
560 graph
562 .add_node(GraphOp::Add, vec![mm_out, bias], vec![add_out], "add")
563 .unwrap();
564 graph
565 .add_node(GraphOp::Relu, vec![mm_out], vec![relu_out], "relu")
566 .unwrap();
567
568 let fusion = OperatorFusion;
569 let changed = fusion.run_on_graph(&mut graph);
570
571 assert!(!changed);
573 assert_eq!(graph.node_count(), 3);
574 }
575
576 #[test]
577 fn no_fusion_standalone_activation() {
578 let mut graph = ComputeGraph::new();
580
581 let input = graph.add_edge(make_tensor("input", &[-1, 256]));
582 let output = graph.add_edge(make_tensor("output", &[-1, 256]));
583
584 graph.inputs = vec![input];
585 graph.outputs = vec![output];
586
587 graph
588 .add_node(GraphOp::Relu, vec![input], vec![output], "relu")
589 .unwrap();
590
591 let fusion = OperatorFusion;
592 let changed = fusion.run_on_graph(&mut graph);
593
594 assert!(!changed);
595 assert_eq!(graph.node_count(), 1);
596 }
597
598 #[test]
599 fn fuse_mul_sigmoid() {
600 let mut graph = ComputeGraph::new();
601
602 let a = graph.add_edge(make_tensor("a", &[-1, 256]));
603 let b = graph.add_edge(make_tensor("b", &[-1, 256]));
604 let mul_out = graph.add_edge(make_tensor("mul_out", &[-1, 256]));
605 let sig_out = graph.add_edge(make_tensor("sig_out", &[-1, 256]));
606
607 graph.inputs = vec![a, b];
608 graph.outputs = vec![sig_out];
609
610 graph
611 .add_node(GraphOp::Mul, vec![a, b], vec![mul_out], "mul")
612 .unwrap();
613 graph
614 .add_node(GraphOp::Sigmoid, vec![mul_out], vec![sig_out], "sigmoid")
615 .unwrap();
616
617 let fusion = OperatorFusion;
618 let changed = fusion.run_on_graph(&mut graph);
619
620 assert!(changed);
621 assert_eq!(graph.node_count(), 1);
622 assert_eq!(
623 graph.nodes[0].op,
624 GraphOp::FusedElementWise {
625 base_op: Box::new(GraphOp::Mul),
626 activation: ActivationFunction::Sigmoid,
627 }
628 );
629 }
630
631 #[test]
632 fn chained_fusion() {
633 let mut graph = ComputeGraph::new();
636
637 let a = graph.add_edge(make_tensor("A", &[-1, 768]));
638 let b = graph.add_edge(make_tensor("B", &[768, 768]));
639 let mm_out = graph.add_edge(make_tensor("mm_out", &[-1, 768]));
640 let bias = graph.add_edge(make_tensor("bias", &[768]));
641 let add_out = graph.add_edge(make_tensor("add_out", &[-1, 768]));
642 let relu_out = graph.add_edge(make_tensor("relu_out", &[-1, 768]));
643
644 graph.inputs = vec![a, b, bias];
645 graph.outputs = vec![relu_out];
646
647 graph
648 .add_node(GraphOp::MatMul, vec![a, b], vec![mm_out], "matmul")
649 .unwrap();
650 graph
651 .add_node(GraphOp::Add, vec![mm_out, bias], vec![add_out], "add")
652 .unwrap();
653 graph
654 .add_node(GraphOp::Relu, vec![add_out], vec![relu_out], "relu")
655 .unwrap();
656
657 let fusion = OperatorFusion;
658 let changed = fusion.run_on_graph(&mut graph);
659
660 assert!(changed);
661 assert_eq!(graph.node_count(), 2);
663 assert!(matches!(
664 graph.nodes[0].op,
665 GraphOp::Gemm { alpha: 1, beta: 1 }
666 ));
667 assert_eq!(graph.nodes[1].op, GraphOp::Relu);
668 }
669
670 #[test]
671 fn conv_bias_no_activation() {
672 let mut graph = ComputeGraph::new();
674
675 let input = graph.add_edge(make_tensor("input", &[-1, 3, 224, 224]));
676 let weight = graph.add_edge(make_tensor("weight", &[64, 3, 3, 3]));
677 let bias = graph.add_edge(make_tensor("bias", &[64]));
678 let conv_out = graph.add_edge(make_tensor("conv_out", &[-1, 64, 222, 222]));
679 let add_out = graph.add_edge(make_tensor("add_out", &[-1, 64, 222, 222]));
680
681 graph.inputs = vec![input, weight, bias];
682 graph.outputs = vec![add_out];
683
684 graph
685 .add_node(GraphOp::Conv2d, vec![input, weight], vec![conv_out], "conv")
686 .unwrap();
687 graph
688 .add_node(
689 GraphOp::Add,
690 vec![conv_out, bias],
691 vec![add_out],
692 "bias_add",
693 )
694 .unwrap();
695
696 let fusion = OperatorFusion;
697 let changed = fusion.run_on_graph(&mut graph);
698
699 assert!(changed);
700 assert_eq!(graph.node_count(), 1);
701 assert_eq!(
702 graph.nodes[0].op,
703 GraphOp::FusedConv2d {
704 activation: ActivationFunction::None
705 }
706 );
707 assert_eq!(graph.nodes[0].inputs, vec![input, weight, bias]);
708 }
709
710 #[test]
711 fn pass_on_empty_graph() {
712 let mut graph = ComputeGraph::new();
713 let fusion = OperatorFusion;
714 let changed = fusion.run_on_graph(&mut graph);
715 assert!(!changed);
716 }
717
718 #[test]
719 fn pass_trait_on_module_is_noop() {
720 let fusion = OperatorFusion;
721 let mut module = nxpu_ir::Module::default();
722 let changed = fusion.run(&mut module);
723 assert!(!changed);
724 }
725
726 #[test]
727 fn fuse_sub_relu() {
728 let mut graph = ComputeGraph::new();
729
730 let a = graph.add_edge(make_tensor("a", &[-1, 256]));
731 let b = graph.add_edge(make_tensor("b", &[-1, 256]));
732 let sub_out = graph.add_edge(make_tensor("sub_out", &[-1, 256]));
733 let relu_out = graph.add_edge(make_tensor("relu_out", &[-1, 256]));
734
735 graph.inputs = vec![a, b];
736 graph.outputs = vec![relu_out];
737
738 graph
739 .add_node(GraphOp::Sub, vec![a, b], vec![sub_out], "sub")
740 .unwrap();
741 graph
742 .add_node(GraphOp::Relu, vec![sub_out], vec![relu_out], "relu")
743 .unwrap();
744
745 let fusion = OperatorFusion;
746 let changed = fusion.run_on_graph(&mut graph);
747
748 assert!(changed);
749 assert_eq!(graph.node_count(), 1);
750 assert_eq!(
751 graph.nodes[0].op,
752 GraphOp::FusedElementWise {
753 base_op: Box::new(GraphOp::Sub),
754 activation: ActivationFunction::Relu,
755 }
756 );
757 }
758
759 #[test]
760 fn fuse_div_sigmoid() {
761 let mut graph = ComputeGraph::new();
762
763 let a = graph.add_edge(make_tensor("a", &[-1, 256]));
764 let b = graph.add_edge(make_tensor("b", &[-1, 256]));
765 let div_out = graph.add_edge(make_tensor("div_out", &[-1, 256]));
766 let sig_out = graph.add_edge(make_tensor("sig_out", &[-1, 256]));
767
768 graph.inputs = vec![a, b];
769 graph.outputs = vec![sig_out];
770
771 graph
772 .add_node(GraphOp::Div, vec![a, b], vec![div_out], "div")
773 .unwrap();
774 graph
775 .add_node(GraphOp::Sigmoid, vec![div_out], vec![sig_out], "sigmoid")
776 .unwrap();
777
778 let fusion = OperatorFusion;
779 let changed = fusion.run_on_graph(&mut graph);
780
781 assert!(changed);
782 assert_eq!(graph.node_count(), 1);
783 assert_eq!(
784 graph.nodes[0].op,
785 GraphOp::FusedElementWise {
786 base_op: Box::new(GraphOp::Div),
787 activation: ActivationFunction::Sigmoid,
788 }
789 );
790 }
791
792 #[test]
793 fn conv_no_consumer_no_fusion() {
794 let mut graph = ComputeGraph::new();
796
797 let input = graph.add_edge(make_tensor("input", &[-1, 3, 224, 224]));
798 let weight = graph.add_edge(make_tensor("weight", &[64, 3, 3, 3]));
799 let conv_out = graph.add_edge(make_tensor("conv_out", &[-1, 64, 222, 222]));
800
801 graph.inputs = vec![input, weight];
802 graph.outputs = vec![conv_out];
803
804 graph
805 .add_node(GraphOp::Conv2d, vec![input, weight], vec![conv_out], "conv")
806 .unwrap();
807
808 let fusion = OperatorFusion;
809 let changed = fusion.run_on_graph(&mut graph);
810
811 assert!(!changed);
812 assert_eq!(graph.node_count(), 1);
813 }
814
815 #[test]
816 fn conv_with_non_fusible_consumer() {
817 let mut graph = ComputeGraph::new();
819
820 let input = graph.add_edge(make_tensor("input", &[-1, 3, 224, 224]));
821 let weight = graph.add_edge(make_tensor("weight", &[64, 3, 3, 3]));
822 let conv_out = graph.add_edge(make_tensor("conv_out", &[-1, 64, 222, 222]));
823 let reshape_out = graph.add_edge(make_tensor("reshape_out", &[-1, 64, 222, 222]));
824
825 graph.inputs = vec![input, weight];
826 graph.outputs = vec![reshape_out];
827
828 graph
829 .add_node(GraphOp::Conv2d, vec![input, weight], vec![conv_out], "conv")
830 .unwrap();
831 graph
832 .add_node(
833 GraphOp::Reshape,
834 vec![conv_out],
835 vec![reshape_out],
836 "reshape",
837 )
838 .unwrap();
839
840 let fusion = OperatorFusion;
841 let changed = fusion.run_on_graph(&mut graph);
842
843 assert!(!changed);
844 assert_eq!(graph.node_count(), 2);
845 }
846
847 #[test]
848 fn conv_sigmoid_fusion() {
849 let mut graph = ComputeGraph::new();
851
852 let input = graph.add_edge(make_tensor("input", &[-1, 3, 224, 224]));
853 let weight = graph.add_edge(make_tensor("weight", &[64, 3, 3, 3]));
854 let conv_out = graph.add_edge(make_tensor("conv_out", &[-1, 64, 222, 222]));
855 let sig_out = graph.add_edge(make_tensor("sig_out", &[-1, 64, 222, 222]));
856
857 graph.inputs = vec![input, weight];
858 graph.outputs = vec![sig_out];
859
860 graph
861 .add_node(GraphOp::Conv2d, vec![input, weight], vec![conv_out], "conv")
862 .unwrap();
863 graph
864 .add_node(GraphOp::Sigmoid, vec![conv_out], vec![sig_out], "sigmoid")
865 .unwrap();
866
867 let fusion = OperatorFusion;
868 let changed = fusion.run_on_graph(&mut graph);
869
870 assert!(changed);
871 assert_eq!(graph.node_count(), 1);
872 assert_eq!(
873 graph.nodes[0].op,
874 GraphOp::FusedConv2d {
875 activation: ActivationFunction::Sigmoid,
876 }
877 );
878 }
879
880 #[test]
881 fn conv_bias_sigmoid_fusion() {
882 let mut graph = ComputeGraph::new();
884
885 let input = graph.add_edge(make_tensor("input", &[-1, 3, 224, 224]));
886 let weight = graph.add_edge(make_tensor("weight", &[64, 3, 3, 3]));
887 let bias = graph.add_edge(make_tensor("bias", &[64]));
888 let conv_out = graph.add_edge(make_tensor("conv_out", &[-1, 64, 222, 222]));
889 let add_out = graph.add_edge(make_tensor("add_out", &[-1, 64, 222, 222]));
890 let sig_out = graph.add_edge(make_tensor("sig_out", &[-1, 64, 222, 222]));
891
892 graph.inputs = vec![input, weight, bias];
893 graph.outputs = vec![sig_out];
894
895 graph
896 .add_node(GraphOp::Conv2d, vec![input, weight], vec![conv_out], "conv")
897 .unwrap();
898 graph
899 .add_node(
900 GraphOp::Add,
901 vec![conv_out, bias],
902 vec![add_out],
903 "bias_add",
904 )
905 .unwrap();
906 graph
907 .add_node(GraphOp::Sigmoid, vec![add_out], vec![sig_out], "sigmoid")
908 .unwrap();
909
910 let fusion = OperatorFusion;
911 let changed = fusion.run_on_graph(&mut graph);
912
913 assert!(changed);
914 assert_eq!(graph.node_count(), 1);
915 assert_eq!(
916 graph.nodes[0].op,
917 GraphOp::FusedConv2d {
918 activation: ActivationFunction::Sigmoid,
919 }
920 );
921 assert_eq!(graph.nodes[0].inputs, vec![input, weight, bias]);
922 }
923
924 #[test]
925 fn matmul_without_add_no_fusion() {
926 let mut graph = ComputeGraph::new();
928
929 let a = graph.add_edge(make_tensor("A", &[-1, 768]));
930 let b = graph.add_edge(make_tensor("B", &[768, 768]));
931 let mm_out = graph.add_edge(make_tensor("mm_out", &[-1, 768]));
932
933 graph.inputs = vec![a, b];
934 graph.outputs = vec![mm_out];
935
936 graph
937 .add_node(GraphOp::MatMul, vec![a, b], vec![mm_out], "matmul")
938 .unwrap();
939
940 let fusion = OperatorFusion;
941 let changed = fusion.run_on_graph(&mut graph);
942
943 assert!(!changed);
944 assert_eq!(graph.node_count(), 1);
945 }
946
947 #[test]
948 fn elementwise_without_activation_no_fusion() {
949 let mut graph = ComputeGraph::new();
951
952 let a = graph.add_edge(make_tensor("a", &[-1, 256]));
953 let b = graph.add_edge(make_tensor("b", &[-1, 256]));
954 let add_out = graph.add_edge(make_tensor("add_out", &[-1, 256]));
955
956 graph.inputs = vec![a, b];
957 graph.outputs = vec![add_out];
958
959 graph
960 .add_node(GraphOp::Add, vec![a, b], vec![add_out], "add")
961 .unwrap();
962
963 let fusion = OperatorFusion;
964 let changed = fusion.run_on_graph(&mut graph);
965
966 assert!(!changed);
967 assert_eq!(graph.node_count(), 1);
968 }
969
970 #[test]
971 fn elementwise_with_non_activation_consumer() {
972 let mut graph = ComputeGraph::new();
974
975 let a = graph.add_edge(make_tensor("a", &[-1, 256]));
976 let b = graph.add_edge(make_tensor("b", &[-1, 256]));
977 let add_out = graph.add_edge(make_tensor("add_out", &[-1, 256]));
978 let reshape_out = graph.add_edge(make_tensor("reshape_out", &[-1, 256]));
979
980 graph.inputs = vec![a, b];
981 graph.outputs = vec![reshape_out];
982
983 graph
984 .add_node(GraphOp::Add, vec![a, b], vec![add_out], "add")
985 .unwrap();
986 graph
987 .add_node(
988 GraphOp::Reshape,
989 vec![add_out],
990 vec![reshape_out],
991 "reshape",
992 )
993 .unwrap();
994
995 let fusion = OperatorFusion;
996 let changed = fusion.run_on_graph(&mut graph);
997
998 assert!(!changed);
999 assert_eq!(graph.node_count(), 2);
1000 }
1001
1002 #[test]
1003 fn conv_multiple_outputs_no_fusion() {
1004 let mut graph = ComputeGraph::new();
1006
1007 let input = graph.add_edge(make_tensor("input", &[-1, 3, 224, 224]));
1008 let weight = graph.add_edge(make_tensor("weight", &[64, 3, 3, 3]));
1009 let conv_out1 = graph.add_edge(make_tensor("conv_out1", &[-1, 64, 222, 222]));
1010 let conv_out2 = graph.add_edge(make_tensor("conv_out2", &[-1, 64, 222, 222]));
1011
1012 graph.inputs = vec![input, weight];
1013 graph.outputs = vec![conv_out1, conv_out2];
1014
1015 graph
1016 .add_node(
1017 GraphOp::Conv2d,
1018 vec![input, weight],
1019 vec![conv_out1, conv_out2],
1020 "conv",
1021 )
1022 .unwrap();
1023
1024 let fusion = OperatorFusion;
1025 let changed = fusion.run_on_graph(&mut graph);
1026
1027 assert!(!changed);
1028 }
1029
1030 #[test]
1031 fn matmul_multiple_outputs_no_fusion() {
1032 let mut graph = ComputeGraph::new();
1034
1035 let a = graph.add_edge(make_tensor("A", &[-1, 768]));
1036 let b = graph.add_edge(make_tensor("B", &[768, 768]));
1037 let mm_out1 = graph.add_edge(make_tensor("mm_out1", &[-1, 768]));
1038 let mm_out2 = graph.add_edge(make_tensor("mm_out2", &[-1, 768]));
1039
1040 graph.inputs = vec![a, b];
1041 graph.outputs = vec![mm_out1, mm_out2];
1042
1043 graph
1044 .add_node(
1045 GraphOp::MatMul,
1046 vec![a, b],
1047 vec![mm_out1, mm_out2],
1048 "matmul",
1049 )
1050 .unwrap();
1051
1052 let fusion = OperatorFusion;
1053 let changed = fusion.run_on_graph(&mut graph);
1054
1055 assert!(!changed);
1056 }
1057
1058 #[test]
1059 fn add_multi_output_no_fusion() {
1060 let mut graph = ComputeGraph::new();
1062
1063 let a = graph.add_edge(make_tensor("a", &[-1, 256]));
1064 let b = graph.add_edge(make_tensor("b", &[-1, 256]));
1065 let add_out1 = graph.add_edge(make_tensor("add_out1", &[-1, 256]));
1066 let add_out2 = graph.add_edge(make_tensor("add_out2", &[-1, 256]));
1067
1068 graph.inputs = vec![a, b];
1069 graph.outputs = vec![add_out1, add_out2];
1070
1071 graph
1072 .add_node(GraphOp::Add, vec![a, b], vec![add_out1, add_out2], "add")
1073 .unwrap();
1074
1075 let fusion = OperatorFusion;
1076 let changed = fusion.run_on_graph(&mut graph);
1077
1078 assert!(!changed);
1079 }
1080
1081 #[test]
1082 fn operator_fusion_name() {
1083 let fusion = OperatorFusion;
1084 assert_eq!(fusion.name(), "operator-fusion");
1085 }
1086}