1use std::collections::{HashMap, HashSet, VecDeque};
76
77use crate::analysis::{dominance_analyse, topo_analyse};
78use crate::error::{GraphError, GraphResult};
79use crate::graph::ComputeGraph;
80use crate::node::{BufferId, GraphNode, KernelConfig, NodeId, NodeKind};
81
82#[derive(Debug, Clone, Copy, PartialEq, Eq)]
93pub enum ReductionPattern {
94 LayerNorm,
96 Softmax,
98 Generic,
100}
101
102impl ReductionPattern {
103 #[must_use]
105 pub fn name(self) -> &'static str {
106 match self {
107 Self::LayerNorm => "layernorm",
108 Self::Softmax => "softmax",
109 Self::Generic => "reduction",
110 }
111 }
112}
113
114impl std::fmt::Display for ReductionPattern {
115 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
116 f.write_str(self.name())
117 }
118}
119
120#[derive(Debug, Clone, PartialEq, Eq)]
129pub struct ReductionFusionGroup {
130 pub id: usize,
132 pub root: NodeId,
134 pub sink: NodeId,
136 pub members: Vec<NodeId>,
138 pub pattern: ReductionPattern,
140 pub config: KernelConfig,
142 pub tag: String,
144}
145
146impl ReductionFusionGroup {
147 #[must_use]
149 pub fn size(&self) -> usize {
150 self.members.len()
151 }
152
153 #[must_use]
156 pub fn launches_saved(&self) -> usize {
157 self.members.len().saturating_sub(1)
158 }
159}
160
161#[derive(Debug, Clone, Default)]
167pub struct ReductionFusionPlan {
168 pub groups: Vec<ReductionFusionGroup>,
170 pub node_to_group: HashMap<NodeId, usize>,
172}
173
174impl ReductionFusionPlan {
175 #[must_use]
177 pub fn fusion_count(&self) -> usize {
178 self.groups.len()
179 }
180
181 #[must_use]
184 pub fn nodes_saved(&self) -> usize {
185 self.groups
186 .iter()
187 .map(ReductionFusionGroup::launches_saved)
188 .sum()
189 }
190
191 #[must_use]
193 pub fn group_of(&self, node: NodeId) -> Option<&ReductionFusionGroup> {
194 self.node_to_group
195 .get(&node)
196 .and_then(|&idx| self.groups.get(idx))
197 }
198
199 #[must_use]
202 pub fn is_absorbed(&self, node: NodeId) -> bool {
203 match self.group_of(node) {
204 Some(g) => g.root != node,
205 None => false,
206 }
207 }
208}
209
210fn kernel_meta(graph: &ComputeGraph, node: NodeId) -> Option<(bool, KernelConfig)> {
216 match &graph.node(node).ok()?.kind {
217 NodeKind::KernelLaunch {
218 fusible, config, ..
219 } => Some((*fusible, *config)),
220 _ => None,
221 }
222}
223
224fn fn_name_lower(graph: &ComputeGraph, node: NodeId) -> String {
226 graph
227 .node(node)
228 .ok()
229 .and_then(|n| n.kind.function_name())
230 .unwrap_or("")
231 .to_ascii_lowercase()
232}
233
234fn configs_compatible(a: &KernelConfig, b: &KernelConfig) -> bool {
236 a.total_threads() == b.total_threads()
237}
238
239fn classify(graph: &ComputeGraph, members: &[NodeId]) -> ReductionPattern {
244 let names: Vec<String> = members.iter().map(|&m| fn_name_lower(graph, m)).collect();
245 let has = |needle: &str| names.iter().any(|n| n.contains(needle));
246
247 let softmax_like = has("exp") && (has("softmax") || has("div") || has("sum") || has("norm"));
249 let layernorm_like = (has("mean") || has("avg"))
251 && (has("var") || has("std") || has("rms") || has("norm") || has("layernorm"));
252
253 if has("softmax") || softmax_like {
254 ReductionPattern::Softmax
255 } else if has("layernorm") || layernorm_like {
256 ReductionPattern::LayerNorm
257 } else {
258 ReductionPattern::Generic
259 }
260}
261
262fn grow_region(
272 graph: &ComputeGraph,
273 root: NodeId,
274 dt: &crate::analysis::DomTree,
275 topo_pos: &HashMap<NodeId, usize>,
276 claimed: &HashSet<NodeId>,
277) -> GraphResult<Option<Vec<NodeId>>> {
278 let root_config = match kernel_meta(graph, root) {
279 Some((true, cfg)) => cfg,
280 _ => return Ok(None),
281 };
282
283 let member_ok = |n: NodeId| -> bool {
289 if n == root {
290 return true;
291 }
292 if claimed.contains(&n) {
293 return false;
294 }
295 if !dt.dominates(root, n) {
296 return false;
297 }
298 match kernel_meta(graph, n) {
299 Some((fusible, cfg)) => fusible && configs_compatible(&root_config, &cfg),
300 None => false,
301 }
302 };
303
304 let mut region: HashSet<NodeId> = HashSet::new();
306 region.insert(root);
307 let mut queue: VecDeque<NodeId> = VecDeque::new();
308 queue.push_back(root);
309 while let Some(cur) = queue.pop_front() {
310 for &succ in graph.successors(cur)? {
311 if region.contains(&succ) {
312 continue;
313 }
314 if member_ok(succ) {
315 region.insert(succ);
316 queue.push_back(succ);
317 }
318 }
319 }
320
321 if region.len() < 3 {
322 return Ok(None);
325 }
326
327 let mut exits: Vec<NodeId> = Vec::new();
331 for &m in ®ion {
332 let leaves = graph.successors(m)?.iter().any(|s| !region.contains(s));
333 let is_graph_sink = graph.successors(m)?.is_empty();
334 if leaves || is_graph_sink {
335 exits.push(m);
336 }
337 }
338
339 let sink = *region
349 .iter()
350 .max_by_key(|&&m| topo_pos.get(&m).copied().unwrap_or(0))
351 .ok_or_else(|| GraphError::Internal("reduction region unexpectedly empty".into()))?;
352
353 if !region_reaches_all(graph, ®ion, sink)? {
356 return Ok(None);
357 }
358
359 for &m in ®ion {
363 if m == sink {
364 continue;
365 }
366 let leaks = graph.successors(m)?.iter().any(|s| !region.contains(s));
367 if leaks {
368 return Ok(None);
369 }
370 }
371
372 let has_fanout = region
375 .iter()
376 .try_fold(false, |acc, &m| -> GraphResult<bool> {
377 if acc {
378 return Ok(true);
379 }
380 let in_region_succ = graph
381 .successors(m)?
382 .iter()
383 .filter(|s| region.contains(s))
384 .count();
385 Ok(in_region_succ >= 2)
386 })?;
387 if !has_fanout {
388 return Ok(None);
389 }
390
391 let mut members: Vec<NodeId> = region.into_iter().collect();
393 members.sort_by_key(|m| topo_pos.get(m).copied().unwrap_or(usize::MAX));
394 Ok(Some(members))
395}
396
397fn region_reaches_all(
400 graph: &ComputeGraph,
401 region: &HashSet<NodeId>,
402 sink: NodeId,
403) -> GraphResult<bool> {
404 let mut reached: HashSet<NodeId> = HashSet::new();
407 reached.insert(sink);
408 let mut queue: VecDeque<NodeId> = VecDeque::new();
409 queue.push_back(sink);
410 while let Some(cur) = queue.pop_front() {
411 for &pred in graph.predecessors(cur)? {
412 if region.contains(&pred) && reached.insert(pred) {
413 queue.push_back(pred);
414 }
415 }
416 }
417 Ok(region.iter().all(|m| reached.contains(m)))
418}
419
420pub fn analyse(graph: &ComputeGraph) -> GraphResult<ReductionFusionPlan> {
434 if graph.is_empty() {
435 return Err(GraphError::EmptyGraph);
436 }
437
438 let topo = topo_analyse(graph)?;
439 let dt = dominance_analyse(graph)?;
440 let topo_pos: HashMap<NodeId, usize> = topo
441 .order
442 .iter()
443 .enumerate()
444 .map(|(p, &id)| (id, p))
445 .collect();
446
447 let mut claimed: HashSet<NodeId> = HashSet::new();
448 let mut groups: Vec<ReductionFusionGroup> = Vec::new();
449 let mut node_to_group: HashMap<NodeId, usize> = HashMap::new();
450
451 for &root in &topo.order {
452 if claimed.contains(&root) {
453 continue;
454 }
455 match kernel_meta(graph, root) {
457 Some((true, _)) => {}
458 _ => continue,
459 }
460
461 let members = match grow_region(graph, root, &dt, &topo_pos, &claimed)? {
462 Some(m) => m,
463 None => continue,
464 };
465
466 let sink = *members.last().ok_or_else(|| {
467 GraphError::Internal("reduction region members unexpectedly empty".into())
468 })?;
469 let pattern = classify(graph, &members);
470 let config = kernel_meta(graph, root)
471 .map(|(_, c)| c)
472 .unwrap_or_else(|| KernelConfig::linear(1, 1, 0));
473
474 let gid = groups.len();
475 let tag = format!(
476 "fused_{}_{}..{}",
477 pattern.name(),
478 graph.node(root)?.display_name(),
479 graph.node(sink)?.display_name()
480 );
481
482 for &m in &members {
483 claimed.insert(m);
484 node_to_group.insert(m, gid);
485 }
486
487 groups.push(ReductionFusionGroup {
488 id: gid,
489 root,
490 sink,
491 members,
492 pattern,
493 config,
494 tag,
495 });
496 }
497
498 Ok(ReductionFusionPlan {
499 groups,
500 node_to_group,
501 })
502}
503
504pub fn rewrite(graph: &ComputeGraph, plan: &ReductionFusionPlan) -> GraphResult<ComputeGraph> {
528 let mut out = ComputeGraph::new();
529
530 for buf in graph.buffers() {
532 out.add_buffer(buf.clone());
533 }
534
535 let mut old_to_new: HashMap<NodeId, NodeId> = HashMap::new();
538
539 for old in graph.nodes() {
542 let oid = old.id;
543 if let Some(group) = plan.group_of(oid) {
544 if group.root != oid {
545 continue;
547 }
548 let region: HashSet<NodeId> = group.members.iter().copied().collect();
550
551 let mut region_outputs: HashSet<BufferId> = HashSet::new();
555 for &m in &group.members {
556 for &b in &graph.node(m)?.outputs {
557 region_outputs.insert(b);
558 }
559 }
560 let mut fused_inputs: Vec<BufferId> = Vec::new();
561 let mut seen_in: HashSet<BufferId> = HashSet::new();
562 for &m in &group.members {
563 for &b in &graph.node(m)?.inputs {
564 if !region_outputs.contains(&b) && seen_in.insert(b) {
566 fused_inputs.push(b);
567 }
568 }
569 }
570 let fused_outputs: Vec<BufferId> = graph.node(group.sink)?.outputs.clone();
572
573 let fn_name = format!(
574 "{}_{}",
575 group.pattern.name(),
576 group
577 .members
578 .iter()
579 .filter_map(|&m| graph.node(m).ok().and_then(|n| n.kind.function_name()))
580 .collect::<Vec<_>>()
581 .join("_")
582 );
583 let cost: u64 = group
584 .members
585 .iter()
586 .filter_map(|&m| graph.node(m).ok().map(|n| n.cost_hint))
587 .sum();
588 let kind = NodeKind::KernelLaunch {
589 function_name: fn_name,
590 config: group.config,
591 fusible: true,
592 };
593 let node = GraphNode::new(NodeId(0), kind)
594 .with_inputs(fused_inputs)
595 .with_outputs(fused_outputs)
596 .with_cost(cost.max(1))
597 .with_name(group.tag.clone());
598 let nid = out.add_node(node);
599 for &m in ®ion {
600 old_to_new.insert(m, nid);
601 }
602 } else {
603 let mut node = GraphNode::new(NodeId(0), old.kind.clone())
605 .with_inputs(old.inputs.iter().copied())
606 .with_outputs(old.outputs.iter().copied())
607 .with_cost(old.cost_hint);
608 if let Some(s) = old.stream_hint {
609 node = node.with_stream(s);
610 }
611 if let Some(name) = &old.name {
612 node = node.with_name(name.clone());
613 }
614 let nid = out.add_node(node);
615 old_to_new.insert(oid, nid);
616 }
617 }
618
619 let mut added: HashSet<(NodeId, NodeId)> = HashSet::new();
621 for (from_old, to_old) in graph.edges() {
622 let from_new = *old_to_new
623 .get(&from_old)
624 .ok_or_else(|| GraphError::Internal("missing node mapping (from)".into()))?;
625 let to_new = *old_to_new
626 .get(&to_old)
627 .ok_or_else(|| GraphError::Internal("missing node mapping (to)".into()))?;
628 if from_new == to_new {
629 continue;
631 }
632 if added.insert((from_new, to_new)) {
633 out.add_edge(from_new, to_new)?;
634 }
635 }
636
637 Ok(out)
638}
639
640#[cfg(test)]
645mod tests {
646 use super::*;
647 use crate::builder::GraphBuilder;
648 use crate::executor::{ExecutionPlan, SequentialExecutor};
649 use crate::node::MemcpyDir;
650
651 fn build_layernorm() -> (ComputeGraph, Vec<NodeId>) {
660 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
661 let mean = b.add_kernel("mean", 4, 256, 0).fusible(true).finish();
662 let sub = b.add_kernel("subtract", 4, 256, 0).fusible(true).finish();
663 let var = b.add_kernel("variance", 4, 256, 0).fusible(true).finish();
664 let norm = b.add_kernel("normalize", 4, 256, 0).fusible(true).finish();
665 let scale = b
666 .add_kernel("scale_shift", 4, 256, 0)
667 .fusible(true)
668 .finish();
669
670 b.dep(mean, sub);
672 b.dep(sub, var);
673 b.dep(sub, norm); b.dep(var, norm);
675 b.dep(norm, scale);
676 let g = b.build().expect("layernorm graph builds");
677 (g, vec![mean, sub, var, norm, scale])
678 }
679
680 fn build_softmax() -> (ComputeGraph, Vec<NodeId>) {
687 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
688 let mx = b.add_kernel("max", 4, 256, 0).fusible(true).finish();
689 let sub = b.add_kernel("subtract", 4, 256, 0).fusible(true).finish();
690 let exp = b.add_kernel("exp", 4, 256, 0).fusible(true).finish();
691 let sum = b.add_kernel("sum", 4, 256, 0).fusible(true).finish();
692 let div = b.add_kernel("divide", 4, 256, 0).fusible(true).finish();
693
694 b.dep(mx, sub);
695 b.dep(sub, exp);
696 b.dep(exp, sum);
697 b.dep(exp, div); b.dep(sum, div);
699 let g = b.build().expect("softmax graph builds");
700 (g, vec![mx, sub, exp, sum, div])
701 }
702
703 #[test]
706 fn reduction_empty_graph_errors() {
707 let g = ComputeGraph::new();
708 assert!(matches!(analyse(&g), Err(GraphError::EmptyGraph)));
709 }
710
711 #[test]
712 fn layernorm_region_detected() {
713 let (g, ids) = build_layernorm();
714 let plan = analyse(&g).expect("reduction fusion analysis succeeds");
715 assert_eq!(plan.fusion_count(), 1);
716 let group = plan
717 .group_of(ids[0])
718 .expect("mean belongs to a fused region");
719 assert_eq!(group.size(), 5);
721 for id in &ids {
722 assert!(group.members.contains(id), "member {id} missing");
723 }
724 assert_eq!(group.root, ids[0]); assert_eq!(group.sink, ids[4]); assert_eq!(group.pattern, ReductionPattern::LayerNorm);
727 assert_eq!(plan.nodes_saved(), 4);
729 }
730
731 #[test]
732 fn softmax_region_detected() {
733 let (g, ids) = build_softmax();
734 let plan = analyse(&g).expect("reduction fusion analysis succeeds");
735 assert_eq!(plan.fusion_count(), 1);
736 let group = plan
737 .group_of(ids[2])
738 .expect("exp belongs to a fused region");
739 assert_eq!(group.size(), 5);
740 assert_eq!(group.root, ids[0]); assert_eq!(group.sink, ids[4]); assert_eq!(group.pattern, ReductionPattern::Softmax);
743 assert_eq!(plan.nodes_saved(), 4);
744 }
745
746 #[test]
747 fn absorbed_members_flagged() {
748 let (g, ids) = build_softmax();
749 let plan = analyse(&g).expect("reduction fusion analysis succeeds");
750 assert!(!plan.is_absorbed(ids[0]));
752 for id in &ids[1..] {
753 assert!(plan.is_absorbed(*id), "member {id} should be absorbed");
754 }
755 }
756
757 #[test]
760 fn linear_chain_not_a_reduction_region() {
761 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
764 let k0 = b.add_kernel("a", 4, 256, 0).fusible(true).finish();
765 let k1 = b.add_kernel("b", 4, 256, 0).fusible(true).finish();
766 let k2 = b.add_kernel("c", 4, 256, 0).fusible(true).finish();
767 b.chain(&[k0, k1, k2]);
768 let g = b.build().expect("chain graph builds");
769 let plan = analyse(&g).expect("reduction fusion analysis succeeds");
770 assert_eq!(plan.fusion_count(), 0);
771 assert_eq!(plan.nodes_saved(), 0);
772 }
773
774 #[test]
775 fn non_fusible_member_breaks_region() {
776 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
779 let mx = b.add_kernel("max", 4, 256, 0).fusible(true).finish();
780 let sub = b.add_kernel("subtract", 4, 256, 0).fusible(true).finish();
781 let exp = b.add_kernel("exp", 4, 256, 0).fusible(false).finish();
782 let sum = b.add_kernel("sum", 4, 256, 0).fusible(true).finish();
783 let div = b.add_kernel("divide", 4, 256, 0).fusible(true).finish();
784 b.dep(mx, sub);
785 b.dep(sub, exp);
786 b.dep(exp, sum);
787 b.dep(exp, div);
788 b.dep(sum, div);
789 let g = b.build().expect("graph builds");
790 let plan = analyse(&g).expect("reduction fusion analysis succeeds");
791 assert_eq!(plan.fusion_count(), 0);
792 }
793
794 #[test]
795 fn open_region_leaking_intermediate_not_fused() {
796 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
804 let mx = b.add_kernel("max", 4, 256, 0).fusible(true).finish();
805 let sub = b.add_kernel("subtract", 4, 256, 0).fusible(true).finish();
806 let exp = b.add_kernel("exp", 4, 256, 0).fusible(true).finish();
807 let sum = b.add_kernel("sum", 4, 256, 0).fusible(true).finish();
808 let div = b.add_kernel("divide", 4, 256, 0).fusible(true).finish();
809 let leak = b.add_memcpy("leak", MemcpyDir::DeviceToHost, 1024);
811 b.dep(mx, sub);
812 b.dep(sub, exp);
813 b.dep(exp, sum);
814 b.dep(exp, div);
815 b.dep(sum, div);
816 b.dep(exp, leak); let g = b.build().expect("graph builds");
818 let plan = analyse(&g).expect("reduction fusion analysis succeeds");
819 assert_eq!(plan.fusion_count(), 0);
820 }
821
822 #[test]
823 fn incompatible_config_member_excluded() {
824 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
827 let mx = b.add_kernel("max", 4, 256, 0).fusible(true).finish();
828 let sub = b.add_kernel("subtract", 4, 256, 0).fusible(true).finish();
829 let exp = b.add_kernel("exp", 4, 256, 0).fusible(true).finish();
830 let sum = b.add_kernel("sum", 4, 256, 0).fusible(true).finish();
831 let div = b.add_kernel("divide", 8, 256, 0).fusible(true).finish(); b.dep(mx, sub);
833 b.dep(sub, exp);
834 b.dep(exp, sum);
835 b.dep(exp, div);
836 b.dep(sum, div);
837 let g = b.build().expect("graph builds");
838 let plan = analyse(&g).expect("reduction fusion analysis succeeds");
839 assert_eq!(plan.fusion_count(), 0);
842 }
843
844 #[test]
847 fn rewrite_collapses_region_to_one_node() {
848 let (g, _ids) = build_layernorm();
849 let plan = analyse(&g).expect("analysis succeeds");
850 let fused = rewrite(&g, &plan).expect("rewrite succeeds");
851 assert_eq!(g.node_count(), 5);
853 assert_eq!(fused.node_count(), 1);
854 let only = fused.node(NodeId(0)).expect("fused node exists");
856 assert!(only.kind.is_compute());
857 assert!(
858 only.kind
859 .function_name()
860 .unwrap_or("")
861 .contains("layernorm")
862 );
863 }
864
865 #[test]
866 fn rewrite_preserves_external_topology() {
867 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
870 let up = b.add_memcpy("up", MemcpyDir::HostToDevice, 1024);
871 let mx = b.add_kernel("max", 4, 256, 0).fusible(true).finish();
872 let sub = b.add_kernel("subtract", 4, 256, 0).fusible(true).finish();
873 let exp = b.add_kernel("exp", 4, 256, 0).fusible(true).finish();
874 let sum = b.add_kernel("sum", 4, 256, 0).fusible(true).finish();
875 let div = b.add_kernel("divide", 4, 256, 0).fusible(true).finish();
876 let dn = b.add_memcpy("dn", MemcpyDir::DeviceToHost, 1024);
877 b.dep(up, mx);
878 b.dep(mx, sub);
879 b.dep(sub, exp);
880 b.dep(exp, sum);
881 b.dep(exp, div);
882 b.dep(sum, div);
883 b.dep(div, dn);
884 let g = b.build().expect("graph builds");
885 let plan = analyse(&g).expect("analysis succeeds");
886 assert_eq!(plan.fusion_count(), 1);
887 let fused = rewrite(&g, &plan).expect("rewrite succeeds");
888 assert_eq!(fused.node_count(), 3);
890 let up_new = fused.sources();
892 assert_eq!(up_new.len(), 1);
893 let dn_new = fused.sinks();
894 assert_eq!(dn_new.len(), 1);
895 assert!(fused.is_reachable(up_new[0], dn_new[0]));
896 assert_eq!(fused.kernel_nodes().len(), 1);
898 }
899
900 #[test]
901 fn rewrite_no_match_is_identity() {
902 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
904 let k0 = b.add_kernel("a", 4, 256, 0).fusible(true).finish();
905 let k1 = b.add_kernel("b", 4, 256, 0).fusible(true).finish();
906 b.chain(&[k0, k1]);
907 let g = b.build().expect("graph builds");
908 let plan = analyse(&g).expect("analysis succeeds");
909 let fused = rewrite(&g, &plan).expect("rewrite succeeds");
910 assert_eq!(fused.node_count(), g.node_count());
911 assert_eq!(fused.edge_count(), g.edge_count());
912 }
913
914 #[test]
915 fn simulator_agrees_before_and_after_fusion() {
916 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
921 let up = b.add_memcpy("up", MemcpyDir::HostToDevice, 4096);
922 let mx = b.add_kernel("max", 4, 256, 0).fusible(true).finish();
923 let sub = b.add_kernel("subtract", 4, 256, 0).fusible(true).finish();
924 let exp = b.add_kernel("exp", 4, 256, 0).fusible(true).finish();
925 let sum = b.add_kernel("sum", 4, 256, 0).fusible(true).finish();
926 let div = b.add_kernel("divide", 4, 256, 0).fusible(true).finish();
927 let dn = b.add_memcpy("dn", MemcpyDir::DeviceToHost, 4096);
928 b.dep(up, mx);
929 b.dep(mx, sub);
930 b.dep(sub, exp);
931 b.dep(exp, sum);
932 b.dep(exp, div);
933 b.dep(sum, div);
934 b.dep(div, dn);
935 let g = b.build().expect("graph builds");
936
937 let plan = analyse(&g).expect("analysis succeeds");
938 assert_eq!(plan.fusion_count(), 1);
939 let fused = rewrite(&g, &plan).expect("rewrite succeeds");
940
941 let before =
942 SequentialExecutor::new(&ExecutionPlan::build(&g, 4).expect("plan(before) builds"))
943 .run()
944 .expect("before runs");
945 let after =
946 SequentialExecutor::new(&ExecutionPlan::build(&fused, 4).expect("plan(after) builds"))
947 .run()
948 .expect("after runs");
949
950 assert_eq!(before.bytes_copied, after.bytes_copied);
953 assert_eq!(before.bytes_copied, 4096 * 2);
954 assert_eq!(before.bytes_set, after.bytes_set);
955 assert_eq!(after.kernels_launched, 1);
957 assert!(after.kernels_launched <= before.kernels_launched);
960 }
961
962 #[test]
963 fn pattern_display_and_name() {
964 assert_eq!(ReductionPattern::LayerNorm.name(), "layernorm");
965 assert_eq!(ReductionPattern::Softmax.to_string(), "softmax");
966 assert_eq!(ReductionPattern::Generic.name(), "reduction");
967 }
968
969 #[test]
970 fn generic_reduction_region_classified() {
971 let mut b = GraphBuilder::new().with_auto_infer_edges(false);
973 let r = b.add_kernel("reduce_op", 4, 256, 0).fusible(true).finish();
974 let a = b.add_kernel("elemwise_a", 4, 256, 0).fusible(true).finish();
975 let c = b.add_kernel("elemwise_c", 4, 256, 0).fusible(true).finish();
976 let join = b.add_kernel("combine", 4, 256, 0).fusible(true).finish();
977 b.dep(r, a);
979 b.dep(r, c);
980 b.dep(a, join);
981 b.dep(c, join);
982 let g = b.build().expect("graph builds");
983 let plan = analyse(&g).expect("analysis succeeds");
984 assert_eq!(plan.fusion_count(), 1);
985 let group = plan.group_of(r).expect("r in a region");
986 assert_eq!(group.pattern, ReductionPattern::Generic);
987 assert_eq!(group.size(), 4);
988 }
989}