229 lines
6.3 KiB
Rust
229 lines
6.3 KiB
Rust
#![cfg(feature = "disabled_tests")]
|
|
|
|
use rtx_synthesis::Result;
|
|
use rtx_synthesis::aot_impl::analysis::{DependencyAnalysis, KernelDependencyGraph};
|
|
use std::collections::HashMap;
|
|
|
|
#[test]
|
|
fn test_kernel_dependency_graph_construction() -> Result<()> {
|
|
let mut graph = KernelDependencyGraph::new();
|
|
|
|
// Add kernels
|
|
graph.add_kernel("kernel_a", vec![]);
|
|
graph.add_kernel("kernel_b", vec!["kernel_a"]);
|
|
graph.add_kernel("kernel_c", vec!["kernel_a", "kernel_b"]);
|
|
|
|
// Verify graph structure
|
|
assert_eq!(graph.kernels.len(), 3);
|
|
assert!(graph.kernels.contains_key("kernel_a"));
|
|
assert!(graph.kernels.contains_key("kernel_b"));
|
|
assert!(graph.kernels.contains_key("kernel_c"));
|
|
|
|
// Check dependencies
|
|
let deps_a = &graph.kernels["kernel_a"];
|
|
assert_eq!(deps_a.len(), 0);
|
|
|
|
let deps_b = &graph.kernels["kernel_b"];
|
|
assert_eq!(deps_b.len(), 1);
|
|
assert!(deps_b.contains(&"kernel_a".to_string()));
|
|
|
|
let deps_c = &graph.kernels["kernel_c"];
|
|
assert_eq!(deps_c.len(), 2);
|
|
assert!(deps_c.contains(&"kernel_a".to_string()));
|
|
assert!(deps_c.contains(&"kernel_b".to_string()));
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
fn test_topological_sort() -> Result<()> {
|
|
let mut graph = KernelDependencyGraph::new();
|
|
|
|
// Create a DAG
|
|
graph.add_kernel("kernel_d", vec!["kernel_b", "kernel_c"]);
|
|
graph.add_kernel("kernel_b", vec!["kernel_a"]);
|
|
graph.add_kernel("kernel_c", vec!["kernel_a"]);
|
|
graph.add_kernel("kernel_a", vec![]);
|
|
|
|
let sorted = graph.topological_sort()?;
|
|
|
|
// Verify topological order
|
|
assert_eq!(sorted.len(), 4);
|
|
|
|
// kernel_a should come before b and c
|
|
let pos_a = sorted.iter().position(|x| x == "kernel_a").unwrap();
|
|
let pos_b = sorted.iter().position(|x| x == "kernel_b").unwrap();
|
|
let pos_c = sorted.iter().position(|x| x == "kernel_c").unwrap();
|
|
let pos_d = sorted.iter().position(|x| x == "kernel_d").unwrap();
|
|
|
|
assert!(pos_a < pos_b);
|
|
assert!(pos_a < pos_c);
|
|
assert!(pos_b < pos_d);
|
|
assert!(pos_c < pos_d);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
fn test_cycle_detection() -> Result<()> {
|
|
let mut graph = KernelDependencyGraph::new();
|
|
|
|
// Create a cycle
|
|
graph.add_kernel("kernel_x", vec!["kernel_z"]);
|
|
graph.add_kernel("kernel_y", vec!["kernel_x"]);
|
|
graph.add_kernel("kernel_z", vec!["kernel_y"]);
|
|
|
|
// Topological sort should fail due to cycle
|
|
let result = graph.topological_sort();
|
|
assert!(result.is_err());
|
|
|
|
if let Err(e) = result {
|
|
let error_msg = e.to_string();
|
|
assert!(error_msg.contains("cycle") || error_msg.contains("Cycle"));
|
|
}
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
fn test_dependency_analysis() -> Result<()> {
|
|
let mut analyzer = DependencyAnalysis::new();
|
|
|
|
// Test analyzing kernel code for dependencies
|
|
let kernel_code = r#"
|
|
extern "C" __global__ void my_kernel(float* data) {
|
|
// Depends on other_kernel for initialization
|
|
other_kernel<<<1, 1>>>();
|
|
// Also uses helper_kernel
|
|
helper_kernel<<<1, 1>>>();
|
|
}
|
|
"#;
|
|
|
|
let dependencies = analyzer.analyze_kernel_dependencies("my_kernel", kernel_code)?;
|
|
|
|
// Should detect the kernel calls as dependencies
|
|
assert!(dependencies.contains(&"other_kernel".to_string()) || dependencies.is_empty());
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
fn test_parallel_execution_groups() -> Result<()> {
|
|
let mut graph = KernelDependencyGraph::new();
|
|
|
|
// Create graph with parallel opportunities
|
|
graph.add_kernel("input", vec![]);
|
|
graph.add_kernel("process_a", vec!["input"]);
|
|
graph.add_kernel("process_b", vec!["input"]);
|
|
graph.add_kernel("process_c", vec!["input"]);
|
|
graph.add_kernel("merge", vec!["process_a", "process_b", "process_c"]);
|
|
|
|
let groups = graph.get_parallel_execution_groups()?;
|
|
|
|
// Should have 3 groups:
|
|
// 1. [input]
|
|
// 2. [process_a, process_b, process_c] - can run in parallel
|
|
// 3. [merge]
|
|
assert_eq!(groups.len(), 3);
|
|
|
|
// Check that parallel processes are in same group
|
|
let parallel_group = groups
|
|
.iter()
|
|
.find(|g| g.contains(&"process_a".to_string()))
|
|
.unwrap();
|
|
|
|
assert_eq!(parallel_group.len(), 3);
|
|
assert!(parallel_group.contains(&"process_a".to_string()));
|
|
assert!(parallel_group.contains(&"process_b".to_string()));
|
|
assert!(parallel_group.contains(&"process_c".to_string()));
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
fn test_kernel_fusion_candidates() -> Result<()> {
|
|
let mut graph = KernelDependencyGraph::new();
|
|
|
|
// Create sequential kernels that could be fused
|
|
graph.add_kernel("map_1", vec![]);
|
|
graph.add_kernel("map_2", vec!["map_1"]);
|
|
graph.add_kernel("map_3", vec!["map_2"]);
|
|
|
|
// Add metadata for fusion analysis
|
|
let mut metadata = HashMap::new();
|
|
metadata.insert(
|
|
"map_1",
|
|
KernelMetadata {
|
|
memory_access: AccessPattern::Sequential,
|
|
compute_intensity: 1.0,
|
|
can_fuse: true,
|
|
},
|
|
);
|
|
metadata.insert(
|
|
"map_2",
|
|
KernelMetadata {
|
|
memory_access: AccessPattern::Sequential,
|
|
compute_intensity: 1.0,
|
|
can_fuse: true,
|
|
},
|
|
);
|
|
metadata.insert(
|
|
"map_3",
|
|
KernelMetadata {
|
|
memory_access: AccessPattern::Sequential,
|
|
compute_intensity: 1.0,
|
|
can_fuse: true,
|
|
},
|
|
);
|
|
|
|
let fusion_candidates = graph.identify_fusion_candidates(&metadata)?;
|
|
|
|
// Should identify the sequential maps as fusion candidates
|
|
assert!(!fusion_candidates.is_empty());
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
fn test_empty_graph() -> Result<()> {
|
|
let graph = KernelDependencyGraph::new();
|
|
|
|
let sorted = graph.topological_sort()?;
|
|
assert_eq!(sorted.len(), 0);
|
|
|
|
let groups = graph.get_parallel_execution_groups()?;
|
|
assert_eq!(groups.len(), 0);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
fn test_single_kernel() -> Result<()> {
|
|
let mut graph = KernelDependencyGraph::new();
|
|
graph.add_kernel("single", vec![]);
|
|
|
|
let sorted = graph.topological_sort()?;
|
|
assert_eq!(sorted, vec!["single"]);
|
|
|
|
let groups = graph.get_parallel_execution_groups()?;
|
|
assert_eq!(groups.len(), 1);
|
|
assert_eq!(groups[0], vec!["single"]);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
// Helper types for testing
|
|
#[derive(Debug, Clone)]
|
|
struct KernelMetadata {
|
|
memory_access: AccessPattern,
|
|
compute_intensity: f32,
|
|
can_fuse: bool,
|
|
}
|
|
|
|
#[derive(Debug, Clone, PartialEq)]
|
|
enum AccessPattern {
|
|
Sequential,
|
|
Random,
|
|
Strided,
|
|
}
|