Skip to main content

nxpu_opt/
fusion.rs

1//! Operator fusion pass.
2//!
3//! Fuses adjacent operations in the compute graph into single fused
4//! operations to reduce memory traffic and kernel launch overhead.
5//!
6//! Supported fusion patterns:
7//! - **Conv2D + Bias + Activation** → `FusedConv2d` with activation attribute
8//! - **MatMul + Bias (Add)** → `Gemm` (with beta=1)
9//! - **ElementWise + Activation** → `FusedElementWise` with activation
10
11use std::collections::HashMap;
12
13use nxpu_ir::graph::{ActivationFunction, ComputeGraph, EdgeId, GraphOp};
14
15use crate::Pass;
16
17/// Fuses adjacent operations in the compute graph into single fused ops.
18#[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        // This pass is a no-op on Module (expression-level IR).
28        // It operates on ComputeGraph via `run_on_graph`.
29        let _ = module;
30        false
31    }
32}
33
34impl OperatorFusion {
35    /// Run the fusion pass on a compute graph.
36    /// Returns `true` if any fusion was performed.
37    pub fn run_on_graph(&self, graph: &mut ComputeGraph) -> bool {
38        let mut changed = false;
39
40        // Keep running until no more fusions are possible (fixed-point).
41        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
57/// Returns `true` if the given `GraphOp` is an activation function.
58fn 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
66/// Returns `true` if the given `GraphOp` is an element-wise binary op.
67fn is_elementwise_binary(op: &GraphOp) -> bool {
68    matches!(
69        op,
70        GraphOp::Add | GraphOp::Sub | GraphOp::Mul | GraphOp::Div
71    )
72}
73
74/// Build a map from edge → list of consumer node indices.
75fn 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
85/// Check if an edge has exactly one consumer in the graph.
86fn 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
93/// Remove a node from the graph by index and rewire edges.
94///
95/// The consumer node is removed, and the graph output edges of the removed
96/// node become the output edges of the surviving (fused) node.
97fn remove_node_and_rewire(
98    graph: &mut ComputeGraph,
99    remove_idx: usize,
100    fused_idx: usize,
101    intermediate_edge: EdgeId,
102) {
103    // The fused node takes the output edges from the removed node.
104    let new_outputs = graph.nodes[remove_idx].outputs.clone();
105    graph.nodes[fused_idx].outputs = new_outputs;
106
107    // Remove the intermediate edge from graph inputs/outputs if present.
108    graph.inputs.retain(|e| *e != intermediate_edge);
109    graph.outputs.retain(|e| *e != intermediate_edge);
110
111    // Remove the consumed node.
112    graph.nodes.remove(remove_idx);
113
114    // Clean up the intermediate edge from the edges map.
115    graph.edges.remove(&intermediate_edge);
116}
117
118/// Try to fuse Conv2D + (optional Add/bias) + (optional activation).
119///
120/// Pattern: Conv2d → Add(conv_out, bias) → Relu/Sigmoid
121/// Result:  FusedConv2d { activation: Relu } with inputs [input, weight, bias]
122fn try_fuse_conv_bias_activation(graph: &mut ComputeGraph) -> bool {
123    let edge_consumer = build_edge_consumer_map(graph);
124
125    // Find Conv2D nodes.
126    for conv_idx in 0..graph.nodes.len() {
127        if graph.nodes[conv_idx].op != GraphOp::Conv2d {
128            continue;
129        }
130
131        // Conv2d must have exactly one output edge.
132        if graph.nodes[conv_idx].outputs.len() != 1 {
133            continue;
134        }
135        let conv_out_edge = graph.nodes[conv_idx].outputs[0];
136
137        // The output must have a single consumer.
138        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        // Check if the consumer is an Add (bias) or an activation.
146        match &graph.nodes[next_idx].op {
147            GraphOp::Add => {
148                // Conv2d + Add (bias pattern). Try to find a subsequent activation.
149                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                // Get the bias input (the non-conv input to Add).
156                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                // Add the bias to the conv inputs.
167                let mut fused_inputs = graph.nodes[conv_idx].inputs.clone();
168                fused_inputs.push(bias_edge);
169
170                // Check if there's an activation after the Add.
171                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                // Create the fused op.
182                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 was found, remove it too.
189                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                // Conv2d + Activation (no bias).
197                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
214/// After fusing Conv+Bias, remove the activation node that follows.
215fn 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
238/// Try to fuse MatMul + Add → Gemm.
239///
240/// Pattern: MatMul(A, B) → Add(matmul_out, C) where C is the bias
241/// Result:  Gemm { alpha: 1, beta: 1 } with inputs [A, B, C]
242fn 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        // Get the bias input (the non-matmul input to Add).
267        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        // Fuse: MatMul + Add → Gemm.
278        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
293/// Try to fuse ElementWise + Activation.
294///
295/// Pattern: Add/Sub/Mul/Div → Relu/Sigmoid
296/// Result:  FusedElementWise { base_op, activation }
297fn 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        // Fuse: ElementWise + Activation → FusedElementWise.
323        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        // The fused node should output relu_out, not conv_out.
396        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        // Should include bias in inputs.
442        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        // MatMul followed by Reshape — should NOT fuse.
518        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        // If the intermediate edge has multiple consumers, do not fuse.
545        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        // mm_out is consumed by both Add and Relu (two consumers).
561        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        // MatMul+Add should NOT fuse because mm_out has 2 consumers.
572        assert!(!changed);
573        assert_eq!(graph.node_count(), 3);
574    }
575
576    #[test]
577    fn no_fusion_standalone_activation() {
578        // A standalone Relu with no preceding fusible op.
579        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        // MatMul → Add → Relu: MatMul+Add fuses to Gemm, but Relu stays
634        // because Gemm is not an elementwise binary op.
635        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        // MatMul+Add fuses to Gemm; Relu remains separate.
662        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        // Conv2D + Add (bias) without activation.
673        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        // Conv2D with no consumer (output is a graph output).
795        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        // Conv2D followed by Reshape (not Add or activation).
818        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        // Conv2D + Sigmoid (no bias).
850        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        // Conv2D + Add(bias) + Sigmoid.
883        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        // MatMul without a subsequent Add.
927        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        // Add without a subsequent activation.
950        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        // Add followed by Reshape (not an activation).
973        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        // Conv2D with multiple outputs -- should not fuse.
1005        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        // MatMul with multiple outputs -- should not fuse.
1033        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        // Add with multiple outputs -- should not fuse.
1061        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}