1use std::collections::{HashMap, HashSet, VecDeque};
18
19use crate::error::{GraphError, GraphResult};
20use crate::node::{BufferDescriptor, BufferId, GraphNode, NodeId};
21
22#[derive(Debug, Clone)]
40pub struct ComputeGraph {
41 nodes: Vec<GraphNode>,
43 successors: Vec<Vec<NodeId>>,
45 predecessors: Vec<Vec<NodeId>>,
47 buffers: Vec<BufferDescriptor>,
49 next_node: u32,
51 next_buf: u32,
53}
54
55impl Default for ComputeGraph {
56 fn default() -> Self {
57 Self::new()
58 }
59}
60
61impl ComputeGraph {
62 #[must_use]
64 pub fn new() -> Self {
65 Self {
66 nodes: Vec::new(),
67 successors: Vec::new(),
68 predecessors: Vec::new(),
69 buffers: Vec::new(),
70 next_node: 0,
71 next_buf: 0,
72 }
73 }
74
75 pub fn add_node(&mut self, mut node: GraphNode) -> NodeId {
83 let id = NodeId(self.next_node);
84 node.id = id;
85 self.next_node += 1;
86 self.nodes.push(node);
87 self.successors.push(Vec::new());
88 self.predecessors.push(Vec::new());
89 id
90 }
91
92 pub fn node(&self, id: NodeId) -> GraphResult<&GraphNode> {
94 self.nodes
95 .get(id.0 as usize)
96 .ok_or(GraphError::NodeNotFound(id))
97 }
98
99 pub fn node_mut(&mut self, id: NodeId) -> GraphResult<&mut GraphNode> {
101 self.nodes
102 .get_mut(id.0 as usize)
103 .ok_or(GraphError::NodeNotFound(id))
104 }
105
106 #[inline]
108 pub fn nodes(&self) -> &[GraphNode] {
109 &self.nodes
110 }
111
112 #[inline]
114 pub fn node_count(&self) -> usize {
115 self.nodes.len()
116 }
117
118 pub fn add_buffer(&mut self, mut buf: BufferDescriptor) -> BufferId {
126 let id = BufferId(self.next_buf);
127 buf.id = id;
128 self.next_buf += 1;
129 self.buffers.push(buf);
130 id
131 }
132
133 pub fn buffer(&self, id: BufferId) -> GraphResult<&BufferDescriptor> {
135 self.buffers
136 .get(id.0 as usize)
137 .ok_or(GraphError::NodeNotFound(NodeId(id.0)))
138 }
139
140 #[inline]
142 pub fn buffers(&self) -> &[BufferDescriptor] {
143 &self.buffers
144 }
145
146 #[inline]
148 pub fn buffer_count(&self) -> usize {
149 self.buffers.len()
150 }
151
152 pub fn add_edge(&mut self, from: NodeId, to: NodeId) -> GraphResult<()> {
165 let n = self.nodes.len();
166 if from.0 as usize >= n {
167 return Err(GraphError::NodeNotFound(from));
168 }
169 if to.0 as usize >= n {
170 return Err(GraphError::NodeNotFound(to));
171 }
172 if from == to {
173 return Err(GraphError::CycleDetected { from, to });
174 }
175 if self.is_reachable(to, from) {
178 return Err(GraphError::CycleDetected { from, to });
179 }
180 if !self.successors[from.0 as usize].contains(&to) {
182 self.successors[from.0 as usize].push(to);
183 self.predecessors[to.0 as usize].push(from);
184 }
185 Ok(())
186 }
187
188 pub fn is_reachable(&self, src: NodeId, dst: NodeId) -> bool {
192 if src == dst {
193 return true;
194 }
195 let n = self.nodes.len();
196 let mut visited = vec![false; n];
197 let mut queue = VecDeque::new();
198 queue.push_back(src);
199 visited[src.0 as usize] = true;
200 while let Some(curr) = queue.pop_front() {
201 for &next in &self.successors[curr.0 as usize] {
202 if next == dst {
203 return true;
204 }
205 if !visited[next.0 as usize] {
206 visited[next.0 as usize] = true;
207 queue.push_back(next);
208 }
209 }
210 }
211 false
212 }
213
214 pub fn successors(&self, id: NodeId) -> GraphResult<&[NodeId]> {
216 if id.0 as usize >= self.nodes.len() {
217 return Err(GraphError::NodeNotFound(id));
218 }
219 Ok(&self.successors[id.0 as usize])
220 }
221
222 pub fn predecessors(&self, id: NodeId) -> GraphResult<&[NodeId]> {
224 if id.0 as usize >= self.nodes.len() {
225 return Err(GraphError::NodeNotFound(id));
226 }
227 Ok(&self.predecessors[id.0 as usize])
228 }
229
230 pub fn edge_count(&self) -> usize {
232 self.successors.iter().map(|v| v.len()).sum()
233 }
234
235 pub fn edges(&self) -> Vec<(NodeId, NodeId)> {
237 let mut edges = Vec::new();
238 for (i, succs) in self.successors.iter().enumerate() {
239 for &to in succs {
240 edges.push((NodeId(i as u32), to));
241 }
242 }
243 edges
244 }
245
246 pub fn sources(&self) -> Vec<NodeId> {
252 self.predecessors
253 .iter()
254 .enumerate()
255 .filter(|(_, preds)| preds.is_empty())
256 .map(|(i, _)| NodeId(i as u32))
257 .collect()
258 }
259
260 pub fn sinks(&self) -> Vec<NodeId> {
262 self.successors
263 .iter()
264 .enumerate()
265 .filter(|(_, succs)| succs.is_empty())
266 .map(|(i, _)| NodeId(i as u32))
267 .collect()
268 }
269
270 pub fn topological_order(&self) -> GraphResult<Vec<NodeId>> {
284 if self.nodes.is_empty() {
285 return Err(GraphError::EmptyGraph);
286 }
287 let n = self.nodes.len();
288 let mut in_degree: Vec<u32> = self.predecessors.iter().map(|p| p.len() as u32).collect();
289 let mut queue: VecDeque<NodeId> = (0..n)
290 .filter(|&i| in_degree[i] == 0)
291 .map(|i| NodeId(i as u32))
292 .collect();
293 let mut order = Vec::with_capacity(n);
294 while let Some(id) = queue.pop_front() {
295 order.push(id);
296 for &succ in &self.successors[id.0 as usize] {
297 let d = &mut in_degree[succ.0 as usize];
298 *d -= 1;
299 if *d == 0 {
300 queue.push_back(succ);
301 }
302 }
303 }
304 debug_assert_eq!(
306 order.len(),
307 n,
308 "topological sort incomplete — internal invariant broken"
309 );
310 Ok(order)
311 }
312
313 pub fn infer_data_edges(&mut self) -> GraphResult<()> {
326 let mut writers: HashMap<BufferId, Vec<NodeId>> = HashMap::new();
328 for node in &self.nodes {
329 for &buf in &node.outputs {
330 writers.entry(buf).or_default().push(node.id);
331 }
332 }
333 let reader_data: Vec<(NodeId, Vec<BufferId>)> = self
335 .nodes
336 .iter()
337 .map(|n| (n.id, n.inputs.clone()))
338 .collect();
339 for (reader_id, inputs) in reader_data {
340 for buf in inputs {
341 if let Some(node_writers) = writers.get(&buf) {
342 for &writer_id in node_writers {
343 if writer_id != reader_id {
344 self.add_edge(writer_id, reader_id)?;
348 }
349 }
350 }
351 }
352 }
353 Ok(())
354 }
355
356 pub fn reachable_from(&self, roots: &[NodeId]) -> HashSet<NodeId> {
362 let mut visited = HashSet::new();
363 let mut stack: Vec<NodeId> = roots.to_vec();
364 while let Some(id) = stack.pop() {
365 if visited.insert(id) {
366 for &s in &self.successors[id.0 as usize] {
367 if !visited.contains(&s) {
368 stack.push(s);
369 }
370 }
371 }
372 }
373 visited
374 }
375
376 pub fn reaching(&self, targets: &[NodeId]) -> HashSet<NodeId> {
378 let mut visited = HashSet::new();
379 let mut stack: Vec<NodeId> = targets.to_vec();
380 while let Some(id) = stack.pop() {
381 if visited.insert(id) {
382 for &p in &self.predecessors[id.0 as usize] {
383 if !visited.contains(&p) {
384 stack.push(p);
385 }
386 }
387 }
388 }
389 visited
390 }
391
392 #[inline]
398 pub fn is_empty(&self) -> bool {
399 self.nodes.is_empty()
400 }
401
402 pub fn critical_path_length(&self) -> GraphResult<usize> {
411 let order = self.topological_order()?;
412 let n = self.nodes.len();
413 let mut dist = vec![0usize; n];
414 let mut max_len = 0usize;
415 for id in &order {
416 let d = dist[id.0 as usize];
417 max_len = max_len.max(d);
418 for &succ in &self.successors[id.0 as usize] {
419 let nd = d + 1;
420 if nd > dist[succ.0 as usize] {
421 dist[succ.0 as usize] = nd;
422 }
423 }
424 }
425 Ok(max_len)
426 }
427
428 pub fn max_in_degree(&self) -> usize {
430 self.predecessors.iter().map(|p| p.len()).max().unwrap_or(0)
431 }
432
433 pub fn max_out_degree(&self) -> usize {
435 self.successors.iter().map(|s| s.len()).max().unwrap_or(0)
436 }
437
438 pub fn parallelism_width(&self) -> GraphResult<usize> {
443 if self.nodes.is_empty() {
446 return Ok(0);
447 }
448 let n = self.nodes.len();
449 let mut level = vec![0usize; n];
450 let mut in_degree: Vec<u32> = self.predecessors.iter().map(|p| p.len() as u32).collect();
451 let mut queue: VecDeque<NodeId> = (0..n)
452 .filter(|&i| in_degree[i] == 0)
453 .map(|i| NodeId(i as u32))
454 .collect();
455 let mut max_width = queue.len();
456 while let Some(id) = queue.pop_front() {
457 for &succ in &self.successors[id.0 as usize] {
458 let nl = level[id.0 as usize] + 1;
459 if nl > level[succ.0 as usize] {
460 level[succ.0 as usize] = nl;
461 }
462 let d = &mut in_degree[succ.0 as usize];
463 *d -= 1;
464 if *d == 0 {
465 queue.push_back(succ);
466 }
467 }
468 }
470 let max_level = *level.iter().max().unwrap_or(&0);
472 let mut width_at_level = vec![0usize; max_level + 1];
473 for &lv in &level {
474 width_at_level[lv] += 1;
475 }
476 max_width = max_width.max(*width_at_level.iter().max().unwrap_or(&0));
477 Ok(max_width)
478 }
479
480 pub fn kernel_nodes(&self) -> Vec<NodeId> {
486 self.nodes
487 .iter()
488 .filter(|n| n.kind.is_compute())
489 .map(|n| n.id)
490 .collect()
491 }
492
493 pub fn fusible_nodes(&self) -> Vec<NodeId> {
495 self.nodes
496 .iter()
497 .filter(|n| n.kind.is_fusible())
498 .map(|n| n.id)
499 .collect()
500 }
501
502 pub fn to_dot(&self) -> String {
508 let mut s = String::from("digraph ComputeGraph {\n rankdir=TB;\n");
509 for node in &self.nodes {
510 let label = node.display_name();
511 let shape = if node.kind.is_compute() {
512 "box"
513 } else if node.kind.is_memory_op() {
514 "parallelogram"
515 } else {
516 "ellipse"
517 };
518 s.push_str(&format!(
519 " {} [label=\"{} ({})\", shape={shape}];\n",
520 node.id.0,
521 label,
522 node.kind.tag()
523 ));
524 }
525 for (i, succs) in self.successors.iter().enumerate() {
526 for &to in succs {
527 s.push_str(&format!(" {} -> {};\n", i, to.0));
528 }
529 }
530 s.push('}');
531 s
532 }
533}
534
535impl std::fmt::Display for ComputeGraph {
536 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
537 write!(
538 f,
539 "ComputeGraph({} nodes, {} edges, {} buffers)",
540 self.node_count(),
541 self.edge_count(),
542 self.buffer_count()
543 )
544 }
545}
546
547#[cfg(test)]
552mod tests {
553 use super::*;
554 use crate::node::{BufferDescriptor, KernelConfig, MemcpyDir, NodeKind};
555
556 fn kernel_node(name: &str) -> GraphNode {
557 GraphNode::new(
558 NodeId(0),
559 NodeKind::KernelLaunch {
560 function_name: name.into(),
561 config: KernelConfig::linear(1, 32, 0),
562 fusible: true,
563 },
564 )
565 }
566
567 fn barrier_node() -> GraphNode {
568 GraphNode::new(NodeId(0), NodeKind::Barrier)
569 }
570
571 fn memcpy_node(dir: MemcpyDir, size: usize) -> GraphNode {
572 GraphNode::new(
573 NodeId(0),
574 NodeKind::Memcpy {
575 dir,
576 size_bytes: size,
577 },
578 )
579 }
580
581 #[test]
584 fn new_graph_is_empty() {
585 let g = ComputeGraph::new();
586 assert!(g.is_empty());
587 assert_eq!(g.node_count(), 0);
588 assert_eq!(g.edge_count(), 0);
589 assert_eq!(g.buffer_count(), 0);
590 }
591
592 #[test]
593 fn default_is_empty() {
594 let g = ComputeGraph::default();
595 assert!(g.is_empty());
596 }
597
598 #[test]
599 fn add_node_assigns_sequential_ids() {
600 let mut g = ComputeGraph::new();
601 let a = g.add_node(barrier_node());
602 let b = g.add_node(barrier_node());
603 let c = g.add_node(barrier_node());
604 assert_eq!(a, NodeId(0));
605 assert_eq!(b, NodeId(1));
606 assert_eq!(c, NodeId(2));
607 assert_eq!(g.node_count(), 3);
608 }
609
610 #[test]
611 fn node_lookup_valid() {
612 let mut g = ComputeGraph::new();
613 let id = g.add_node(kernel_node("add"));
614 assert!(g.node(id).is_ok());
615 assert_eq!(
616 g.node(id)
617 .expect("node registered in graph")
618 .kind
619 .function_name(),
620 Some("add")
621 );
622 }
623
624 #[test]
625 fn node_lookup_invalid() {
626 let g = ComputeGraph::new();
627 assert!(matches!(
628 g.node(NodeId(0)),
629 Err(GraphError::NodeNotFound(_))
630 ));
631 }
632
633 #[test]
634 fn node_mut_allows_modification() {
635 let mut g = ComputeGraph::new();
636 let id = g.add_node(barrier_node());
637 g.node_mut(id).expect("node registered in graph").cost_hint = 42;
638 assert_eq!(g.node(id).expect("node registered in graph").cost_hint, 42);
639 }
640
641 #[test]
644 fn add_buffer_assigns_ids() {
645 let mut g = ComputeGraph::new();
646 let b0 = g.add_buffer(BufferDescriptor::new(BufferId(0), 1024));
647 let b1 = g.add_buffer(BufferDescriptor::new(BufferId(0), 2048));
648 assert_eq!(b0, BufferId(0));
649 assert_eq!(b1, BufferId(1));
650 assert_eq!(g.buffer_count(), 2);
651 }
652
653 #[test]
654 fn buffer_lookup_invalid() {
655 let g = ComputeGraph::new();
656 assert!(g.buffer(BufferId(0)).is_err());
657 }
658
659 #[test]
662 fn add_edge_valid() {
663 let mut g = ComputeGraph::new();
664 let a = g.add_node(kernel_node("a"));
665 let b = g.add_node(kernel_node("b"));
666 assert!(g.add_edge(a, b).is_ok());
667 assert_eq!(g.edge_count(), 1);
668 assert_eq!(g.successors(a).expect("node a registered in graph"), &[b]);
669 assert_eq!(g.predecessors(b).expect("node b registered in graph"), &[a]);
670 }
671
672 #[test]
673 fn add_edge_self_loop_rejected() {
674 let mut g = ComputeGraph::new();
675 let a = g.add_node(barrier_node());
676 assert!(matches!(
677 g.add_edge(a, a),
678 Err(GraphError::CycleDetected { .. })
679 ));
680 }
681
682 #[test]
683 fn add_edge_cycle_rejected() {
684 let mut g = ComputeGraph::new();
685 let a = g.add_node(barrier_node());
686 let b = g.add_node(barrier_node());
687 g.add_edge(a, b).expect("valid DAG edge from a to b");
688 assert!(matches!(
689 g.add_edge(b, a),
690 Err(GraphError::CycleDetected { .. })
691 ));
692 }
693
694 #[test]
695 fn add_edge_invalid_node_rejected() {
696 let mut g = ComputeGraph::new();
697 let a = g.add_node(barrier_node());
698 assert!(matches!(
699 g.add_edge(a, NodeId(99)),
700 Err(GraphError::NodeNotFound(_))
701 ));
702 assert!(matches!(
703 g.add_edge(NodeId(99), a),
704 Err(GraphError::NodeNotFound(_))
705 ));
706 }
707
708 #[test]
709 fn add_edge_duplicate_is_idempotent() {
710 let mut g = ComputeGraph::new();
711 let a = g.add_node(barrier_node());
712 let b = g.add_node(barrier_node());
713 g.add_edge(a, b).expect("valid DAG edge from a to b");
714 g.add_edge(a, b).expect("valid DAG edge from a to b"); assert_eq!(g.edge_count(), 1);
716 }
717
718 #[test]
719 fn edges_returns_all() {
720 let mut g = ComputeGraph::new();
721 let a = g.add_node(barrier_node());
722 let b = g.add_node(barrier_node());
723 let c = g.add_node(barrier_node());
724 g.add_edge(a, b).expect("valid DAG edge from a to b");
725 g.add_edge(a, c).expect("valid DAG edge from a to c");
726 let mut edges = g.edges();
727 edges.sort();
728 assert_eq!(edges, vec![(a, b), (a, c)]);
729 }
730
731 #[test]
734 fn is_reachable_direct() {
735 let mut g = ComputeGraph::new();
736 let a = g.add_node(barrier_node());
737 let b = g.add_node(barrier_node());
738 g.add_edge(a, b).expect("valid DAG edge from a to b");
739 assert!(g.is_reachable(a, b));
740 assert!(!g.is_reachable(b, a));
741 }
742
743 #[test]
744 fn is_reachable_transitive() {
745 let mut g = ComputeGraph::new();
746 let a = g.add_node(barrier_node());
747 let b = g.add_node(barrier_node());
748 let c = g.add_node(barrier_node());
749 g.add_edge(a, b).expect("valid DAG edge from a to b");
750 g.add_edge(b, c).expect("valid DAG edge from b to c");
751 assert!(g.is_reachable(a, c));
752 assert!(!g.is_reachable(c, a));
753 }
754
755 #[test]
756 fn is_reachable_disconnected() {
757 let mut g = ComputeGraph::new();
758 let a = g.add_node(barrier_node());
759 let b = g.add_node(barrier_node());
760 assert!(!g.is_reachable(a, b));
761 assert!(!g.is_reachable(b, a));
762 }
763
764 #[test]
767 fn sources_and_sinks_linear_chain() {
768 let mut g = ComputeGraph::new();
769 let a = g.add_node(barrier_node());
770 let b = g.add_node(barrier_node());
771 let c = g.add_node(barrier_node());
772 g.add_edge(a, b).expect("valid DAG edge from a to b");
773 g.add_edge(b, c).expect("valid DAG edge from b to c");
774 let sources = g.sources();
775 let sinks = g.sinks();
776 assert_eq!(sources, vec![a]);
777 assert_eq!(sinks, vec![c]);
778 }
779
780 #[test]
781 fn sources_and_sinks_diamond() {
782 let mut g = ComputeGraph::new();
783 let a = g.add_node(barrier_node()); let b = g.add_node(barrier_node());
785 let c = g.add_node(barrier_node());
786 let d = g.add_node(barrier_node()); g.add_edge(a, b).expect("valid DAG edge from a to b");
788 g.add_edge(a, c).expect("valid DAG edge from a to c");
789 g.add_edge(b, d).expect("valid DAG edge from b to d");
790 g.add_edge(c, d).expect("valid DAG edge from c to d");
791 assert_eq!(g.sources(), vec![a]);
792 assert_eq!(g.sinks(), vec![d]);
793 }
794
795 #[test]
798 fn topological_order_linear() {
799 let mut g = ComputeGraph::new();
800 let a = g.add_node(barrier_node());
801 let b = g.add_node(barrier_node());
802 let c = g.add_node(barrier_node());
803 g.add_edge(a, b).expect("valid DAG edge from a to b");
804 g.add_edge(b, c).expect("valid DAG edge from b to c");
805 let order = g
806 .topological_order()
807 .expect("topological sort of valid DAG");
808 let pos_a = order
809 .iter()
810 .position(|&x| x == a)
811 .expect("node a present in topological order");
812 let pos_b = order
813 .iter()
814 .position(|&x| x == b)
815 .expect("node b present in topological order");
816 let pos_c = order
817 .iter()
818 .position(|&x| x == c)
819 .expect("node c present in topological order");
820 assert!(pos_a < pos_b && pos_b < pos_c);
821 }
822
823 #[test]
824 fn topological_order_diamond() {
825 let mut g = ComputeGraph::new();
826 let a = g.add_node(barrier_node());
827 let b = g.add_node(barrier_node());
828 let c = g.add_node(barrier_node());
829 let d = g.add_node(barrier_node());
830 g.add_edge(a, b).expect("valid DAG edge from a to b");
831 g.add_edge(a, c).expect("valid DAG edge from a to c");
832 g.add_edge(b, d).expect("valid DAG edge from b to d");
833 g.add_edge(c, d).expect("valid DAG edge from c to d");
834 let order = g
835 .topological_order()
836 .expect("topological sort of valid DAG");
837 assert_eq!(order.len(), 4);
838 let pos = |n: NodeId| {
839 order
840 .iter()
841 .position(|&x| x == n)
842 .expect("node present in topological order")
843 };
844 assert!(pos(a) < pos(b));
845 assert!(pos(a) < pos(c));
846 assert!(pos(b) < pos(d));
847 assert!(pos(c) < pos(d));
848 }
849
850 #[test]
851 fn topological_order_empty_graph() {
852 let g = ComputeGraph::new();
853 assert!(matches!(g.topological_order(), Err(GraphError::EmptyGraph)));
854 }
855
856 #[test]
857 fn topological_order_isolated_nodes() {
858 let mut g = ComputeGraph::new();
859 g.add_node(barrier_node());
860 g.add_node(barrier_node());
861 g.add_node(barrier_node());
862 let order = g
863 .topological_order()
864 .expect("topological sort of valid DAG");
865 assert_eq!(order.len(), 3);
866 }
867
868 #[test]
871 fn infer_data_edges_connects_writer_to_reader() {
872 let mut g = ComputeGraph::new();
873 let buf = g.add_buffer(BufferDescriptor::new(BufferId(0), 1024));
874 let writer = g.add_node(GraphNode::new(NodeId(0), NodeKind::Barrier).with_outputs([buf]));
875 let reader = g.add_node(GraphNode::new(NodeId(0), NodeKind::Barrier).with_inputs([buf]));
876 g.infer_data_edges()
877 .expect("data edge inference on valid graph");
878 assert!(g.is_reachable(writer, reader));
879 }
880
881 #[test]
882 fn infer_data_edges_multiple_readers() {
883 let mut g = ComputeGraph::new();
884 let buf = g.add_buffer(BufferDescriptor::new(BufferId(0), 1024));
885 let writer = g.add_node(GraphNode::new(NodeId(0), NodeKind::Barrier).with_outputs([buf]));
886 let r1 = g.add_node(GraphNode::new(NodeId(0), NodeKind::Barrier).with_inputs([buf]));
887 let r2 = g.add_node(GraphNode::new(NodeId(0), NodeKind::Barrier).with_inputs([buf]));
888 g.infer_data_edges()
889 .expect("data edge inference on valid graph");
890 assert!(g.is_reachable(writer, r1));
891 assert!(g.is_reachable(writer, r2));
892 }
893
894 #[test]
897 fn critical_path_linear_chain() {
898 let mut g = ComputeGraph::new();
899 let a = g.add_node(barrier_node());
900 let b = g.add_node(barrier_node());
901 let c = g.add_node(barrier_node());
902 let d = g.add_node(barrier_node());
903 g.add_edge(a, b).expect("valid DAG edge from a to b");
904 g.add_edge(b, c).expect("valid DAG edge from b to c");
905 g.add_edge(c, d).expect("valid DAG edge from c to d");
906 assert_eq!(
907 g.critical_path_length()
908 .expect("critical path length of valid DAG"),
909 3
910 );
911 }
912
913 #[test]
914 fn critical_path_diamond() {
915 let mut g = ComputeGraph::new();
916 let a = g.add_node(barrier_node());
917 let b = g.add_node(barrier_node());
918 let c = g.add_node(barrier_node());
919 let d = g.add_node(barrier_node());
920 g.add_edge(a, b).expect("valid DAG edge from a to b");
921 g.add_edge(a, c).expect("valid DAG edge from a to c");
922 g.add_edge(b, d).expect("valid DAG edge from b to d");
923 g.add_edge(c, d).expect("valid DAG edge from c to d");
924 assert_eq!(
925 g.critical_path_length()
926 .expect("critical path length of valid DAG"),
927 2
928 );
929 }
930
931 #[test]
932 fn max_degrees_computed_correctly() {
933 let mut g = ComputeGraph::new();
934 let a = g.add_node(barrier_node());
935 let b = g.add_node(barrier_node());
936 let c = g.add_node(barrier_node());
937 let d = g.add_node(barrier_node());
938 g.add_edge(a, b).expect("valid DAG edge from a to b");
939 g.add_edge(a, c).expect("valid DAG edge from a to c");
940 g.add_edge(a, d).expect("valid DAG edge from a to d");
941 assert_eq!(g.max_out_degree(), 3);
942 assert_eq!(g.max_in_degree(), 1);
943 }
944
945 #[test]
948 fn reachable_from_set() {
949 let mut g = ComputeGraph::new();
950 let a = g.add_node(barrier_node());
951 let b = g.add_node(barrier_node());
952 let c = g.add_node(barrier_node());
953 let d = g.add_node(barrier_node());
954 g.add_edge(a, b).expect("valid DAG edge from a to b");
955 g.add_edge(b, c).expect("valid DAG edge from b to c");
956 let reach = g.reachable_from(&[a]);
958 assert!(reach.contains(&a));
959 assert!(reach.contains(&b));
960 assert!(reach.contains(&c));
961 assert!(!reach.contains(&d));
962 }
963
964 #[test]
965 fn reaching_set() {
966 let mut g = ComputeGraph::new();
967 let a = g.add_node(barrier_node());
968 let b = g.add_node(barrier_node());
969 let c = g.add_node(barrier_node());
970 g.add_edge(a, b).expect("valid DAG edge from a to b");
971 g.add_edge(b, c).expect("valid DAG edge from b to c");
972 let reaching = g.reaching(&[c]);
973 assert!(reaching.contains(&a));
974 assert!(reaching.contains(&b));
975 assert!(reaching.contains(&c));
976 }
977
978 #[test]
981 fn kernel_nodes_returns_compute_only() {
982 let mut g = ComputeGraph::new();
983 g.add_node(kernel_node("k0"));
984 g.add_node(barrier_node());
985 g.add_node(memcpy_node(MemcpyDir::HostToDevice, 1024));
986 g.add_node(kernel_node("k1"));
987 let kernels = g.kernel_nodes();
988 assert_eq!(kernels.len(), 2);
989 }
990
991 #[test]
992 fn fusible_nodes_returns_fusible_only() {
993 let mut g = ComputeGraph::new();
994 g.add_node(kernel_node("fusible")); g.add_node(GraphNode::new(
996 NodeId(0),
997 NodeKind::KernelLaunch {
998 function_name: "custom".into(),
999 config: KernelConfig::linear(1, 32, 0),
1000 fusible: false,
1001 },
1002 ));
1003 assert_eq!(g.fusible_nodes().len(), 1);
1004 }
1005
1006 #[test]
1009 fn parallelism_width_linear_is_one() {
1010 let mut g = ComputeGraph::new();
1011 let a = g.add_node(barrier_node());
1012 let b = g.add_node(barrier_node());
1013 let c = g.add_node(barrier_node());
1014 g.add_edge(a, b).expect("valid DAG edge from a to b");
1015 g.add_edge(b, c).expect("valid DAG edge from b to c");
1016 assert_eq!(
1017 g.parallelism_width()
1018 .expect("parallelism width of valid DAG"),
1019 1
1020 );
1021 }
1022
1023 #[test]
1024 fn parallelism_width_fork_join() {
1025 let mut g = ComputeGraph::new();
1026 let src = g.add_node(barrier_node());
1027 let b = g.add_node(barrier_node());
1028 let c = g.add_node(barrier_node());
1029 let d = g.add_node(barrier_node());
1030 let sink = g.add_node(barrier_node());
1031 g.add_edge(src, b).expect("valid DAG edge from src to b");
1032 g.add_edge(src, c).expect("valid DAG edge from src to c");
1033 g.add_edge(src, d).expect("valid DAG edge from src to d");
1034 g.add_edge(b, sink).expect("valid DAG edge from b to sink");
1035 g.add_edge(c, sink).expect("valid DAG edge from c to sink");
1036 g.add_edge(d, sink).expect("valid DAG edge from d to sink");
1037 assert_eq!(
1038 g.parallelism_width()
1039 .expect("parallelism width of valid DAG"),
1040 3
1041 );
1042 }
1043
1044 #[test]
1047 fn to_dot_contains_node_labels() {
1048 let mut g = ComputeGraph::new();
1049 let a = g.add_node(kernel_node("my_kernel").with_name("k0"));
1050 let b = g.add_node(barrier_node());
1051 g.add_edge(a, b).expect("valid DAG edge from a to b");
1052 let dot = g.to_dot();
1053 assert!(dot.contains("digraph"));
1054 assert!(dot.contains("k0"));
1055 assert!(dot.contains("->"));
1056 }
1057
1058 #[test]
1061 fn display_shows_counts() {
1062 let mut g = ComputeGraph::new();
1063 g.add_node(barrier_node());
1064 g.add_node(barrier_node());
1065 g.add_buffer(BufferDescriptor::new(BufferId(0), 1));
1066 let s = g.to_string();
1067 assert!(s.contains("2 nodes"));
1068 assert!(s.contains("1 buffers"));
1069 }
1070}