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

314 lines
7.9 KiB
Rust

use rtx_polygraph::{
error::PolygraphError,
ir::{DataType, IRGraph, IRNode, IRNodeType, NodeId, Shape},
optimizer::{DeadCodeEliminationPass, FusionPass, MemoryOptimizationPass, OptimizationPass},
};
#[test]
fn test_fusion_pass_application() {
let mut graph = IRGraph::new();
// Create fusible sequence: MatMul -> MatMul
let node1 = IRNode::new(
NodeId(1),
IRNodeType::MatMul {
transpose_a: false,
transpose_b: false,
},
vec![NodeId(0)],
vec![DataType::F32],
vec![Shape::new(vec![64, 128])],
);
let node2 = IRNode::new(
NodeId(2),
IRNodeType::MatMul {
transpose_a: false,
transpose_b: false,
},
vec![NodeId(1)],
vec![DataType::F32],
vec![Shape::new(vec![64, 256])],
);
graph.add_node(node1);
graph.add_node(node2);
let original_node_count = graph.node_count();
assert_eq!(original_node_count, 2);
let mut pass = FusionPass::new();
let result = pass.apply(&mut graph);
assert!(result.is_ok());
let optimized_node_count = graph.node_count();
// After fusion, should have fewer nodes
assert!(optimized_node_count < original_node_count);
}
#[test]
fn test_dead_code_elimination_pass() {
let mut graph = IRGraph::new();
// Create a node with no users
let unused_node = IRNode::new(
NodeId(1),
IRNodeType::MatMul {
transpose_a: false,
transpose_b: false,
},
vec![NodeId(0)],
vec![DataType::F32],
vec![Shape::new(vec![64, 128])],
);
// Create a used node
let used_node = IRNode::new(
NodeId(2),
IRNodeType::MatMul {
transpose_a: false,
transpose_b: false,
},
vec![NodeId(0)],
vec![DataType::F32],
vec![Shape::new(vec![64, 256])],
);
// Mark used_node as output
graph.add_node(unused_node);
graph.add_node(used_node);
graph.mark_output(NodeId(2));
let original_count = graph.node_count();
assert_eq!(original_count, 2);
let mut pass = DeadCodeEliminationPass::new();
let result = pass.apply(&mut graph);
assert!(result.is_ok());
// Dead node should be eliminated
let final_count = graph.node_count();
assert_eq!(final_count, 1);
assert!(graph.contains_node(NodeId(2))); // Used node remains
assert!(!graph.contains_node(NodeId(1))); // Unused node removed
}
#[test]
fn test_memory_optimization_pass() {
let mut graph = IRGraph::new();
// Create nodes with overlapping lifetimes that can share memory
let node1 = IRNode::new(
NodeId(1),
IRNodeType::MatMul {
transpose_a: false,
transpose_b: false,
},
vec![NodeId(0)],
vec![DataType::F32],
vec![Shape::new(vec![1024, 1024])], // Large tensor
);
let node2 = IRNode::new(
NodeId(2),
IRNodeType::Add,
vec![NodeId(1)],
vec![DataType::F32],
vec![Shape::new(vec![1024, 1024])],
);
let node3 = IRNode::new(
NodeId(3),
IRNodeType::MatMul {
transpose_a: false,
transpose_b: false,
},
vec![NodeId(2)],
vec![DataType::F32],
vec![Shape::new(vec![1024, 512])],
);
graph.add_node(node1);
graph.add_node(node2);
graph.add_node(node3);
graph.mark_output(NodeId(3));
let original_memory = graph.estimated_memory_usage();
let mut pass = MemoryOptimizationPass::new();
let result = pass.apply(&mut graph);
assert!(result.is_ok());
let optimized_memory = graph.estimated_memory_usage();
assert!(optimized_memory < original_memory);
}
#[test]
fn test_optimization_pass_chain() {
let mut graph = IRGraph::new();
// Create complex graph with fusion opportunities and dead code
let nodes = vec![
IRNode::new(
NodeId(1),
IRNodeType::MatMul {
transpose_a: false,
transpose_b: false,
},
vec![NodeId(0)],
vec![DataType::F32],
vec![Shape::new(vec![64, 128])],
),
IRNode::new(
NodeId(2),
IRNodeType::MatMul {
transpose_a: false,
transpose_b: false,
},
vec![NodeId(1)],
vec![DataType::F32],
vec![Shape::new(vec![64, 256])],
),
IRNode::new(
NodeId(3),
IRNodeType::Add, // Unused node
vec![NodeId(0)],
vec![DataType::F32],
vec![Shape::new(vec![64, 64])],
),
IRNode::new(
NodeId(4),
IRNodeType::ReLU,
vec![NodeId(2)],
vec![DataType::F32],
vec![Shape::new(vec![64, 256])],
),
];
for node in nodes {
graph.add_node(node);
}
graph.mark_output(NodeId(4));
let original_count = graph.node_count();
assert_eq!(original_count, 4);
// Apply optimization passes in sequence
let mut fusion_pass = FusionPass::new();
let mut dce_pass = DeadCodeEliminationPass::new();
let mut memory_pass = MemoryOptimizationPass::new();
fusion_pass.apply(&mut graph).expect("Fusion pass failed");
dce_pass.apply(&mut graph).expect("DCE pass failed");
memory_pass.apply(&mut graph).expect("Memory pass failed");
let final_count = graph.node_count();
assert!(final_count < original_count);
}
#[test]
fn test_optimization_pass_preserves_semantics() {
let mut graph = IRGraph::new();
let node1 = IRNode::new(
NodeId(1),
IRNodeType::MatMul {
transpose_a: false,
transpose_b: false,
},
vec![NodeId(0)],
vec![DataType::F32],
vec![Shape::new(vec![64, 128])],
);
let node2 = IRNode::new(
NodeId(2),
IRNodeType::Add,
vec![NodeId(1)],
vec![DataType::F32],
vec![Shape::new(vec![64, 128])],
);
graph.add_node(node1);
graph.add_node(node2);
graph.mark_output(NodeId(2));
// Capture original semantics
let original_outputs = graph.output_shapes();
let original_dtypes = graph.output_dtypes();
let mut pass = FusionPass::new();
pass.apply(&mut graph).expect("Pass failed");
// Verify semantics preserved
let final_outputs = graph.output_shapes();
let final_dtypes = graph.output_dtypes();
assert_eq!(original_outputs, final_outputs);
assert_eq!(original_dtypes, final_dtypes);
}
#[test]
fn test_optimization_pass_error_handling() {
let mut graph = IRGraph::new();
// Create invalid graph structure
let node = IRNode::new(
NodeId(1),
IRNodeType::MatMul {
transpose_a: false,
transpose_b: false,
},
vec![NodeId(999)], // Non-existent input
vec![DataType::F32],
vec![Shape::new(vec![64, 128])],
);
graph.add_node(node);
graph.mark_output(NodeId(1));
let mut pass = FusionPass::new();
let result = pass.apply(&mut graph);
assert!(result.is_err());
assert!(matches!(
result.unwrap_err(),
PolygraphError::InvalidGraphStructure { .. }
));
}
#[test]
fn test_custom_optimization_pass() {
struct NoOpPass;
impl OptimizationPass for NoOpPass {
fn name(&self) -> &str {
"noop"
}
fn apply(&mut self, graph: &mut IRGraph) -> Result<(), PolygraphError> {
// Do nothing - just validate we can implement custom passes
Ok(())
}
}
let mut graph = IRGraph::new();
let node = IRNode::new(
NodeId(1),
IRNodeType::Add,
vec![NodeId(0)],
vec![DataType::F32],
vec![Shape::new(vec![64, 128])],
);
graph.add_node(node);
let mut pass = NoOpPass;
let result = pass.apply(&mut graph);
assert!(result.is_ok());
assert_eq!(pass.name(), "noop");
}