1use std::collections::{HashMap, HashSet, VecDeque};
68
69use crate::GpuOptimError;
70
71#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
77pub enum OpKind {
78 Add,
80 Mul,
82 Sub,
84 Div,
86 Relu,
88 Sigmoid,
90 Tanh,
92 Scale,
94 Axpy,
96 MatMul,
98 Reduce,
100 Transpose,
102}
103
104impl OpKind {
105 pub fn is_fusible(&self) -> bool {
108 matches!(
109 self,
110 OpKind::Add
111 | OpKind::Mul
112 | OpKind::Sub
113 | OpKind::Div
114 | OpKind::Relu
115 | OpKind::Sigmoid
116 | OpKind::Tanh
117 | OpKind::Scale
118 | OpKind::Axpy
119 )
120 }
121
122 pub fn is_barrier(&self) -> bool {
125 !self.is_fusible()
126 }
127}
128
129#[derive(Debug, Clone)]
135pub struct FusionOp {
136 pub id: usize,
138 pub kind: OpKind,
140 pub inputs: Vec<usize>,
142 pub output_shape: Vec<usize>,
144 pub dtype_bytes: usize,
146}
147
148impl FusionOp {
149 pub fn output_elements(&self) -> Result<u64, GpuOptimError> {
153 let mut elements: u64 = 1;
154 for &dim in &self.output_shape {
155 elements = elements.checked_mul(dim as u64).ok_or_else(|| {
156 GpuOptimError::UnsupportedOperation(format!(
157 "op {} output shape {:?} overflows the element counter",
158 self.id, self.output_shape
159 ))
160 })?;
161 }
162 Ok(elements)
163 }
164
165 pub fn output_bytes(&self) -> Result<u64, GpuOptimError> {
169 let elements = self.output_elements()?;
170 elements
171 .checked_mul(self.dtype_bytes as u64)
172 .ok_or_else(|| {
173 GpuOptimError::UnsupportedOperation(format!(
174 "op {} output ({} elements x {} bytes) overflows the byte counter",
175 self.id, elements, self.dtype_bytes
176 ))
177 })
178 }
179}
180
181#[derive(Debug, Default, Clone)]
183pub struct FusionGraph {
184 ops: Vec<FusionOp>,
185}
186
187impl FusionGraph {
188 pub fn new() -> Self {
190 Self { ops: Vec::new() }
191 }
192
193 pub fn add_op(
200 &mut self,
201 kind: OpKind,
202 inputs: Vec<usize>,
203 output_shape: Vec<usize>,
204 dtype_bytes: usize,
205 ) -> usize {
206 let id = self.ops.len();
207 self.ops.push(FusionOp {
208 id,
209 kind,
210 inputs,
211 output_shape,
212 dtype_bytes,
213 });
214 id
215 }
216
217 pub fn ops(&self) -> &[FusionOp] {
219 &self.ops
220 }
221
222 pub fn num_ops(&self) -> usize {
224 self.ops.len()
225 }
226
227 pub fn is_empty(&self) -> bool {
229 self.ops.is_empty()
230 }
231
232 pub fn validate(&self) -> Result<(), GpuOptimError> {
239 let n = self.ops.len();
240 for op in &self.ops {
241 if op.dtype_bytes == 0 {
242 return Err(GpuOptimError::InvalidState(format!(
243 "op {} has dtype_bytes == 0",
244 op.id
245 )));
246 }
247 for &producer in &op.inputs {
248 if producer >= n {
249 return Err(GpuOptimError::InvalidState(format!(
250 "op {} references non-existent input op {}",
251 op.id, producer
252 )));
253 }
254 if producer == op.id {
255 return Err(GpuOptimError::InvalidState(format!(
256 "op {} references itself as an input",
257 op.id
258 )));
259 }
260 }
261 op.output_bytes()?;
263 }
264 self.topological_order()?;
265 Ok(())
266 }
267
268 fn topological_order(&self) -> Result<Vec<usize>, GpuOptimError> {
273 let n = self.ops.len();
274 let mut indegree = vec![0usize; n];
275 let mut adjacency: Vec<Vec<usize>> = vec![Vec::new(); n];
276 for consumer in &self.ops {
277 let mut seen: HashSet<usize> = HashSet::new();
278 for &producer in &consumer.inputs {
279 if producer >= n {
280 return Err(GpuOptimError::InvalidState(format!(
281 "op {} references non-existent input op {}",
282 consumer.id, producer
283 )));
284 }
285 if !seen.insert(producer) {
287 continue;
288 }
289 adjacency[producer].push(consumer.id);
290 indegree[consumer.id] += 1;
291 }
292 }
293
294 let mut queue: VecDeque<usize> = (0..n).filter(|&i| indegree[i] == 0).collect();
295 let mut order = Vec::with_capacity(n);
296 while let Some(node) = queue.pop_front() {
297 order.push(node);
298 for &consumer in &adjacency[node] {
299 indegree[consumer] -= 1;
300 if indegree[consumer] == 0 {
301 queue.push_back(consumer);
302 }
303 }
304 }
305
306 if order.len() != n {
307 return Err(GpuOptimError::InvalidState(
308 "operation graph contains a cycle".to_string(),
309 ));
310 }
311 Ok(order)
312 }
313}
314
315#[derive(Debug, Clone)]
323pub struct FusionGroup {
324 pub members: Vec<usize>,
326 pub external_inputs: Vec<usize>,
328 pub external_outputs: Vec<usize>,
330 pub internal_intermediates: Vec<usize>,
332 pub bytes_read: u64,
334 pub bytes_written: u64,
336}
337
338impl FusionGroup {
339 pub fn bytes_fused(&self) -> u64 {
341 self.bytes_read + self.bytes_written
342 }
343
344 pub fn is_fused(&self) -> bool {
346 self.members.len() > 1
347 }
348}
349
350#[derive(Debug, Clone)]
353pub struct FusionPlan {
354 pub groups: Vec<FusionGroup>,
356 pub bytes_unfused: u64,
358 pub bytes_fused: u64,
360 pub bytes_saved: u64,
362 pub speedup_estimate: f64,
364}
365
366impl FusionPlan {
367 pub fn num_groups(&self) -> usize {
369 self.groups.len()
370 }
371}
372
373#[derive(Debug, Clone)]
375pub struct FusionPlanner {
376 allow_broadcast: bool,
377}
378
379impl Default for FusionPlanner {
380 fn default() -> Self {
381 Self::new()
382 }
383}
384
385impl FusionPlanner {
386 pub fn new() -> Self {
388 Self {
389 allow_broadcast: true,
390 }
391 }
392
393 pub fn with_broadcast(mut self, allow_broadcast: bool) -> Self {
396 self.allow_broadcast = allow_broadcast;
397 self
398 }
399
400 fn shapes_fuse_compatible(&self, producer: &[usize], consumer: &[usize]) -> bool {
409 if producer == consumer {
410 return true;
411 }
412 if !self.allow_broadcast {
413 return false;
414 }
415 if producer.len() > consumer.len() {
416 return false;
417 }
418 let offset = consumer.len() - producer.len();
419 for (i, &producer_dim) in producer.iter().enumerate() {
420 let consumer_dim = consumer[offset + i];
421 if producer_dim != consumer_dim && producer_dim != 1 {
422 return false;
423 }
424 }
425 true
426 }
427
428 pub fn plan(&self, graph: &FusionGraph) -> Result<FusionPlan, GpuOptimError> {
433 graph.validate()?;
434 let ops = graph.ops();
435 let n = ops.len();
436
437 let mut bytes: Vec<u64> = Vec::with_capacity(n);
439 for op in ops {
440 bytes.push(op.output_bytes()?);
441 }
442
443 let topo = graph.topological_order()?;
444
445 let mut consumers: Vec<Vec<usize>> = vec![Vec::new(); n];
447 for consumer in ops {
448 let mut seen: HashSet<usize> = HashSet::new();
449 for &producer in &consumer.inputs {
450 if seen.insert(producer) {
451 consumers[producer].push(consumer.id);
452 }
453 }
454 }
455
456 let mut group_of: Vec<usize> = (0..n).collect();
458 for &consumer_id in &topo {
459 let consumer = &ops[consumer_id];
460 if !consumer.kind.is_fusible() {
461 continue;
462 }
463 let mut seen: HashSet<usize> = HashSet::new();
464 for &producer_id in &consumer.inputs {
465 if !seen.insert(producer_id) {
466 continue;
467 }
468 let producer = &ops[producer_id];
469 if !producer.kind.is_fusible() {
470 continue;
471 }
472 if !self.shapes_fuse_compatible(&producer.output_shape, &consumer.output_shape) {
473 continue;
474 }
475 let group_producer = group_of[producer_id];
476 let group_consumer = group_of[consumer_id];
477 if group_producer == group_consumer {
478 continue;
479 }
480 if merge_keeps_acyclic(ops, &group_of, group_producer, group_consumer) {
482 for label in group_of.iter_mut() {
483 if *label == group_consumer {
484 *label = group_producer;
485 }
486 }
487 }
488 }
489 }
490
491 let mut label_to_members: HashMap<usize, Vec<usize>> = HashMap::new();
493 for &id in &topo {
494 label_to_members.entry(group_of[id]).or_default().push(id);
495 }
496 let mut raw_groups: Vec<Vec<usize>> = label_to_members.into_values().collect();
497 for members in raw_groups.iter_mut() {
498 members.sort_unstable();
499 }
500 raw_groups.sort_by_key(|members| members[0]);
501
502 let mut groups: Vec<FusionGroup> = Vec::with_capacity(raw_groups.len());
504 let mut bytes_fused: u64 = 0;
505 for members in raw_groups {
506 let member_set: HashSet<usize> = members.iter().copied().collect();
507
508 let mut external_inputs: Vec<usize> = Vec::new();
509 let mut external_input_seen: HashSet<usize> = HashSet::new();
510 for &member in &members {
511 let mut seen: HashSet<usize> = HashSet::new();
512 for &producer in &ops[member].inputs {
513 if !seen.insert(producer) {
514 continue;
515 }
516 if !member_set.contains(&producer) && external_input_seen.insert(producer) {
517 external_inputs.push(producer);
518 }
519 }
520 }
521 external_inputs.sort_unstable();
522
523 let mut external_outputs: Vec<usize> = Vec::new();
524 let mut internal_intermediates: Vec<usize> = Vec::new();
525 for &member in &members {
526 let consumed_externally = consumers[member].iter().any(|c| !member_set.contains(c));
527 let is_terminal = consumers[member].is_empty();
528 if consumed_externally || is_terminal {
529 external_outputs.push(member);
530 } else {
531 internal_intermediates.push(member);
532 }
533 }
534
535 let bytes_read: u64 = external_inputs.iter().map(|&p| bytes[p]).sum();
536 let bytes_written: u64 = external_outputs.iter().map(|&m| bytes[m]).sum();
537 bytes_fused += bytes_read + bytes_written;
538
539 groups.push(FusionGroup {
540 members,
541 external_inputs,
542 external_outputs,
543 internal_intermediates,
544 bytes_read,
545 bytes_written,
546 });
547 }
548
549 let mut bytes_unfused: u64 = 0;
551 for op in ops {
552 let mut seen: HashSet<usize> = HashSet::new();
553 let mut read: u64 = 0;
554 for &producer in &op.inputs {
555 if seen.insert(producer) {
556 read += bytes[producer];
557 }
558 }
559 bytes_unfused += read + bytes[op.id];
560 }
561
562 let bytes_saved = bytes_unfused.saturating_sub(bytes_fused);
563 let speedup_estimate = if bytes_fused == 0 {
564 1.0
565 } else {
566 bytes_unfused as f64 / bytes_fused as f64
567 };
568
569 Ok(FusionPlan {
570 groups,
571 bytes_unfused,
572 bytes_fused,
573 bytes_saved,
574 speedup_estimate,
575 })
576 }
577}
578
579fn merge_keeps_acyclic(
586 ops: &[FusionOp],
587 group_of: &[usize],
588 group_a: usize,
589 group_b: usize,
590) -> bool {
591 let label = |op_id: usize| -> usize {
592 let group = group_of[op_id];
593 if group == group_b {
594 group_a
595 } else {
596 group
597 }
598 };
599
600 let mut adjacency: HashMap<usize, HashSet<usize>> = HashMap::new();
601 let mut nodes: HashSet<usize> = HashSet::new();
602 for consumer in ops {
603 let consumer_label = label(consumer.id);
604 nodes.insert(consumer_label);
605 for &producer in &consumer.inputs {
606 let producer_label = label(producer);
607 nodes.insert(producer_label);
608 if producer_label != consumer_label {
609 adjacency
610 .entry(producer_label)
611 .or_default()
612 .insert(consumer_label);
613 }
614 }
615 }
616
617 let mut indegree: HashMap<usize, usize> = nodes.iter().map(|&node| (node, 0usize)).collect();
618 for targets in adjacency.values() {
619 for &target in targets {
620 if let Some(degree) = indegree.get_mut(&target) {
621 *degree += 1;
622 }
623 }
624 }
625
626 let mut queue: VecDeque<usize> = indegree
627 .iter()
628 .filter_map(|(&node, °ree)| if degree == 0 { Some(node) } else { None })
629 .collect();
630 let mut visited = 0usize;
631 while let Some(node) = queue.pop_front() {
632 visited += 1;
633 if let Some(targets) = adjacency.get(&node) {
634 for &target in targets {
635 if let Some(degree) = indegree.get_mut(&target) {
636 *degree -= 1;
637 if *degree == 0 {
638 queue.push_back(target);
639 }
640 }
641 }
642 }
643 }
644
645 visited == nodes.len()
646}
647
648#[cfg(test)]
649mod tests {
650 use super::*;
651
652 fn find_group(plan: &FusionPlan, op_id: usize) -> &FusionGroup {
653 plan.groups
654 .iter()
655 .find(|g| g.members.contains(&op_id))
656 .expect("every op must belong to exactly one group")
657 }
658
659 #[test]
660 fn is_fusible_classification() {
661 assert!(OpKind::Add.is_fusible());
662 assert!(OpKind::Mul.is_fusible());
663 assert!(OpKind::Axpy.is_fusible());
664 assert!(OpKind::Scale.is_fusible());
665 assert!(!OpKind::MatMul.is_fusible());
666 assert!(!OpKind::Reduce.is_fusible());
667 assert!(!OpKind::Transpose.is_fusible());
668 assert!(OpKind::MatMul.is_barrier());
669 assert!(!OpKind::Relu.is_barrier());
670 }
671
672 #[test]
673 fn linear_chain_fuses_into_one_group() {
674 let mut graph = FusionGraph::new();
675 let a = graph.add_op(OpKind::Relu, vec![], vec![256], 4);
676 let b = graph.add_op(OpKind::Sigmoid, vec![a], vec![256], 4);
677 let c = graph.add_op(OpKind::Tanh, vec![b], vec![256], 4);
678
679 let plan = FusionPlanner::new()
680 .plan(&graph)
681 .expect("fusible chain must plan");
682
683 assert_eq!(plan.groups.len(), 1);
684 let group = &plan.groups[0];
685 assert_eq!(group.members, vec![a, b, c]);
686 assert!(group.external_inputs.is_empty());
687 assert_eq!(group.external_outputs, vec![c]);
688 assert_eq!(group.internal_intermediates, vec![a, b]);
689
690 let tensor = 256u64 * 4;
691 assert_eq!(plan.bytes_unfused, 5 * tensor);
692 assert_eq!(plan.bytes_fused, tensor);
693 assert!(plan.bytes_fused < plan.bytes_unfused);
694 assert_eq!(plan.bytes_saved, 4 * tensor);
695 assert!((plan.speedup_estimate - 5.0).abs() < 1e-9);
696 }
697
698 #[test]
699 fn barrier_splits_into_three_groups() {
700 let mut graph = FusionGraph::new();
701 let a = graph.add_op(OpKind::Relu, vec![], vec![128], 4);
702 let b = graph.add_op(OpKind::Sigmoid, vec![a], vec![128], 4);
703 let c = graph.add_op(OpKind::MatMul, vec![b], vec![128], 4); let d = graph.add_op(OpKind::Relu, vec![c], vec![128], 4);
705 let e = graph.add_op(OpKind::Tanh, vec![d], vec![128], 4);
706
707 let plan = FusionPlanner::new()
708 .plan(&graph)
709 .expect("graph with a barrier must plan");
710
711 assert_eq!(plan.groups.len(), 3);
712 assert_eq!(find_group(&plan, a).members, vec![a, b]);
713 assert_eq!(find_group(&plan, c).members, vec![c]);
714 assert_eq!(find_group(&plan, d).members, vec![d, e]);
715 }
716
717 #[test]
718 fn external_consumer_materializes_intermediate() {
719 let mut graph = FusionGraph::new();
720 let a = graph.add_op(OpKind::Relu, vec![], vec![256], 4);
721 let b = graph.add_op(OpKind::Sigmoid, vec![a], vec![256], 4);
722 let c = graph.add_op(OpKind::Tanh, vec![b], vec![256], 4);
723 let d = graph.add_op(OpKind::MatMul, vec![b], vec![256], 4);
725
726 let plan = FusionPlanner::new()
727 .plan(&graph)
728 .expect("diamond graph must plan");
729
730 assert_eq!(plan.groups.len(), 2);
731 let group = find_group(&plan, b);
732 assert_eq!(group.members, vec![a, b, c]);
733 assert!(
734 group.external_outputs.contains(&b),
735 "b is consumed outside the group and must be materialized"
736 );
737 assert!(!group.internal_intermediates.contains(&b));
738 assert!(group.external_outputs.contains(&c));
739 assert_eq!(group.internal_intermediates, vec![a]);
740 assert_eq!(find_group(&plan, d).members, vec![d]);
741
742 let tensor = 256u64 * 4;
744 assert_eq!(plan.bytes_unfused, 7 * tensor);
745 assert_eq!(plan.bytes_fused, 4 * tensor);
746 assert_eq!(plan.bytes_saved, 3 * tensor);
747 assert!((plan.speedup_estimate - 1.75).abs() < 1e-9);
748 }
749
750 #[test]
751 fn shared_external_input_counted_once() {
752 let mut graph = FusionGraph::new();
753 let x = graph.add_op(OpKind::MatMul, vec![], vec![16], 4);
755 let r = graph.add_op(OpKind::Relu, vec![x], vec![16], 4);
756 let s = graph.add_op(OpKind::Add, vec![r, x], vec![16], 4);
757
758 let plan = FusionPlanner::new().plan(&graph).expect("graph must plan");
759
760 assert_eq!(plan.groups.len(), 2);
761 let group = find_group(&plan, r);
762 assert_eq!(group.members, vec![r, s]);
763 assert_eq!(group.external_inputs, vec![x]);
765
766 let tensor = 16u64 * 4; assert_eq!(plan.bytes_unfused, 6 * tensor);
769 assert_eq!(plan.bytes_fused, 3 * tensor);
771 assert_eq!(plan.bytes_saved, 3 * tensor);
772 assert!((plan.speedup_estimate - 2.0).abs() < 1e-9);
773 }
774
775 #[test]
776 fn hand_computed_bytes_exact() {
777 let mut graph = FusionGraph::new();
778 let a = graph.add_op(OpKind::Mul, vec![], vec![10], 4); let b = graph.add_op(OpKind::Add, vec![a], vec![10], 4);
780
781 let plan = FusionPlanner::new()
782 .plan(&graph)
783 .expect("two-op chain must plan");
784
785 assert_eq!(plan.groups.len(), 1);
786 let group = &plan.groups[0];
787 assert!(group.external_inputs.is_empty());
788 assert_eq!(group.external_outputs, vec![b]);
789 assert_eq!(group.internal_intermediates, vec![a]);
790 assert_eq!(group.bytes_read, 0);
791 assert_eq!(group.bytes_written, 40);
792
793 assert_eq!(plan.bytes_unfused, 120);
795 assert_eq!(plan.bytes_fused, 40);
797 assert_eq!(plan.bytes_saved, 80);
798 assert!((plan.speedup_estimate - 3.0).abs() < 1e-9);
799 }
800
801 #[test]
802 fn fusion_avoids_introducing_cycle() {
803 let mut graph = FusionGraph::new();
806 let a = graph.add_op(OpKind::Relu, vec![], vec![16], 4);
807 let b = graph.add_op(OpKind::Sigmoid, vec![a], vec![16], 4);
808 let c = graph.add_op(OpKind::MatMul, vec![b], vec![16], 4); let d = graph.add_op(OpKind::Add, vec![a, c], vec![16], 4);
810
811 let plan = FusionPlanner::new().plan(&graph).expect("graph must plan");
812
813 assert_eq!(plan.groups.len(), 3);
814 assert_eq!(find_group(&plan, a).members, vec![a, b]);
815 assert_eq!(find_group(&plan, c).members, vec![c]);
816 assert_eq!(find_group(&plan, d).members, vec![d]);
817 assert!(find_group(&plan, a).external_outputs.contains(&a));
819 }
820
821 #[test]
822 fn broadcast_compatible_edge_fuses() {
823 let mut graph = FusionGraph::new();
824 let a = graph.add_op(OpKind::Relu, vec![], vec![1], 4); let b = graph.add_op(OpKind::Add, vec![a], vec![32], 4); let plan = FusionPlanner::new()
828 .plan(&graph)
829 .expect("broadcast chain must plan");
830 assert_eq!(plan.groups.len(), 1);
831 assert_eq!(plan.groups[0].members, vec![a, b]);
832
833 let strict = FusionPlanner::new()
835 .with_broadcast(false)
836 .plan(&graph)
837 .expect("strict planner must plan");
838 assert_eq!(strict.groups.len(), 2);
839 }
840
841 #[test]
842 fn bytes_saved_and_speedup_invariants() {
843 let mut graph = FusionGraph::new();
844 let a = graph.add_op(OpKind::Scale, vec![], vec![64, 64], 4);
845 let b = graph.add_op(OpKind::Relu, vec![a], vec![64, 64], 4);
846 let c = graph.add_op(OpKind::Sigmoid, vec![b], vec![64, 64], 4);
847 graph.add_op(OpKind::Tanh, vec![c], vec![64, 64], 4);
848
849 let plan = FusionPlanner::new()
850 .plan(&graph)
851 .expect("fusible chain must plan");
852
853 assert_eq!(plan.groups.len(), 1);
854 assert_eq!(plan.bytes_saved, plan.bytes_unfused - plan.bytes_fused);
855 assert!(plan.speedup_estimate >= 1.0);
856 assert!(plan.bytes_fused < plan.bytes_unfused);
857 }
858
859 #[test]
860 fn cyclic_graph_rejected() {
861 let mut graph = FusionGraph::new();
862 let _a = graph.add_op(OpKind::Add, vec![1], vec![8], 4); let _b = graph.add_op(OpKind::Add, vec![0], vec![8], 4); let error = graph
866 .validate()
867 .expect_err("a cyclic graph must be rejected");
868 assert!(matches!(error, GpuOptimError::InvalidState(_)));
869 assert!(FusionPlanner::new().plan(&graph).is_err());
870 }
871
872 #[test]
873 fn dangling_input_rejected() {
874 let mut graph = FusionGraph::new();
875 let _a = graph.add_op(OpKind::Relu, vec![99], vec![8], 4); let error = graph
878 .validate()
879 .expect_err("a dangling input must be rejected");
880 assert!(matches!(error, GpuOptimError::InvalidState(_)));
881 }
882
883 #[test]
884 fn zero_dtype_rejected() {
885 let mut graph = FusionGraph::new();
886 let _a = graph.add_op(OpKind::Relu, vec![], vec![8], 0);
887 assert!(graph.validate().is_err());
888 }
889
890 #[test]
891 fn empty_graph_is_valid() {
892 let graph = FusionGraph::new();
893 assert!(graph.is_empty());
894 let plan = FusionPlanner::new()
895 .plan(&graph)
896 .expect("empty graph must plan");
897 assert_eq!(plan.num_groups(), 0);
898 assert_eq!(plan.bytes_unfused, 0);
899 assert_eq!(plan.bytes_fused, 0);
900 assert_eq!(plan.bytes_saved, 0);
901 assert!((plan.speedup_estimate - 1.0).abs() < 1e-9);
902 }
903}