1use std::collections::{BTreeSet, HashSet, VecDeque};
10use std::fmt;
11
12use nxpu_ir::{Expression, Function, Handle, Statement};
13
14#[derive(Debug, thiserror::Error)]
16pub enum DataflowError {
17 #[error("cycle detected in dataflow graph ({visited} of {total} nodes visited)")]
19 CycleDetected { visited: usize, total: usize },
20}
21
22#[derive(Clone, Debug, PartialEq, Eq)]
24pub enum DfgNodeKind {
25 Store,
27 Call,
29 Atomic,
31 Barrier,
33 Emit,
35 If,
37 Loop,
39 Return,
41 Break,
43 Continue,
45}
46
47impl fmt::Display for DfgNodeKind {
48 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
49 f.write_str(match self {
50 Self::Store => "Store",
51 Self::Call => "Call",
52 Self::Atomic => "Atomic",
53 Self::Barrier => "Barrier",
54 Self::Emit => "Emit",
55 Self::If => "If",
56 Self::Loop => "Loop",
57 Self::Return => "Return",
58 Self::Break => "Break",
59 Self::Continue => "Continue",
60 })
61 }
62}
63
64#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
66pub enum DependencyKind {
67 DataFlow,
69 AntiDependency,
71 OutputDependency,
73 Control,
75}
76
77impl fmt::Display for DependencyKind {
78 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
79 f.write_str(match self {
80 Self::DataFlow => "RAW",
81 Self::AntiDependency => "WAR",
82 Self::OutputDependency => "WAW",
83 Self::Control => "Control",
84 })
85 }
86}
87
88#[derive(Clone, Debug)]
90pub struct DfgNode {
91 pub id: usize,
93 pub kind: DfgNodeKind,
95 pub reads: Vec<Handle<Expression>>,
97 pub writes: Vec<Handle<Expression>>,
99}
100
101#[derive(Clone, Debug)]
103pub struct DfgEdge {
104 pub from: usize,
106 pub to: usize,
108 pub kind: DependencyKind,
110}
111
112#[derive(Clone, Debug)]
117pub struct DataflowGraph {
118 nodes: Vec<DfgNode>,
120 edges: Vec<DfgEdge>,
122}
123
124impl DataflowGraph {
125 pub fn build(func: &Function) -> Self {
134 let mut nodes = Vec::new();
135 let mut edges = Vec::new();
136
137 for stmt in &func.body {
139 let id = nodes.len();
140 let (kind, reads, writes) = classify_statement(stmt, &func.expressions);
141 nodes.push(DfgNode {
142 id,
143 kind,
144 reads,
145 writes,
146 });
147 }
148
149 for (i, node_i) in nodes.iter().enumerate() {
151 for (j, node_j) in nodes.iter().enumerate().skip(i + 1) {
152 for w in &node_i.writes {
154 if node_j.reads.contains(w) {
155 edges.push(DfgEdge {
156 from: i,
157 to: j,
158 kind: DependencyKind::DataFlow,
159 });
160 }
161 }
162
163 for r in &node_i.reads {
165 if node_j.writes.contains(r) {
166 edges.push(DfgEdge {
167 from: i,
168 to: j,
169 kind: DependencyKind::AntiDependency,
170 });
171 }
172 }
173
174 for w in &node_i.writes {
176 if node_j.writes.contains(w) {
177 edges.push(DfgEdge {
178 from: i,
179 to: j,
180 kind: DependencyKind::OutputDependency,
181 });
182 }
183 }
184 }
185 }
186
187 for (i, node_i) in nodes.iter().enumerate() {
189 if node_i.kind == DfgNodeKind::Barrier {
190 for (j, node_j) in nodes.iter().enumerate().skip(i + 1) {
191 let has_memory_effect = !node_j.reads.is_empty()
192 || !node_j.writes.is_empty()
193 || node_j.kind == DfgNodeKind::Barrier;
194 if has_memory_effect {
195 edges.push(DfgEdge {
196 from: i,
197 to: j,
198 kind: DependencyKind::Control,
199 });
200 }
201 }
202 }
203 }
204
205 Self { nodes, edges }
206 }
207
208 pub fn nodes(&self) -> &[DfgNode] {
210 &self.nodes
211 }
212
213 pub fn edges(&self) -> &[DfgEdge] {
215 &self.edges
216 }
217
218 pub fn node_count(&self) -> usize {
220 self.nodes.len()
221 }
222
223 pub fn edge_count(&self) -> usize {
225 self.edges.len()
226 }
227
228 pub fn successors(&self, node_id: usize) -> Vec<&DfgEdge> {
230 self.edges.iter().filter(|e| e.from == node_id).collect()
231 }
232
233 pub fn predecessors(&self, node_id: usize) -> Vec<&DfgEdge> {
235 self.edges.iter().filter(|e| e.to == node_id).collect()
236 }
237
238 pub fn topological_sort(&self) -> Result<Vec<usize>, DataflowError> {
246 let n = self.nodes.len();
247 if n == 0 {
248 return Ok(Vec::new());
249 }
250
251 let mut in_degree = vec![0usize; n];
253 let mut successors: Vec<Vec<usize>> = vec![Vec::new(); n];
254
255 for edge in &self.edges {
256 in_degree[edge.to] += 1;
257 successors[edge.from].push(edge.to);
258 }
259
260 let mut ready: BTreeSet<usize> = BTreeSet::new();
262 for (i, °) in in_degree.iter().enumerate() {
263 if deg == 0 {
264 ready.insert(i);
265 }
266 }
267
268 let mut result = Vec::with_capacity(n);
269
270 while let Some(&node_id) = ready.iter().next() {
271 ready.remove(&node_id);
272 result.push(node_id);
273
274 for &succ in &successors[node_id] {
275 in_degree[succ] -= 1;
276 if in_degree[succ] == 0 {
277 ready.insert(succ);
278 }
279 }
280 }
281
282 if result.len() != n {
283 return Err(DataflowError::CycleDetected {
284 visited: result.len(),
285 total: n,
286 });
287 }
288
289 Ok(result)
290 }
291
292 pub fn critical_path(&self) -> CriticalPathResult {
300 self.critical_path_with_costs(&vec![1usize; self.nodes.len()])
301 }
302
303 pub fn critical_path_with_costs(&self, costs: &[usize]) -> CriticalPathResult {
309 let n = self.nodes.len();
310 if n == 0 {
311 return CriticalPathResult {
312 node_distances: Vec::new(),
313 critical_path_length: 0,
314 critical_path: Vec::new(),
315 };
316 }
317
318 let topo = match self.topological_sort() {
320 Ok(order) => order,
321 Err(_) => {
322 return CriticalPathResult {
323 node_distances: vec![0; n],
324 critical_path_length: 0,
325 critical_path: Vec::new(),
326 };
327 }
328 };
329
330 let mut preds: Vec<Vec<usize>> = vec![Vec::new(); n];
332 for edge in &self.edges {
333 preds[edge.to].push(edge.from);
334 }
335
336 let mut dist = vec![0usize; n];
338 let mut prev: Vec<Option<usize>> = vec![None; n];
340
341 for &node_id in &topo {
342 let node_cost = costs.get(node_id).copied().unwrap_or(1);
343 let mut best_dist = 0;
344 let mut best_pred = None;
345
346 for &pred_id in &preds[node_id] {
347 if dist[pred_id] > best_dist {
348 best_dist = dist[pred_id];
349 best_pred = Some(pred_id);
350 }
351 }
352
353 dist[node_id] = best_dist + node_cost;
354 prev[node_id] = best_pred;
355 }
356
357 let mut max_dist = 0;
359 let mut max_node = 0;
360 for (i, &d) in dist.iter().enumerate() {
361 if d > max_dist {
362 max_dist = d;
363 max_node = i;
364 }
365 }
366
367 let mut path = VecDeque::new();
369 let mut current = Some(max_node);
370 while let Some(node_id) = current {
371 path.push_front(node_id);
372 current = prev[node_id];
373 }
374
375 CriticalPathResult {
376 node_distances: dist,
377 critical_path_length: max_dist,
378 critical_path: path.into(),
379 }
380 }
381
382 pub fn parallel_groups(&self) -> Vec<Vec<usize>> {
387 let n = self.nodes.len();
388 if n == 0 {
389 return Vec::new();
390 }
391
392 let topo = match self.topological_sort() {
393 Ok(order) => order,
394 Err(_) => return Vec::new(),
395 };
396
397 let mut has_pred: Vec<HashSet<usize>> = vec![HashSet::new(); n];
399 for edge in &self.edges {
400 has_pred[edge.to].insert(edge.from);
401 }
402
403 let mut asap = vec![0usize; n];
405 for &node_id in &topo {
406 let mut max_pred_time = 0;
407 for &pred_id in &has_pred[node_id] {
408 let pred_finish = asap[pred_id] + 1;
409 if pred_finish > max_pred_time {
410 max_pred_time = pred_finish;
411 }
412 }
413 asap[node_id] = max_pred_time;
414 }
415
416 let max_time = asap.iter().copied().max().unwrap_or(0);
418 let mut groups: Vec<Vec<usize>> = vec![Vec::new(); max_time + 1];
419 for &node_id in &topo {
420 groups[asap[node_id]].push(node_id);
421 }
422
423 groups.retain(|g| !g.is_empty());
425 groups
426 }
427}
428
429#[derive(Clone, Debug)]
431pub struct CriticalPathResult {
432 pub node_distances: Vec<usize>,
434 pub critical_path_length: usize,
436 pub critical_path: Vec<usize>,
438}
439
440fn classify_statement(
442 stmt: &Statement,
443 expressions: &nxpu_ir::Arena<Expression>,
444) -> (
445 DfgNodeKind,
446 Vec<Handle<Expression>>,
447 Vec<Handle<Expression>>,
448) {
449 match stmt {
450 Statement::Store { pointer, value } => {
451 let mut reads = Vec::new();
452 reads.push(*value);
453 collect_expr_reads(*value, expressions, &mut reads);
454 collect_expr_reads(*pointer, expressions, &mut reads);
455 let writes = vec![*pointer];
456 (DfgNodeKind::Store, reads, writes)
457 }
458 Statement::Call {
459 arguments, result, ..
460 } => {
461 let mut reads: Vec<Handle<Expression>> = arguments.clone();
462 for &arg in arguments {
463 collect_expr_reads(arg, expressions, &mut reads);
464 }
465 let writes = result.iter().copied().collect();
466 (DfgNodeKind::Call, reads, writes)
467 }
468 Statement::Atomic {
469 pointer,
470 value,
471 result,
472 fun,
473 } => {
474 let mut reads = vec![*pointer, *value];
475 collect_expr_reads(*pointer, expressions, &mut reads);
476 collect_expr_reads(*value, expressions, &mut reads);
477 if let nxpu_ir::AtomicFunction::Exchange {
478 compare: Some(cmp), ..
479 } = fun
480 {
481 reads.push(*cmp);
482 collect_expr_reads(*cmp, expressions, &mut reads);
483 }
484 let mut writes = vec![*pointer];
485 if let Some(r) = result {
486 writes.push(*r);
487 }
488 (DfgNodeKind::Atomic, reads, writes)
489 }
490 Statement::Barrier(_) => (DfgNodeKind::Barrier, Vec::new(), Vec::new()),
491 Statement::Emit(range) => {
492 let mut reads = Vec::new();
493 let mut writes = Vec::new();
494 let idx_range = range.index_range();
495 for (handle, expr) in expressions.iter() {
497 let idx = handle.index() as u32;
498 if idx >= idx_range.start && idx < idx_range.end {
499 writes.push(handle);
500 for operand in expression_operands(expr) {
501 if !reads.contains(&operand) {
502 reads.push(operand);
503 }
504 }
505 }
506 }
507 (DfgNodeKind::Emit, reads, writes)
508 }
509 Statement::If {
510 condition,
511 accept,
512 reject,
513 } => {
514 let mut reads = vec![*condition];
515 collect_expr_reads(*condition, expressions, &mut reads);
516 for child in accept.iter().chain(reject.iter()) {
518 let (_, child_reads, _) = classify_statement(child, expressions);
519 for r in child_reads {
520 if !reads.contains(&r) {
521 reads.push(r);
522 }
523 }
524 }
525 let mut writes = Vec::new();
526 for child in accept.iter().chain(reject.iter()) {
527 let (_, _, child_writes) = classify_statement(child, expressions);
528 for w in child_writes {
529 if !writes.contains(&w) {
530 writes.push(w);
531 }
532 }
533 }
534 (DfgNodeKind::If, reads, writes)
535 }
536 Statement::Loop {
537 body,
538 continuing,
539 break_if,
540 } => {
541 let mut reads = Vec::new();
542 if let Some(brk) = break_if {
543 reads.push(*brk);
544 collect_expr_reads(*brk, expressions, &mut reads);
545 }
546 for child in body.iter().chain(continuing.iter()) {
547 let (_, child_reads, _) = classify_statement(child, expressions);
548 for r in child_reads {
549 if !reads.contains(&r) {
550 reads.push(r);
551 }
552 }
553 }
554 let mut writes = Vec::new();
555 for child in body.iter().chain(continuing.iter()) {
556 let (_, _, child_writes) = classify_statement(child, expressions);
557 for w in child_writes {
558 if !writes.contains(&w) {
559 writes.push(w);
560 }
561 }
562 }
563 (DfgNodeKind::Loop, reads, writes)
564 }
565 Statement::Return { value } => {
566 let mut reads = Vec::new();
567 if let Some(v) = value {
568 reads.push(*v);
569 collect_expr_reads(*v, expressions, &mut reads);
570 }
571 (DfgNodeKind::Return, reads, Vec::new())
572 }
573 Statement::Break => (DfgNodeKind::Break, Vec::new(), Vec::new()),
574 Statement::Continue => (DfgNodeKind::Continue, Vec::new(), Vec::new()),
575 }
576}
577
578fn collect_expr_reads(
580 handle: Handle<Expression>,
581 expressions: &nxpu_ir::Arena<Expression>,
582 reads: &mut Vec<Handle<Expression>>,
583) {
584 if let Some(expr) = expressions.try_get(handle) {
585 for operand in expression_operands(expr) {
586 if !reads.contains(&operand) {
587 reads.push(operand);
588 collect_expr_reads(operand, expressions, reads);
589 }
590 }
591 }
592}
593
594fn expression_operands(expr: &Expression) -> Vec<Handle<Expression>> {
596 match expr {
597 Expression::Literal(_)
598 | Expression::FunctionArgument(_)
599 | Expression::GlobalVariable(_)
600 | Expression::LocalVariable(_)
601 | Expression::CallResult(_)
602 | Expression::AtomicResult { .. }
603 | Expression::ZeroValue(_) => vec![],
604
605 Expression::Load { pointer } => vec![*pointer],
606 Expression::Unary { expr, .. } => vec![*expr],
607 Expression::ArrayLength(e) => vec![*e],
608 Expression::Splat { value, .. } => vec![*value],
609 Expression::As { expr, .. } => vec![*expr],
610
611 Expression::Binary { left, right, .. } => vec![*left, *right],
612 Expression::Access { base, index } => vec![*base, *index],
613 Expression::AccessIndex { base, .. } => vec![*base],
614 Expression::Select {
615 condition,
616 accept,
617 reject,
618 } => vec![*condition, *accept, *reject],
619 Expression::Swizzle { vector, .. } => vec![*vector],
620
621 Expression::Compose { components, .. } => components.clone(),
622 Expression::Math {
623 arg,
624 arg1,
625 arg2,
626 arg3,
627 ..
628 } => {
629 let mut ops = vec![*arg];
630 if let Some(a) = arg1 {
631 ops.push(*a);
632 }
633 if let Some(a) = arg2 {
634 ops.push(*a);
635 }
636 if let Some(a) = arg3 {
637 ops.push(*a);
638 }
639 ops
640 }
641 }
642}
643
644impl fmt::Display for DataflowGraph {
645 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
646 writeln!(
647 f,
648 "DataflowGraph ({} nodes, {} edges):",
649 self.nodes.len(),
650 self.edges.len()
651 )?;
652 for node in &self.nodes {
653 writeln!(
654 f,
655 " node[{}]: {} (reads: {}, writes: {})",
656 node.id,
657 node.kind,
658 node.reads.len(),
659 node.writes.len()
660 )?;
661 }
662 for edge in &self.edges {
663 writeln!(f, " edge: {} -> {} ({})", edge.from, edge.to, edge.kind)?;
664 }
665 Ok(())
666 }
667}
668
669#[cfg(test)]
670mod tests {
671 use super::*;
672 use nxpu_ir::{Expression, Function, Literal, Statement};
673
674 fn make_store_chain() -> Function {
677 let mut func = Function::new("test");
678
679 let gv0 = {
681 let mut arena = nxpu_ir::Arena::new();
682 arena.append(nxpu_ir::GlobalVariable {
683 name: Some("a".into()),
684 space: nxpu_ir::AddressSpace::Storage {
685 access: nxpu_ir::StorageAccess::LOAD | nxpu_ir::StorageAccess::STORE,
686 },
687 binding: None,
688 ty: {
689 let mut types = nxpu_ir::UniqueArena::new();
690 types.insert(nxpu_ir::Type {
691 name: None,
692 inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
693 })
694 },
695 init: None,
696 layout: None,
697 })
698 };
699 let gv1 = {
700 let mut arena = nxpu_ir::Arena::new();
701 let _ = arena.append(nxpu_ir::GlobalVariable {
702 name: None,
703 space: nxpu_ir::AddressSpace::Storage {
704 access: nxpu_ir::StorageAccess::LOAD | nxpu_ir::StorageAccess::STORE,
705 },
706 binding: None,
707 ty: {
708 let mut types = nxpu_ir::UniqueArena::new();
709 types.insert(nxpu_ir::Type {
710 name: None,
711 inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
712 })
713 },
714 init: None,
715 layout: None,
716 });
717 arena.append(nxpu_ir::GlobalVariable {
718 name: Some("b".into()),
719 space: nxpu_ir::AddressSpace::Storage {
720 access: nxpu_ir::StorageAccess::LOAD | nxpu_ir::StorageAccess::STORE,
721 },
722 binding: None,
723 ty: {
724 let mut types = nxpu_ir::UniqueArena::new();
725 types.insert(nxpu_ir::Type {
726 name: None,
727 inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
728 })
729 },
730 init: None,
731 layout: None,
732 })
733 };
734
735 let ptr_a = func.expressions.append(Expression::GlobalVariable(gv0));
736 let ptr_b = func.expressions.append(Expression::GlobalVariable(gv1));
737 let val = func
738 .expressions
739 .append(Expression::Literal(Literal::F32(42.0)));
740
741 func.body.push(Statement::Store {
743 pointer: ptr_a,
744 value: val,
745 });
746
747 func.body.push(Statement::Store {
749 pointer: ptr_b,
750 value: val,
751 });
752
753 func
754 }
755
756 #[test]
757 fn build_dfg_empty_function() {
758 let func = Function::new("empty");
759 let dfg = DataflowGraph::build(&func);
760 assert_eq!(dfg.node_count(), 0);
761 assert_eq!(dfg.edge_count(), 0);
762 }
763
764 #[test]
765 fn build_dfg_simple_stores() {
766 let func = make_store_chain();
767 let dfg = DataflowGraph::build(&func);
768 assert_eq!(dfg.node_count(), 2);
769 assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Store);
774 assert_eq!(dfg.nodes()[1].kind, DfgNodeKind::Store);
775 }
776
777 #[test]
778 fn build_dfg_raw_dependency() {
779 let mut func = Function::new("test");
782
783 let gv = {
784 let mut arena = nxpu_ir::Arena::new();
785 arena.append(nxpu_ir::GlobalVariable {
786 name: Some("x".into()),
787 space: nxpu_ir::AddressSpace::Storage {
788 access: nxpu_ir::StorageAccess::LOAD | nxpu_ir::StorageAccess::STORE,
789 },
790 binding: None,
791 ty: {
792 let mut types = nxpu_ir::UniqueArena::new();
793 types.insert(nxpu_ir::Type {
794 name: None,
795 inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
796 })
797 },
798 init: None,
799 layout: None,
800 })
801 };
802
803 let ptr = func.expressions.append(Expression::GlobalVariable(gv));
804 let val = func
805 .expressions
806 .append(Expression::Literal(Literal::F32(1.0)));
807
808 func.body.push(Statement::Store {
810 pointer: ptr,
811 value: val,
812 });
813 func.body.push(Statement::Store {
815 pointer: ptr,
816 value: val,
817 });
818
819 let dfg = DataflowGraph::build(&func);
820 assert_eq!(dfg.node_count(), 2);
821
822 let waw_edges: Vec<_> = dfg
824 .edges()
825 .iter()
826 .filter(|e| e.kind == DependencyKind::OutputDependency)
827 .collect();
828 assert!(
829 !waw_edges.is_empty(),
830 "expected WAW dependency for two stores to same pointer"
831 );
832 }
833
834 #[test]
835 fn build_dfg_war_dependency() {
836 let mut func = Function::new("test");
838
839 let gv = {
840 let mut arena = nxpu_ir::Arena::new();
841 arena.append(nxpu_ir::GlobalVariable {
842 name: Some("x".into()),
843 space: nxpu_ir::AddressSpace::Storage {
844 access: nxpu_ir::StorageAccess::LOAD | nxpu_ir::StorageAccess::STORE,
845 },
846 binding: None,
847 ty: {
848 let mut types = nxpu_ir::UniqueArena::new();
849 types.insert(nxpu_ir::Type {
850 name: None,
851 inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
852 })
853 },
854 init: None,
855 layout: None,
856 })
857 };
858
859 let ptr = func.expressions.append(Expression::GlobalVariable(gv));
860 let val = func
861 .expressions
862 .append(Expression::Literal(Literal::F32(1.0)));
863
864 func.body.push(Statement::Return { value: Some(ptr) });
866 func.body.push(Statement::Store {
868 pointer: ptr,
869 value: val,
870 });
871
872 let dfg = DataflowGraph::build(&func);
873 assert_eq!(dfg.node_count(), 2);
874
875 let war_edges: Vec<_> = dfg
877 .edges()
878 .iter()
879 .filter(|e| e.kind == DependencyKind::AntiDependency)
880 .collect();
881 assert!(!war_edges.is_empty(), "expected WAR dependency");
882 }
883
884 #[test]
885 fn topological_sort_linear_chain() {
886 let mut func = Function::new("test");
888
889 let gv = {
890 let mut arena = nxpu_ir::Arena::new();
891 arena.append(nxpu_ir::GlobalVariable {
892 name: Some("x".into()),
893 space: nxpu_ir::AddressSpace::Storage {
894 access: nxpu_ir::StorageAccess::LOAD | nxpu_ir::StorageAccess::STORE,
895 },
896 binding: None,
897 ty: {
898 let mut types = nxpu_ir::UniqueArena::new();
899 types.insert(nxpu_ir::Type {
900 name: None,
901 inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
902 })
903 },
904 init: None,
905 layout: None,
906 })
907 };
908
909 let ptr = func.expressions.append(Expression::GlobalVariable(gv));
910 let val = func
911 .expressions
912 .append(Expression::Literal(Literal::F32(1.0)));
913
914 func.body.push(Statement::Store {
916 pointer: ptr,
917 value: val,
918 });
919 func.body.push(Statement::Store {
920 pointer: ptr,
921 value: val,
922 });
923 func.body.push(Statement::Store {
924 pointer: ptr,
925 value: val,
926 });
927
928 let dfg = DataflowGraph::build(&func);
929 let order = dfg.topological_sort().unwrap();
930 assert_eq!(order.len(), 3);
931 assert_eq!(order[0], 0);
933 assert_eq!(order[1], 1);
934 assert_eq!(order[2], 2);
935 }
936
937 #[test]
938 fn topological_sort_empty() {
939 let func = Function::new("empty");
940 let dfg = DataflowGraph::build(&func);
941 let order = dfg.topological_sort().unwrap();
942 assert_eq!(order.len(), 0);
943 }
944
945 #[test]
946 fn topological_sort_independent_nodes() {
947 let mut func = Function::new("test");
948
949 let gv0 = {
951 let mut arena = nxpu_ir::Arena::new();
952 arena.append(nxpu_ir::GlobalVariable {
953 name: Some("a".into()),
954 space: nxpu_ir::AddressSpace::Storage {
955 access: nxpu_ir::StorageAccess::STORE,
956 },
957 binding: None,
958 ty: {
959 let mut types = nxpu_ir::UniqueArena::new();
960 types.insert(nxpu_ir::Type {
961 name: None,
962 inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
963 })
964 },
965 init: None,
966 layout: None,
967 })
968 };
969 let gv1 = {
970 let mut arena = nxpu_ir::Arena::new();
971 let _ = arena.append(nxpu_ir::GlobalVariable {
972 name: None,
973 space: nxpu_ir::AddressSpace::Storage {
974 access: nxpu_ir::StorageAccess::STORE,
975 },
976 binding: None,
977 ty: {
978 let mut types = nxpu_ir::UniqueArena::new();
979 types.insert(nxpu_ir::Type {
980 name: None,
981 inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
982 })
983 },
984 init: None,
985 layout: None,
986 });
987 arena.append(nxpu_ir::GlobalVariable {
988 name: Some("b".into()),
989 space: nxpu_ir::AddressSpace::Storage {
990 access: nxpu_ir::StorageAccess::STORE,
991 },
992 binding: None,
993 ty: {
994 let mut types = nxpu_ir::UniqueArena::new();
995 types.insert(nxpu_ir::Type {
996 name: None,
997 inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
998 })
999 },
1000 init: None,
1001 layout: None,
1002 })
1003 };
1004
1005 let ptr_a = func.expressions.append(Expression::GlobalVariable(gv0));
1006 let val_a = func
1007 .expressions
1008 .append(Expression::Literal(Literal::F32(1.0)));
1009 let ptr_b = func.expressions.append(Expression::GlobalVariable(gv1));
1010 let val_b = func
1011 .expressions
1012 .append(Expression::Literal(Literal::F32(2.0)));
1013
1014 func.body.push(Statement::Store {
1015 pointer: ptr_a,
1016 value: val_a,
1017 });
1018 func.body.push(Statement::Store {
1019 pointer: ptr_b,
1020 value: val_b,
1021 });
1022
1023 let dfg = DataflowGraph::build(&func);
1024 let order = dfg.topological_sort().unwrap();
1025 assert_eq!(order.len(), 2);
1026 }
1028
1029 #[test]
1030 fn critical_path_linear_chain() {
1031 let mut func = Function::new("test");
1032
1033 let gv = {
1034 let mut arena = nxpu_ir::Arena::new();
1035 arena.append(nxpu_ir::GlobalVariable {
1036 name: Some("x".into()),
1037 space: nxpu_ir::AddressSpace::Storage {
1038 access: nxpu_ir::StorageAccess::LOAD | nxpu_ir::StorageAccess::STORE,
1039 },
1040 binding: None,
1041 ty: {
1042 let mut types = nxpu_ir::UniqueArena::new();
1043 types.insert(nxpu_ir::Type {
1044 name: None,
1045 inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
1046 })
1047 },
1048 init: None,
1049 layout: None,
1050 })
1051 };
1052
1053 let ptr = func.expressions.append(Expression::GlobalVariable(gv));
1054 let val = func
1055 .expressions
1056 .append(Expression::Literal(Literal::F32(1.0)));
1057
1058 func.body.push(Statement::Store {
1060 pointer: ptr,
1061 value: val,
1062 });
1063 func.body.push(Statement::Store {
1064 pointer: ptr,
1065 value: val,
1066 });
1067 func.body.push(Statement::Store {
1068 pointer: ptr,
1069 value: val,
1070 });
1071
1072 let dfg = DataflowGraph::build(&func);
1073 let cp = dfg.critical_path();
1074 assert_eq!(cp.critical_path_length, 3);
1076 assert_eq!(cp.critical_path.len(), 3);
1077 }
1078
1079 #[test]
1080 fn critical_path_empty() {
1081 let func = Function::new("empty");
1082 let dfg = DataflowGraph::build(&func);
1083 let cp = dfg.critical_path();
1084 assert_eq!(cp.critical_path_length, 0);
1085 assert_eq!(cp.critical_path.len(), 0);
1086 }
1087
1088 #[test]
1089 fn barrier_creates_control_edges() {
1090 let mut func = Function::new("test");
1091
1092 let gv = {
1093 let mut arena = nxpu_ir::Arena::new();
1094 arena.append(nxpu_ir::GlobalVariable {
1095 name: Some("x".into()),
1096 space: nxpu_ir::AddressSpace::Storage {
1097 access: nxpu_ir::StorageAccess::LOAD | nxpu_ir::StorageAccess::STORE,
1098 },
1099 binding: None,
1100 ty: {
1101 let mut types = nxpu_ir::UniqueArena::new();
1102 types.insert(nxpu_ir::Type {
1103 name: None,
1104 inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
1105 })
1106 },
1107 init: None,
1108 layout: None,
1109 })
1110 };
1111
1112 let ptr = func.expressions.append(Expression::GlobalVariable(gv));
1113 let val = func
1114 .expressions
1115 .append(Expression::Literal(Literal::F32(1.0)));
1116
1117 func.body.push(Statement::Store {
1119 pointer: ptr,
1120 value: val,
1121 });
1122 func.body
1124 .push(Statement::Barrier(nxpu_ir::Barrier::STORAGE));
1125 func.body.push(Statement::Store {
1127 pointer: ptr,
1128 value: val,
1129 });
1130
1131 let dfg = DataflowGraph::build(&func);
1132 assert_eq!(dfg.node_count(), 3);
1133
1134 let control_edges: Vec<_> = dfg
1136 .edges()
1137 .iter()
1138 .filter(|e| e.kind == DependencyKind::Control)
1139 .collect();
1140 assert!(
1141 !control_edges.is_empty(),
1142 "expected control dependency from barrier"
1143 );
1144 assert!(control_edges.iter().any(|e| e.from == 1 && e.to == 2));
1146 }
1147
1148 #[test]
1149 fn parallel_groups_independent_ops() {
1150 let mut func = Function::new("test");
1151
1152 let gv0 = {
1153 let mut arena = nxpu_ir::Arena::new();
1154 arena.append(nxpu_ir::GlobalVariable {
1155 name: Some("a".into()),
1156 space: nxpu_ir::AddressSpace::Storage {
1157 access: nxpu_ir::StorageAccess::STORE,
1158 },
1159 binding: None,
1160 ty: {
1161 let mut types = nxpu_ir::UniqueArena::new();
1162 types.insert(nxpu_ir::Type {
1163 name: None,
1164 inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
1165 })
1166 },
1167 init: None,
1168 layout: None,
1169 })
1170 };
1171 let gv1 = {
1172 let mut arena = nxpu_ir::Arena::new();
1173 let _ = arena.append(nxpu_ir::GlobalVariable {
1174 name: None,
1175 space: nxpu_ir::AddressSpace::Storage {
1176 access: nxpu_ir::StorageAccess::STORE,
1177 },
1178 binding: None,
1179 ty: {
1180 let mut types = nxpu_ir::UniqueArena::new();
1181 types.insert(nxpu_ir::Type {
1182 name: None,
1183 inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
1184 })
1185 },
1186 init: None,
1187 layout: None,
1188 });
1189 arena.append(nxpu_ir::GlobalVariable {
1190 name: Some("b".into()),
1191 space: nxpu_ir::AddressSpace::Storage {
1192 access: nxpu_ir::StorageAccess::STORE,
1193 },
1194 binding: None,
1195 ty: {
1196 let mut types = nxpu_ir::UniqueArena::new();
1197 types.insert(nxpu_ir::Type {
1198 name: None,
1199 inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
1200 })
1201 },
1202 init: None,
1203 layout: None,
1204 })
1205 };
1206
1207 let ptr_a = func.expressions.append(Expression::GlobalVariable(gv0));
1208 let val_a = func
1209 .expressions
1210 .append(Expression::Literal(Literal::F32(1.0)));
1211 let ptr_b = func.expressions.append(Expression::GlobalVariable(gv1));
1212 let val_b = func
1213 .expressions
1214 .append(Expression::Literal(Literal::F32(2.0)));
1215
1216 func.body.push(Statement::Store {
1217 pointer: ptr_a,
1218 value: val_a,
1219 });
1220 func.body.push(Statement::Store {
1221 pointer: ptr_b,
1222 value: val_b,
1223 });
1224
1225 let dfg = DataflowGraph::build(&func);
1226 let groups = dfg.parallel_groups();
1227
1228 let max_group_size = groups.iter().map(|g| g.len()).max().unwrap_or(0);
1231 assert!(
1232 max_group_size >= 2,
1233 "expected independent ops to be grouped together, got groups: {groups:?}"
1234 );
1235 }
1236
1237 #[test]
1238 fn display_dfg() {
1239 let func = Function::new("empty");
1240 let dfg = DataflowGraph::build(&func);
1241 let s = format!("{dfg}");
1242 assert!(s.contains("DataflowGraph"));
1243 }
1244
1245 #[test]
1246 fn critical_path_with_custom_costs() {
1247 let mut func = Function::new("test");
1248
1249 let gv = {
1250 let mut arena = nxpu_ir::Arena::new();
1251 arena.append(nxpu_ir::GlobalVariable {
1252 name: Some("x".into()),
1253 space: nxpu_ir::AddressSpace::Storage {
1254 access: nxpu_ir::StorageAccess::LOAD | nxpu_ir::StorageAccess::STORE,
1255 },
1256 binding: None,
1257 ty: {
1258 let mut types = nxpu_ir::UniqueArena::new();
1259 types.insert(nxpu_ir::Type {
1260 name: None,
1261 inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
1262 })
1263 },
1264 init: None,
1265 layout: None,
1266 })
1267 };
1268
1269 let ptr = func.expressions.append(Expression::GlobalVariable(gv));
1270 let val = func
1271 .expressions
1272 .append(Expression::Literal(Literal::F32(1.0)));
1273
1274 func.body.push(Statement::Store {
1275 pointer: ptr,
1276 value: val,
1277 });
1278 func.body.push(Statement::Store {
1279 pointer: ptr,
1280 value: val,
1281 });
1282
1283 let dfg = DataflowGraph::build(&func);
1284 let cp = dfg.critical_path_with_costs(&[5, 3]);
1285 assert_eq!(cp.critical_path_length, 8);
1287 }
1288
1289 fn make_gv(index: usize) -> Handle<nxpu_ir::GlobalVariable> {
1291 let mut arena = nxpu_ir::Arena::new();
1292 let mut handle = None;
1294 for i in 0..=index {
1295 let h = arena.append(nxpu_ir::GlobalVariable {
1296 name: Some(format!("gv_{i}")),
1297 space: nxpu_ir::AddressSpace::Storage {
1298 access: nxpu_ir::StorageAccess::LOAD | nxpu_ir::StorageAccess::STORE,
1299 },
1300 binding: None,
1301 ty: {
1302 let mut types = nxpu_ir::UniqueArena::new();
1303 types.insert(nxpu_ir::Type {
1304 name: None,
1305 inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
1306 })
1307 },
1308 init: None,
1309 layout: None,
1310 });
1311 handle = Some(h);
1312 }
1313 handle.unwrap()
1314 }
1315
1316 #[test]
1319 fn display_dfg_with_nodes_and_edges() {
1320 let mut func = Function::new("test");
1321 let gv = make_gv(0);
1322 let ptr = func.expressions.append(Expression::GlobalVariable(gv));
1323 let val = func
1324 .expressions
1325 .append(Expression::Literal(Literal::F32(1.0)));
1326
1327 func.body.push(Statement::Store {
1329 pointer: ptr,
1330 value: val,
1331 });
1332 func.body.push(Statement::Store {
1333 pointer: ptr,
1334 value: val,
1335 });
1336
1337 let dfg = DataflowGraph::build(&func);
1338 let s = format!("{dfg}");
1339 assert!(s.contains("DataflowGraph (2 nodes, "));
1340 assert!(s.contains("node[0]: Store"));
1341 assert!(s.contains("node[1]: Store"));
1342 assert!(s.contains("edge:"));
1343 assert!(s.contains("WAW"));
1344 }
1345
1346 #[test]
1347 fn dfg_node_kind_display_all_variants() {
1348 assert_eq!(format!("{}", DfgNodeKind::Store), "Store");
1349 assert_eq!(format!("{}", DfgNodeKind::Call), "Call");
1350 assert_eq!(format!("{}", DfgNodeKind::Atomic), "Atomic");
1351 assert_eq!(format!("{}", DfgNodeKind::Barrier), "Barrier");
1352 assert_eq!(format!("{}", DfgNodeKind::Emit), "Emit");
1353 assert_eq!(format!("{}", DfgNodeKind::If), "If");
1354 assert_eq!(format!("{}", DfgNodeKind::Loop), "Loop");
1355 assert_eq!(format!("{}", DfgNodeKind::Return), "Return");
1356 assert_eq!(format!("{}", DfgNodeKind::Break), "Break");
1357 assert_eq!(format!("{}", DfgNodeKind::Continue), "Continue");
1358 }
1359
1360 #[test]
1361 fn dependency_kind_display_all_variants() {
1362 assert_eq!(format!("{}", DependencyKind::DataFlow), "RAW");
1363 assert_eq!(format!("{}", DependencyKind::AntiDependency), "WAR");
1364 assert_eq!(format!("{}", DependencyKind::OutputDependency), "WAW");
1365 assert_eq!(format!("{}", DependencyKind::Control), "Control");
1366 }
1367
1368 #[test]
1369 fn dataflow_error_display() {
1370 let err = DataflowError::CycleDetected {
1371 visited: 3,
1372 total: 5,
1373 };
1374 let msg = format!("{err}");
1375 assert!(msg.contains("cycle detected"));
1376 assert!(msg.contains("3 of 5"));
1377 }
1378
1379 #[test]
1382 fn successors_and_predecessors() {
1383 let mut func = Function::new("test");
1384 let gv = make_gv(0);
1385 let ptr = func.expressions.append(Expression::GlobalVariable(gv));
1386 let val = func
1387 .expressions
1388 .append(Expression::Literal(Literal::F32(1.0)));
1389
1390 func.body.push(Statement::Store {
1392 pointer: ptr,
1393 value: val,
1394 });
1395 func.body.push(Statement::Store {
1396 pointer: ptr,
1397 value: val,
1398 });
1399 func.body.push(Statement::Store {
1400 pointer: ptr,
1401 value: val,
1402 });
1403
1404 let dfg = DataflowGraph::build(&func);
1405
1406 let succs_0 = dfg.successors(0);
1408 assert_ne!(succs_0.len(), 0);
1409
1410 let preds_2 = dfg.predecessors(2);
1412 assert_ne!(preds_2.len(), 0);
1413
1414 let preds_0 = dfg.predecessors(0);
1416 assert_eq!(preds_0.len(), 0);
1417 }
1418
1419 #[test]
1422 fn break_and_continue_nodes() {
1423 let mut func = Function::new("test");
1424 func.body.push(Statement::Break);
1425 func.body.push(Statement::Continue);
1426
1427 let dfg = DataflowGraph::build(&func);
1428 assert_eq!(dfg.node_count(), 2);
1429 assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Break);
1430 assert_eq!(dfg.nodes()[1].kind, DfgNodeKind::Continue);
1431 assert_eq!(dfg.edge_count(), 0);
1433 }
1434
1435 #[test]
1438 fn return_no_value() {
1439 let mut func = Function::new("test");
1440 func.body.push(Statement::Return { value: None });
1441
1442 let dfg = DataflowGraph::build(&func);
1443 assert_eq!(dfg.node_count(), 1);
1444 assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Return);
1445 assert_eq!(dfg.nodes()[0].reads.len(), 0);
1446 assert_eq!(dfg.nodes()[0].writes.len(), 0);
1447 }
1448
1449 #[test]
1452 fn return_with_value() {
1453 let mut func = Function::new("test");
1454 let val = func
1455 .expressions
1456 .append(Expression::Literal(Literal::F32(42.0)));
1457 func.body.push(Statement::Return { value: Some(val) });
1458
1459 let dfg = DataflowGraph::build(&func);
1460 assert_eq!(dfg.node_count(), 1);
1461 assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Return);
1462 assert!(dfg.nodes()[0].reads.contains(&val));
1463 assert_eq!(dfg.nodes()[0].writes.len(), 0);
1464 }
1465
1466 #[test]
1469 fn if_statement_classification() {
1470 let mut func = Function::new("test");
1471 let gv = make_gv(0);
1472 let ptr = func.expressions.append(Expression::GlobalVariable(gv));
1473 let cond = func
1474 .expressions
1475 .append(Expression::Literal(Literal::Bool(true)));
1476 let val = func
1477 .expressions
1478 .append(Expression::Literal(Literal::F32(1.0)));
1479
1480 func.body.push(Statement::If {
1482 condition: cond,
1483 accept: vec![Statement::Store {
1484 pointer: ptr,
1485 value: val,
1486 }],
1487 reject: vec![],
1488 });
1489
1490 let dfg = DataflowGraph::build(&func);
1491 assert_eq!(dfg.node_count(), 1);
1492 assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::If);
1493 assert!(dfg.nodes()[0].reads.contains(&cond));
1495 assert!(dfg.nodes()[0].writes.contains(&ptr));
1497 }
1498
1499 #[test]
1502 fn if_statement_with_reject_branch() {
1503 let mut func = Function::new("test");
1504 let gv0 = make_gv(0);
1505 let gv1 = make_gv(1);
1506 let ptr_a = func.expressions.append(Expression::GlobalVariable(gv0));
1507 let ptr_b = func.expressions.append(Expression::GlobalVariable(gv1));
1508 let cond = func
1509 .expressions
1510 .append(Expression::Literal(Literal::Bool(true)));
1511 let val = func
1512 .expressions
1513 .append(Expression::Literal(Literal::F32(1.0)));
1514
1515 func.body.push(Statement::If {
1516 condition: cond,
1517 accept: vec![Statement::Store {
1518 pointer: ptr_a,
1519 value: val,
1520 }],
1521 reject: vec![Statement::Store {
1522 pointer: ptr_b,
1523 value: val,
1524 }],
1525 });
1526
1527 let dfg = DataflowGraph::build(&func);
1528 assert_eq!(dfg.node_count(), 1);
1529 assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::If);
1530 assert!(dfg.nodes()[0].writes.contains(&ptr_a));
1532 assert!(dfg.nodes()[0].writes.contains(&ptr_b));
1533 }
1534
1535 #[test]
1538 fn loop_statement_classification() {
1539 let mut func = Function::new("test");
1540 let gv = make_gv(0);
1541 let ptr = func.expressions.append(Expression::GlobalVariable(gv));
1542 let val = func
1543 .expressions
1544 .append(Expression::Literal(Literal::F32(1.0)));
1545 let break_cond = func
1546 .expressions
1547 .append(Expression::Literal(Literal::Bool(false)));
1548
1549 func.body.push(Statement::Loop {
1550 body: vec![Statement::Store {
1551 pointer: ptr,
1552 value: val,
1553 }],
1554 continuing: vec![],
1555 break_if: Some(break_cond),
1556 });
1557
1558 let dfg = DataflowGraph::build(&func);
1559 assert_eq!(dfg.node_count(), 1);
1560 assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Loop);
1561 assert!(dfg.nodes()[0].reads.contains(&break_cond));
1563 assert!(dfg.nodes()[0].writes.contains(&ptr));
1565 }
1566
1567 #[test]
1570 fn loop_with_continuing_block() {
1571 let mut func = Function::new("test");
1572 let gv0 = make_gv(0);
1573 let gv1 = make_gv(1);
1574 let ptr_body = func.expressions.append(Expression::GlobalVariable(gv0));
1575 let ptr_cont = func.expressions.append(Expression::GlobalVariable(gv1));
1576 let val = func
1577 .expressions
1578 .append(Expression::Literal(Literal::F32(1.0)));
1579
1580 func.body.push(Statement::Loop {
1581 body: vec![Statement::Store {
1582 pointer: ptr_body,
1583 value: val,
1584 }],
1585 continuing: vec![Statement::Store {
1586 pointer: ptr_cont,
1587 value: val,
1588 }],
1589 break_if: None,
1590 });
1591
1592 let dfg = DataflowGraph::build(&func);
1593 assert_eq!(dfg.node_count(), 1);
1594 assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Loop);
1595 assert!(dfg.nodes()[0].writes.contains(&ptr_body));
1597 assert!(dfg.nodes()[0].writes.contains(&ptr_cont));
1598 }
1599
1600 #[test]
1603 fn call_statement_classification() {
1604 let mut func = Function::new("test");
1605 let arg_val = func
1606 .expressions
1607 .append(Expression::Literal(Literal::F32(1.0)));
1608 let result_val = func
1609 .expressions
1610 .append(Expression::Literal(Literal::F32(0.0)));
1611
1612 let callee_handle = {
1614 let mut funcs = nxpu_ir::Arena::new();
1615 funcs.append(Function::new("callee"))
1616 };
1617
1618 func.body.push(Statement::Call {
1619 function: callee_handle,
1620 arguments: vec![arg_val],
1621 result: Some(result_val),
1622 });
1623
1624 let dfg = DataflowGraph::build(&func);
1625 assert_eq!(dfg.node_count(), 1);
1626 assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Call);
1627 assert!(dfg.nodes()[0].reads.contains(&arg_val));
1628 assert!(dfg.nodes()[0].writes.contains(&result_val));
1629 }
1630
1631 #[test]
1634 fn call_no_result() {
1635 let mut func = Function::new("test");
1636 let arg_val = func
1637 .expressions
1638 .append(Expression::Literal(Literal::F32(1.0)));
1639
1640 let callee_handle = {
1641 let mut funcs = nxpu_ir::Arena::new();
1642 funcs.append(Function::new("callee"))
1643 };
1644
1645 func.body.push(Statement::Call {
1646 function: callee_handle,
1647 arguments: vec![arg_val],
1648 result: None,
1649 });
1650
1651 let dfg = DataflowGraph::build(&func);
1652 assert_eq!(dfg.node_count(), 1);
1653 assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Call);
1654 assert_eq!(dfg.nodes()[0].writes.len(), 0);
1655 }
1656
1657 #[test]
1660 fn atomic_statement_classification() {
1661 let mut func = Function::new("test");
1662 let gv = make_gv(0);
1663 let ptr = func.expressions.append(Expression::GlobalVariable(gv));
1664 let val = func
1665 .expressions
1666 .append(Expression::Literal(Literal::U32(1)));
1667 let result_expr = func
1668 .expressions
1669 .append(Expression::Literal(Literal::U32(0)));
1670
1671 func.body.push(Statement::Atomic {
1672 pointer: ptr,
1673 fun: nxpu_ir::AtomicFunction::Add,
1674 value: val,
1675 result: Some(result_expr),
1676 });
1677
1678 let dfg = DataflowGraph::build(&func);
1679 assert_eq!(dfg.node_count(), 1);
1680 assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Atomic);
1681 assert!(dfg.nodes()[0].reads.contains(&ptr));
1682 assert!(dfg.nodes()[0].reads.contains(&val));
1683 assert!(dfg.nodes()[0].writes.contains(&ptr));
1684 assert!(dfg.nodes()[0].writes.contains(&result_expr));
1685 }
1686
1687 #[test]
1690 fn atomic_exchange_with_compare() {
1691 let mut func = Function::new("test");
1692 let gv = make_gv(0);
1693 let ptr = func.expressions.append(Expression::GlobalVariable(gv));
1694 let val = func
1695 .expressions
1696 .append(Expression::Literal(Literal::U32(1)));
1697 let cmp = func
1698 .expressions
1699 .append(Expression::Literal(Literal::U32(0)));
1700
1701 func.body.push(Statement::Atomic {
1702 pointer: ptr,
1703 fun: nxpu_ir::AtomicFunction::Exchange { compare: Some(cmp) },
1704 value: val,
1705 result: None,
1706 });
1707
1708 let dfg = DataflowGraph::build(&func);
1709 assert_eq!(dfg.node_count(), 1);
1710 assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Atomic);
1711 assert!(dfg.nodes()[0].reads.contains(&cmp));
1713 }
1714
1715 #[test]
1718 fn atomic_no_result() {
1719 let mut func = Function::new("test");
1720 let gv = make_gv(0);
1721 let ptr = func.expressions.append(Expression::GlobalVariable(gv));
1722 let val = func
1723 .expressions
1724 .append(Expression::Literal(Literal::U32(1)));
1725
1726 func.body.push(Statement::Atomic {
1727 pointer: ptr,
1728 fun: nxpu_ir::AtomicFunction::Add,
1729 value: val,
1730 result: None,
1731 });
1732
1733 let dfg = DataflowGraph::build(&func);
1734 assert_eq!(dfg.nodes()[0].writes.len(), 1);
1735 assert!(dfg.nodes()[0].writes.contains(&ptr));
1736 }
1737
1738 #[test]
1741 fn emit_statement_classification() {
1742 let mut func = Function::new("test");
1743
1744 let lit_a = func
1746 .expressions
1747 .append(Expression::Literal(Literal::F32(1.0)));
1748 let lit_b = func
1749 .expressions
1750 .append(Expression::Literal(Literal::F32(2.0)));
1751 let add = func.expressions.append(Expression::Binary {
1752 op: nxpu_ir::BinaryOp::Add,
1753 left: lit_a,
1754 right: lit_b,
1755 });
1756
1757 let range = nxpu_ir::Range::from_index_range(0..3);
1759 func.body.push(Statement::Emit(range));
1760
1761 let dfg = DataflowGraph::build(&func);
1762 assert_eq!(dfg.node_count(), 1);
1763 assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Emit);
1764 assert!(dfg.nodes()[0].writes.contains(&lit_a));
1766 assert!(dfg.nodes()[0].writes.contains(&lit_b));
1767 assert!(dfg.nodes()[0].writes.contains(&add));
1768 }
1769
1770 #[test]
1773 fn barrier_statement_no_reads_writes() {
1774 let mut func = Function::new("test");
1775 func.body
1776 .push(Statement::Barrier(nxpu_ir::Barrier::STORAGE));
1777
1778 let dfg = DataflowGraph::build(&func);
1779 assert_eq!(dfg.node_count(), 1);
1780 assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Barrier);
1781 assert_eq!(dfg.nodes()[0].reads.len(), 0);
1782 assert_eq!(dfg.nodes()[0].writes.len(), 0);
1783 }
1784
1785 #[test]
1788 fn barrier_to_barrier_control_edge() {
1789 let mut func = Function::new("test");
1790 func.body
1791 .push(Statement::Barrier(nxpu_ir::Barrier::STORAGE));
1792 func.body
1793 .push(Statement::Barrier(nxpu_ir::Barrier::WORKGROUP));
1794
1795 let dfg = DataflowGraph::build(&func);
1796 assert_eq!(dfg.node_count(), 2);
1797 let control_edges: Vec<_> = dfg
1799 .edges()
1800 .iter()
1801 .filter(|e| e.kind == DependencyKind::Control)
1802 .collect();
1803 assert!(
1804 control_edges.iter().any(|e| e.from == 0 && e.to == 1),
1805 "expected control edge from barrier[0] to barrier[1]"
1806 );
1807 }
1808
1809 #[test]
1812 fn barrier_no_control_edge_to_break() {
1813 let mut func = Function::new("test");
1814 func.body
1815 .push(Statement::Barrier(nxpu_ir::Barrier::STORAGE));
1816 func.body.push(Statement::Break);
1817
1818 let dfg = DataflowGraph::build(&func);
1819 assert_eq!(dfg.node_count(), 2);
1820 let control_edges: Vec<_> = dfg
1822 .edges()
1823 .iter()
1824 .filter(|e| e.kind == DependencyKind::Control)
1825 .collect();
1826 assert!(
1827 control_edges.is_empty(),
1828 "should not have control edge from barrier to break"
1829 );
1830 }
1831
1832 #[test]
1835 fn explicit_raw_dependency() {
1836 let mut func = Function::new("test");
1837 let gv = make_gv(0);
1838 let ptr = func.expressions.append(Expression::GlobalVariable(gv));
1839 let val = func
1840 .expressions
1841 .append(Expression::Literal(Literal::F32(1.0)));
1842 let load = func.expressions.append(Expression::Load { pointer: ptr });
1843
1844 func.body.push(Statement::Store {
1846 pointer: ptr,
1847 value: val,
1848 });
1849 func.body.push(Statement::Return { value: Some(load) });
1851
1852 let dfg = DataflowGraph::build(&func);
1853 assert_eq!(dfg.node_count(), 2);
1854
1855 let raw_edges: Vec<_> = dfg
1857 .edges()
1858 .iter()
1859 .filter(|e| e.kind == DependencyKind::DataFlow)
1860 .collect();
1861 assert!(
1862 !raw_edges.is_empty(),
1863 "expected RAW dependency: store then read"
1864 );
1865 }
1866
1867 #[test]
1870 fn combined_war_waw_raw_dependencies() {
1871 let mut func = Function::new("test");
1872 let gv = make_gv(0);
1873 let ptr = func.expressions.append(Expression::GlobalVariable(gv));
1874 let val = func
1875 .expressions
1876 .append(Expression::Literal(Literal::F32(1.0)));
1877 let load = func.expressions.append(Expression::Load { pointer: ptr });
1878
1879 func.body.push(Statement::Store {
1881 pointer: ptr,
1882 value: val,
1883 });
1884 func.body.push(Statement::Store {
1889 pointer: ptr,
1890 value: load,
1891 });
1892
1893 let dfg = DataflowGraph::build(&func);
1894 assert_eq!(dfg.node_count(), 2);
1895
1896 let has_raw = dfg
1897 .edges()
1898 .iter()
1899 .any(|e| e.kind == DependencyKind::DataFlow);
1900 let has_waw = dfg
1901 .edges()
1902 .iter()
1903 .any(|e| e.kind == DependencyKind::OutputDependency);
1904 assert!(has_raw, "expected RAW dependency");
1905 assert!(has_waw, "expected WAW dependency");
1906 }
1907
1908 #[test]
1911 fn parallel_groups_with_chain() {
1912 let mut func = Function::new("test");
1913 let gv = make_gv(0);
1914 let ptr = func.expressions.append(Expression::GlobalVariable(gv));
1915 let val = func
1916 .expressions
1917 .append(Expression::Literal(Literal::F32(1.0)));
1918
1919 func.body.push(Statement::Store {
1921 pointer: ptr,
1922 value: val,
1923 });
1924 func.body.push(Statement::Store {
1925 pointer: ptr,
1926 value: val,
1927 });
1928 func.body.push(Statement::Store {
1929 pointer: ptr,
1930 value: val,
1931 });
1932
1933 let dfg = DataflowGraph::build(&func);
1934 let groups = dfg.parallel_groups();
1935
1936 assert_eq!(groups.len(), 3, "expected 3 sequential groups");
1938 for g in &groups {
1939 assert_eq!(g.len(), 1, "each group should have exactly 1 node");
1940 }
1941 }
1942
1943 #[test]
1946 fn parallel_groups_empty() {
1947 let func = Function::new("empty");
1948 let dfg = DataflowGraph::build(&func);
1949 let groups = dfg.parallel_groups();
1950 assert_eq!(groups.len(), 0);
1951 }
1952
1953 #[test]
1956 fn critical_path_diamond_graph() {
1957 let mut func = Function::new("test");
1961 let gv0 = make_gv(0);
1962 let gv1 = make_gv(1);
1963
1964 let ptr_a = func.expressions.append(Expression::GlobalVariable(gv0));
1965 let ptr_b = func.expressions.append(Expression::GlobalVariable(gv1));
1966 let val = func
1967 .expressions
1968 .append(Expression::Literal(Literal::F32(1.0)));
1969 let load_a = func.expressions.append(Expression::Load { pointer: ptr_a });
1970 let load_b = func.expressions.append(Expression::Load { pointer: ptr_b });
1971
1972 func.body.push(Statement::Store {
1974 pointer: ptr_a,
1975 value: val,
1976 });
1977 func.body.push(Statement::Store {
1979 pointer: ptr_b,
1980 value: val,
1981 });
1982 func.body.push(Statement::Store {
1984 pointer: ptr_a,
1985 value: load_a,
1986 });
1987 func.body.push(Statement::Store {
1989 pointer: ptr_b,
1990 value: load_b,
1991 });
1992
1993 let dfg = DataflowGraph::build(&func);
1994 let cp = dfg.critical_path();
1995
1996 assert_eq!(cp.critical_path_length, 2);
1998 assert_eq!(cp.critical_path.len(), 2);
1999 assert_eq!(cp.node_distances.len(), 4);
2000 }
2001
2002 #[test]
2005 fn critical_path_node_distances() {
2006 let mut func = Function::new("test");
2007 let gv = make_gv(0);
2008 let ptr = func.expressions.append(Expression::GlobalVariable(gv));
2009 let val = func
2010 .expressions
2011 .append(Expression::Literal(Literal::F32(1.0)));
2012
2013 func.body.push(Statement::Store {
2015 pointer: ptr,
2016 value: val,
2017 });
2018 func.body.push(Statement::Store {
2019 pointer: ptr,
2020 value: val,
2021 });
2022 func.body.push(Statement::Store {
2023 pointer: ptr,
2024 value: val,
2025 });
2026
2027 let dfg = DataflowGraph::build(&func);
2028 let cp = dfg.critical_path();
2029 assert_eq!(cp.node_distances.len(), 3);
2030 assert_eq!(cp.node_distances[0], 1);
2031 assert_eq!(cp.node_distances[1], 2);
2032 assert_eq!(cp.node_distances[2], 3);
2033 assert_eq!(cp.critical_path, vec![0, 1, 2]);
2034 }
2035
2036 #[test]
2039 fn expression_operands_access_and_select() {
2040 let mut func = Function::new("test");
2041 let gv = make_gv(0);
2042 let ptr = func.expressions.append(Expression::GlobalVariable(gv));
2043 let idx = func
2044 .expressions
2045 .append(Expression::Literal(Literal::U32(0)));
2046 let access = func.expressions.append(Expression::Access {
2047 base: ptr,
2048 index: idx,
2049 });
2050 let val_a = func
2051 .expressions
2052 .append(Expression::Literal(Literal::F32(1.0)));
2053 let val_b = func
2054 .expressions
2055 .append(Expression::Literal(Literal::F32(2.0)));
2056 let cond = func
2057 .expressions
2058 .append(Expression::Literal(Literal::Bool(true)));
2059 let select = func.expressions.append(Expression::Select {
2060 condition: cond,
2061 accept: val_a,
2062 reject: val_b,
2063 });
2064
2065 let range = nxpu_ir::Range::from_index_range(0..func.expressions.len() as u32);
2067 func.body.push(Statement::Emit(range));
2068
2069 let dfg = DataflowGraph::build(&func);
2070 assert_eq!(dfg.node_count(), 1);
2071 assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Emit);
2072 assert!(dfg.nodes()[0].writes.contains(&access));
2074 assert!(dfg.nodes()[0].writes.contains(&select));
2075 }
2076
2077 #[test]
2080 fn expression_operands_splat_as_arraylength_compose_math() {
2081 let mut func = Function::new("test");
2082 let lit = func
2083 .expressions
2084 .append(Expression::Literal(Literal::F32(1.0)));
2085 let splat = func.expressions.append(Expression::Splat {
2086 size: nxpu_ir::VectorSize::Quad,
2087 value: lit,
2088 });
2089 let cast = func.expressions.append(Expression::As {
2090 expr: lit,
2091 kind: nxpu_ir::ScalarKind::Uint,
2092 convert: Some(4),
2093 });
2094 let arr_len = func.expressions.append(Expression::ArrayLength(lit));
2095
2096 let ty_handle = {
2097 let mut types = nxpu_ir::UniqueArena::new();
2098 types.insert(nxpu_ir::Type {
2099 name: None,
2100 inner: nxpu_ir::TypeInner::Vector {
2101 size: nxpu_ir::VectorSize::Quad,
2102 scalar: nxpu_ir::Scalar::F32,
2103 },
2104 })
2105 };
2106 let compose = func.expressions.append(Expression::Compose {
2107 ty: ty_handle,
2108 components: vec![lit],
2109 });
2110
2111 let lit2 = func
2112 .expressions
2113 .append(Expression::Literal(Literal::F32(2.0)));
2114 let lit3 = func
2115 .expressions
2116 .append(Expression::Literal(Literal::F32(3.0)));
2117 let math = func.expressions.append(Expression::Math {
2118 fun: nxpu_ir::MathFunction::Clamp,
2119 arg: lit,
2120 arg1: Some(lit2),
2121 arg2: Some(lit3),
2122 arg3: None,
2123 });
2124
2125 let range = nxpu_ir::Range::from_index_range(0..func.expressions.len() as u32);
2126 func.body.push(Statement::Emit(range));
2127
2128 let dfg = DataflowGraph::build(&func);
2129 assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Emit);
2130 assert!(dfg.nodes()[0].writes.contains(&splat));
2131 assert!(dfg.nodes()[0].writes.contains(&cast));
2132 assert!(dfg.nodes()[0].writes.contains(&arr_len));
2133 assert!(dfg.nodes()[0].writes.contains(&compose));
2134 assert!(dfg.nodes()[0].writes.contains(&math));
2135 }
2136
2137 #[test]
2140 fn expression_operands_swizzle_unary_accessindex() {
2141 let mut func = Function::new("test");
2142 let lit = func
2143 .expressions
2144 .append(Expression::Literal(Literal::F32(1.0)));
2145 let neg = func.expressions.append(Expression::Unary {
2146 op: nxpu_ir::UnaryOp::Negate,
2147 expr: lit,
2148 });
2149 let gv = make_gv(0);
2150 let ptr = func.expressions.append(Expression::GlobalVariable(gv));
2151 let access_idx = func.expressions.append(Expression::AccessIndex {
2152 base: ptr,
2153 index: 0,
2154 });
2155
2156 let vec_lit = func
2157 .expressions
2158 .append(Expression::Literal(Literal::F32(0.0)));
2159 let swizzle = func.expressions.append(Expression::Swizzle {
2160 size: nxpu_ir::VectorSize::Bi,
2161 vector: vec_lit,
2162 pattern: [
2163 nxpu_ir::SwizzleComponent::X,
2164 nxpu_ir::SwizzleComponent::Y,
2165 nxpu_ir::SwizzleComponent::Z,
2166 nxpu_ir::SwizzleComponent::W,
2167 ],
2168 });
2169
2170 let range = nxpu_ir::Range::from_index_range(0..func.expressions.len() as u32);
2171 func.body.push(Statement::Emit(range));
2172
2173 let dfg = DataflowGraph::build(&func);
2174 assert!(dfg.nodes()[0].writes.contains(&neg));
2175 assert!(dfg.nodes()[0].writes.contains(&access_idx));
2176 assert!(dfg.nodes()[0].writes.contains(&swizzle));
2177 }
2178
2179 #[test]
2182 fn expression_operands_leaf_nodes() {
2183 let mut func = Function::new("test");
2184 let _func_arg = func.expressions.append(Expression::FunctionArgument(0));
2185 let gv = make_gv(0);
2186 let _gv_expr = func.expressions.append(Expression::GlobalVariable(gv));
2187
2188 let callee_handle = {
2189 let mut funcs = nxpu_ir::Arena::new();
2190 funcs.append(Function::new("callee"))
2191 };
2192 let _call_result = func
2193 .expressions
2194 .append(Expression::CallResult(callee_handle));
2195
2196 let ty_handle = {
2197 let mut types = nxpu_ir::UniqueArena::new();
2198 types.insert(nxpu_ir::Type {
2199 name: None,
2200 inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::U32),
2201 })
2202 };
2203 let _atomic_result = func.expressions.append(Expression::AtomicResult {
2204 ty: ty_handle,
2205 comparison: false,
2206 });
2207 let _zero_val = func.expressions.append(Expression::ZeroValue(ty_handle));
2208
2209 let range = nxpu_ir::Range::from_index_range(0..func.expressions.len() as u32);
2210 func.body.push(Statement::Emit(range));
2211
2212 let dfg = DataflowGraph::build(&func);
2213 assert_eq!(dfg.node_count(), 1);
2214 }
2217
2218 #[test]
2221 fn topological_sort_diamond_shape() {
2222 let mut func = Function::new("test");
2223 let gv0 = make_gv(0);
2224 let gv1 = make_gv(1);
2225 let ptr_a = func.expressions.append(Expression::GlobalVariable(gv0));
2226 let ptr_b = func.expressions.append(Expression::GlobalVariable(gv1));
2227 let val = func
2228 .expressions
2229 .append(Expression::Literal(Literal::F32(1.0)));
2230 let load_a = func.expressions.append(Expression::Load { pointer: ptr_a });
2231
2232 func.body.push(Statement::Store {
2234 pointer: ptr_a,
2235 value: val,
2236 });
2237 func.body.push(Statement::Store {
2239 pointer: ptr_b,
2240 value: val,
2241 });
2242 func.body.push(Statement::Store {
2244 pointer: ptr_b,
2245 value: load_a,
2246 });
2247
2248 let dfg = DataflowGraph::build(&func);
2249 let order = dfg.topological_sort().unwrap();
2250 assert_eq!(order.len(), 3);
2251
2252 let pos_0 = order.iter().position(|&n| n == 0).unwrap();
2255 let pos_1 = order.iter().position(|&n| n == 1).unwrap();
2256 let pos_2 = order.iter().position(|&n| n == 2).unwrap();
2257 assert!(pos_0 < pos_2);
2258 assert!(pos_1 < pos_2);
2259 }
2260
2261 #[test]
2264 fn parallel_groups_mixed() {
2265 let mut func = Function::new("test");
2266 let gv0 = make_gv(0);
2267 let gv1 = make_gv(1);
2268 let ptr_a = func.expressions.append(Expression::GlobalVariable(gv0));
2269 let ptr_b = func.expressions.append(Expression::GlobalVariable(gv1));
2270 let val_a = func
2271 .expressions
2272 .append(Expression::Literal(Literal::F32(1.0)));
2273 let val_b = func
2274 .expressions
2275 .append(Expression::Literal(Literal::F32(2.0)));
2276
2277 func.body.push(Statement::Store {
2279 pointer: ptr_a,
2280 value: val_a,
2281 });
2282 func.body.push(Statement::Store {
2283 pointer: ptr_b,
2284 value: val_b,
2285 });
2286 func.body.push(Statement::Store {
2288 pointer: ptr_a,
2289 value: val_b,
2290 });
2291
2292 let dfg = DataflowGraph::build(&func);
2293 let groups = dfg.parallel_groups();
2294
2295 assert_ne!(groups.len(), 0);
2297 let first_group = &groups[0];
2298 assert!(
2299 first_group.len() >= 2,
2300 "expected first parallel group to contain independent ops"
2301 );
2302 }
2303
2304 #[test]
2307 fn single_node_graph() {
2308 let mut func = Function::new("test");
2309 func.body.push(Statement::Break);
2310
2311 let dfg = DataflowGraph::build(&func);
2312 assert_eq!(dfg.node_count(), 1);
2313 assert_eq!(dfg.edge_count(), 0);
2314
2315 let order = dfg.topological_sort().unwrap();
2316 assert_eq!(order, vec![0]);
2317
2318 let cp = dfg.critical_path();
2319 assert_eq!(cp.critical_path_length, 1);
2320 assert_eq!(cp.critical_path, vec![0]);
2321
2322 let groups = dfg.parallel_groups();
2323 assert_eq!(groups.len(), 1);
2324 assert_eq!(groups[0], vec![0]);
2325 }
2326
2327 #[test]
2330 fn display_all_edge_kinds() {
2331 let mut func = Function::new("test");
2332 let gv = make_gv(0);
2333 let ptr = func.expressions.append(Expression::GlobalVariable(gv));
2334 let val = func
2335 .expressions
2336 .append(Expression::Literal(Literal::F32(1.0)));
2337 let load = func.expressions.append(Expression::Load { pointer: ptr });
2338
2339 func.body.push(Statement::Store {
2341 pointer: ptr,
2342 value: val,
2343 });
2344 func.body.push(Statement::Store {
2346 pointer: ptr,
2347 value: load,
2348 });
2349
2350 let dfg = DataflowGraph::build(&func);
2351 let s = format!("{dfg}");
2352 assert!(s.contains("RAW") || s.contains("WAW") || s.contains("WAR"));
2353 }
2354
2355 #[test]
2358 fn math_expression_four_args() {
2359 let mut func = Function::new("test");
2360 let a = func
2361 .expressions
2362 .append(Expression::Literal(Literal::F32(1.0)));
2363 let b = func
2364 .expressions
2365 .append(Expression::Literal(Literal::F32(2.0)));
2366 let c = func
2367 .expressions
2368 .append(Expression::Literal(Literal::F32(3.0)));
2369 let d = func
2370 .expressions
2371 .append(Expression::Literal(Literal::F32(4.0)));
2372 let _math = func.expressions.append(Expression::Math {
2373 fun: nxpu_ir::MathFunction::Fma,
2374 arg: a,
2375 arg1: Some(b),
2376 arg2: Some(c),
2377 arg3: Some(d),
2378 });
2379
2380 let range = nxpu_ir::Range::from_index_range(0..func.expressions.len() as u32);
2381 func.body.push(Statement::Emit(range));
2382
2383 let dfg = DataflowGraph::build(&func);
2384 let reads = &dfg.nodes()[0].reads;
2386 assert!(reads.contains(&a), "math should read arg");
2387 assert!(reads.contains(&b), "math should read arg1");
2388 assert!(reads.contains(&c), "math should read arg2");
2389 assert!(reads.contains(&d), "math should read arg3");
2390 }
2391}