Files
rustytorch/crates/specialized/rtx-synthesis/tests/aot_analysis_test.rs
T
2026-03-04 00:08:42 +00:00

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,
}