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

273 lines
6.9 KiB
Rust

use rtx_polygraph::{
cache::{CacheKey, CachedKernel, KernelCache},
error::PolygraphError,
fusion::FusionOpportunity,
ir::{DataType, IRNode, IRNodeType, NodeId, Shape},
};
#[test]
fn test_cache_key_generation() {
let node = IRNode::new(
NodeId(1),
IRNodeType::MatMul {
transpose_a: false,
transpose_b: true,
},
vec![NodeId(0)],
vec![DataType::F32],
vec![Shape::new(vec![64, 128])],
);
let key = CacheKey::from_node(&node);
assert!(!key.hash().is_empty());
// Same configuration should produce same key
let node2 = IRNode::new(
NodeId(999), // Different ID shouldn't affect key
IRNodeType::MatMul {
transpose_a: false,
transpose_b: true,
},
vec![NodeId(888)],
vec![DataType::F32],
vec![Shape::new(vec![64, 128])],
);
let key2 = CacheKey::from_node(&node2);
assert_eq!(key.hash(), key2.hash());
}
#[test]
fn test_cache_key_sensitivity() {
let node1 = IRNode::new(
NodeId(1),
IRNodeType::MatMul {
transpose_a: false,
transpose_b: false, // Different from above
},
vec![NodeId(0)],
vec![DataType::F32],
vec![Shape::new(vec![64, 128])],
);
let node2 = IRNode::new(
NodeId(1),
IRNodeType::MatMul {
transpose_a: false,
transpose_b: true, // Different transpose setting
},
vec![NodeId(0)],
vec![DataType::F32],
vec![Shape::new(vec![64, 128])],
);
let key1 = CacheKey::from_node(&node1);
let key2 = CacheKey::from_node(&node2);
assert_ne!(key1.hash(), key2.hash());
}
#[test]
fn test_kernel_cache_hit() {
let mut cache = KernelCache::new();
let node = IRNode::new(
NodeId(1),
IRNodeType::MatMul {
transpose_a: false,
transpose_b: true,
},
vec![NodeId(0)],
vec![DataType::F32],
vec![Shape::new(vec![64, 128])],
);
let key = CacheKey::from_node(&node);
// First lookup should be a miss
assert!(cache.get(&key).is_none());
// Insert kernel
let kernel = CachedKernel::new(
"test_kernel".to_string(),
vec![0u8; 1024], // Mock binary
std::time::Duration::from_millis(150),
);
cache.insert(key.clone(), kernel.clone());
// Second lookup should be a hit
let cached = cache.get(&key).unwrap();
assert_eq!(cached.name(), kernel.name());
assert_eq!(cached.compile_time(), kernel.compile_time());
}
#[test]
fn test_kernel_cache_miss_different_shapes() {
let mut cache = KernelCache::new();
let node1 = IRNode::new(
NodeId(1),
IRNodeType::MatMul {
transpose_a: false,
transpose_b: true,
},
vec![NodeId(0)],
vec![DataType::F32],
vec![Shape::new(vec![64, 128])],
);
let node2 = IRNode::new(
NodeId(1),
IRNodeType::MatMul {
transpose_a: false,
transpose_b: true,
},
vec![NodeId(0)],
vec![DataType::F32],
vec![Shape::new(vec![32, 128])], // Different shape
);
let key1 = CacheKey::from_node(&node1);
let key2 = CacheKey::from_node(&node2);
let kernel = CachedKernel::new(
"test_kernel".to_string(),
vec![0u8; 1024],
std::time::Duration::from_millis(150),
);
cache.insert(key1.clone(), kernel);
// Different shape should be cache miss
assert!(cache.get(&key1).is_some());
assert!(cache.get(&key2).is_none());
}
#[test]
fn test_fusion_cache_key() {
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])],
);
let fusion_op = FusionOpportunity::new(
vec![node1, node2],
rtx_polygraph::fusion::FusionType::DenseDense,
2.5, // speedup estimate
);
let key = CacheKey::from_fusion(&fusion_op);
assert!(!key.hash().is_empty());
}
#[test]
fn test_cache_eviction_policy() {
let mut cache = KernelCache::with_capacity(2); // Small capacity for testing
let nodes = (0..3)
.map(|i| {
IRNode::new(
NodeId(i as u32),
IRNodeType::MatMul {
transpose_a: false,
transpose_b: false,
},
vec![NodeId((i + 10) as u32)],
vec![DataType::F32],
vec![Shape::new(vec![64 + i, 128])], // Different shapes
)
})
.collect::<Vec<_>>();
let keys: Vec<_> = nodes.iter().map(CacheKey::from_node).collect();
// Fill cache to capacity
for (i, key) in keys.iter().take(2).enumerate() {
let kernel = CachedKernel::new(
format!("kernel_{}", i),
vec![i as u8; 1024],
std::time::Duration::from_millis(100),
);
cache.insert(key.clone(), kernel);
}
assert_eq!(cache.len(), 2);
assert!(cache.get(&keys[0]).is_some());
assert!(cache.get(&keys[1]).is_some());
// Insert third item should evict first (LRU)
let kernel = CachedKernel::new(
"kernel_2".to_string(),
vec![2u8; 1024],
std::time::Duration::from_millis(100),
);
cache.insert(keys[2].clone(), kernel);
assert_eq!(cache.len(), 2);
assert!(cache.get(&keys[0]).is_none()); // Evicted
assert!(cache.get(&keys[1]).is_some()); // Still present
assert!(cache.get(&keys[2]).is_some()); // Newly inserted
}
#[test]
fn test_cache_statistics() {
let mut cache = KernelCache::new();
let node = IRNode::new(
NodeId(1),
IRNodeType::MatMul {
transpose_a: false,
transpose_b: true,
},
vec![NodeId(0)],
vec![DataType::F32],
vec![Shape::new(vec![64, 128])],
);
let key = CacheKey::from_node(&node);
// Initial stats
let stats = cache.statistics();
assert_eq!(stats.hits(), 0);
assert_eq!(stats.misses(), 0);
assert_eq!(stats.hit_rate(), 0.0);
// First lookup - miss
cache.get(&key);
let stats = cache.statistics();
assert_eq!(stats.misses(), 1);
assert_eq!(stats.hit_rate(), 0.0);
// Insert and lookup - hit
let kernel = CachedKernel::new(
"test_kernel".to_string(),
vec![0u8; 1024],
std::time::Duration::from_millis(150),
);
cache.insert(key.clone(), kernel);
cache.get(&key);
let stats = cache.statistics();
assert_eq!(stats.hits(), 1);
assert_eq!(stats.misses(), 1);
assert_eq!(stats.hit_rate(), 0.5);
}