Files
rustytorch/crates/production/rtx-inference/tests/inference_engine.rs
T
2026-03-04 00:08:42 +00:00

684 lines
22 KiB
Rust

//! Comprehensive tests for the core InferenceEngine
//!
//! This test suite validates the complete inference engine functionality
//! including model loading, graph optimization, batch processing, and
//! integration with all subsystems.
use futures::future;
use rtx_inference::*;
use rtx_tensor::{DType, Device, Shape, Tensor};
use std::collections::HashMap;
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::Mutex;
use tokio_stream::StreamExt;
use uuid::Uuid;
/// Mock model weights for testing
#[derive(Debug, Clone)]
struct MockModelWeights {
layers: HashMap<String, Tensor>,
config: ModelConfig,
}
fn create_test_config() -> InferenceEngineConfig {
InferenceEngineConfig {
model_path: "test_model.safetensors".to_string(),
device: Device::cpu(),
max_batch_size: 32,
max_sequence_length: 2048,
memory_pool_size: 1024 * 1024 * 1024, // 1GB
optimization_level: 2,
enable_caching: true,
enable_quantization: false,
request_manager_config: RequestManagerConfig::default(),
scheduler_config: BatchSchedulerConfig::default(),
cache_config: KvCacheConfig::default(),
}
}
impl MockModelWeights {
fn new() -> InferenceResult<Self> {
let config = ModelConfig {
vocab_size: 32000,
hidden_size: 4096,
num_layers: 32,
num_heads: 32,
max_position_embeddings: 2048,
layer_norm_epsilon: 1e-6,
};
let mut layers = HashMap::new();
// Create mock embedding layer
let embedding_weight =
Tensor::randn(&[config.vocab_size, config.hidden_size], &Device::cpu())?;
layers.insert("embedding.weight".to_string(), embedding_weight);
// Create mock transformer layers
for layer_idx in 0..config.num_layers {
// Attention weights
let attn_q = Tensor::randn(&[config.hidden_size, config.hidden_size], &Device::cpu())?;
let attn_k = Tensor::randn(&[config.hidden_size, config.hidden_size], &Device::cpu())?;
let attn_v = Tensor::randn(&[config.hidden_size, config.hidden_size], &Device::cpu())?;
let attn_o = Tensor::randn(&[config.hidden_size, config.hidden_size], &Device::cpu())?;
layers.insert(
format!("layers.{}.attention.q_proj.weight", layer_idx),
attn_q,
);
layers.insert(
format!("layers.{}.attention.k_proj.weight", layer_idx),
attn_k,
);
layers.insert(
format!("layers.{}.attention.v_proj.weight", layer_idx),
attn_v,
);
layers.insert(
format!("layers.{}.attention.o_proj.weight", layer_idx),
attn_o,
);
// MLP weights
let mlp_gate = Tensor::randn(
&[config.hidden_size, config.hidden_size * 4],
&Device::cpu(),
)?;
let mlp_up = Tensor::randn(
&[config.hidden_size, config.hidden_size * 4],
&Device::cpu(),
)?;
let mlp_down = Tensor::randn(
&[config.hidden_size * 4, config.hidden_size],
&Device::cpu(),
)?;
layers.insert(
format!("layers.{}.mlp.gate_proj.weight", layer_idx),
mlp_gate,
);
layers.insert(format!("layers.{}.mlp.up_proj.weight", layer_idx), mlp_up);
layers.insert(
format!("layers.{}.mlp.down_proj.weight", layer_idx),
mlp_down,
);
// Layer norm weights
let input_layernorm = Tensor::ones(&[config.hidden_size], &Device::cpu())?;
let post_attention_layernorm = Tensor::ones(&[config.hidden_size], &Device::cpu())?;
layers.insert(
format!("layers.{}.input_layernorm.weight", layer_idx),
input_layernorm,
);
layers.insert(
format!("layers.{}.post_attention_layernorm.weight", layer_idx),
post_attention_layernorm,
);
}
// Output layer
let lm_head = Tensor::randn(&[config.hidden_size, config.vocab_size], &Device::cpu())?;
layers.insert("lm_head.weight".to_string(), lm_head);
Ok(Self { layers, config })
}
fn get_layer(&self, name: &str) -> Option<&Tensor> {
self.layers.get(name)
}
fn layer_count(&self) -> usize {
self.layers.len()
}
}
// Core InferenceEngine tests - these should fail initially
#[tokio::test]
async fn test_inference_engine_creation() {
let config = create_test_config();
let engine = InferenceEngine::new(config.clone()).await;
assert!(
engine.is_ok(),
"Failed to create inference engine: {:?}",
engine.err()
);
let engine = engine.unwrap();
assert_eq!(engine.device(), &config.device);
assert_eq!(engine.max_batch_size(), config.max_batch_size);
}
#[tokio::test]
async fn test_model_loading_and_registration() {
let config = create_test_config();
let mut engine = InferenceEngine::new(config).await.unwrap();
// Create mock model weights
let mock_weights = MockModelWeights::new().unwrap();
// Load model should fail initially (not implemented)
let result = engine
.load_model("test_model", &mock_weights.layers, &mock_weights.config)
.await;
assert!(result.is_ok(), "Model loading should succeed");
// Check model is registered
let models = engine.list_models().await;
assert!(models.contains(&"test_model".to_string()));
// Get model info
let info = engine.get_model_info("test_model").await.unwrap();
assert_eq!(info.name, "test_model");
assert_eq!(info.layer_count, mock_weights.layer_count());
assert!(info.memory_usage > 0);
}
#[tokio::test]
async fn test_single_inference_request() {
let config = create_test_config();
let mut engine = InferenceEngine::new(config).await.unwrap();
// Load model
let mock_weights = MockModelWeights::new().unwrap();
engine
.load_model("test_model", &mock_weights.layers, &mock_weights.config)
.await
.unwrap();
// Create inference request
let input_tokens = vec![1, 2, 3, 4, 5];
let request = InferenceRequest::new("test_model".to_string(), input_tokens.clone(), 10);
let request_id = request.id.clone();
// Run inference
let start_time = Instant::now();
let result = engine.infer(request).await;
let inference_time = start_time.elapsed();
assert!(result.is_ok(), "Inference should succeed");
let result = result.unwrap();
// Validate output
assert_eq!(result.request_id, request_id);
assert!(!result.output_tokens.is_empty());
assert!(result.output_tokens.len() <= 10);
assert!(matches!(
result.finish_reason,
FinishReason::Length | FinishReason::EndOfSequence
));
// Validate performance metrics
let metrics = result.metrics.unwrap();
assert!(metrics.total_time > Duration::ZERO);
assert!(metrics.generation_time > Duration::ZERO);
assert_eq!(metrics.input_token_count, input_tokens.len());
assert_eq!(metrics.output_token_count, result.output_tokens.len());
assert!(metrics.tokens_per_second > 0.0);
println!(
"Inference completed in {:?}, {:.2} tokens/sec",
inference_time, metrics.tokens_per_second
);
}
#[tokio::test]
async fn test_batch_inference() {
let mut config = create_test_config();
config.max_batch_size = 4;
let mut engine = InferenceEngine::new(config).await.unwrap();
// Load model
let mock_weights = MockModelWeights::new().unwrap();
engine
.load_model("test_model", &mock_weights.layers, &mock_weights.config)
.await
.unwrap();
// Create multiple requests
let requests = vec![
InferenceRequest::new("test_model".to_string(), vec![1, 2, 3], 5),
InferenceRequest::new("test_model".to_string(), vec![4, 5, 6, 7], 5),
InferenceRequest::new("test_model".to_string(), vec![8, 9], 5),
InferenceRequest::new("test_model".to_string(), vec![10, 11, 12, 13, 14], 5),
];
// Submit all requests concurrently
let mut handles = vec![];
for request in requests {
let engine_clone = engine.clone();
let handle = tokio::spawn(async move { engine_clone.infer(request).await });
handles.push(handle);
}
// Wait for all results
let results: Vec<_> = future::join_all(handles).await;
// Validate all succeeded
for result in results {
let inference_result = result.unwrap().unwrap();
assert!(!inference_result.output_tokens.is_empty());
assert!(inference_result.metrics.is_some());
}
}
#[tokio::test]
async fn test_streaming_inference() {
let config = create_test_config();
let mut engine = InferenceEngine::new(config).await.unwrap();
// Load model
let mock_weights = MockModelWeights::new().unwrap();
engine
.load_model("test_model", &mock_weights.layers, &mock_weights.config)
.await
.unwrap();
// Create streaming request
let mut request = InferenceRequest::new("test_model".to_string(), vec![1, 2, 3], 20);
request.stream = true;
// Start streaming inference
let mut stream = engine.infer_stream(request).await.unwrap();
let mut token_count = 0;
let mut total_time = Duration::ZERO;
// Collect streaming tokens
while let Some(token_result) = stream.next().await {
assert!(token_result.is_ok());
let token = token_result.unwrap();
assert!(token.id > 0);
assert!(!token.text.is_empty());
assert!(token.probability > 0.0);
token_count += 1;
total_time += token.generation_time;
if token_count >= 20 {
break;
}
}
assert!(token_count > 0);
assert!(total_time > Duration::ZERO);
println!(
"Generated {} tokens via streaming in {:?}",
token_count, total_time
);
}
#[tokio::test]
async fn test_memory_management() {
let mut config = create_test_config();
config.memory_pool_size = 100 * 1024 * 1024; // 100MB
let mut engine = InferenceEngine::new(config).await.unwrap();
// Load model
let mock_weights = MockModelWeights::new().unwrap();
engine
.load_model("test_model", &mock_weights.layers, &mock_weights.config)
.await
.unwrap();
// Get initial memory stats
let initial_stats = engine.get_memory_stats().await;
assert!(initial_stats.total_allocated > 0);
assert!(initial_stats.total_allocated <= 100 * 1024 * 1024);
// Run inference to allocate memory
let request = InferenceRequest::new("test_model".to_string(), vec![1, 2, 3], 10);
let _result = engine.infer(request).await.unwrap();
// Check memory usage increased
let after_inference_stats = engine.get_memory_stats().await;
assert!(after_inference_stats.total_allocated >= initial_stats.total_allocated);
// Force garbage collection
engine.gc().await.unwrap();
// Memory should be reclaimed
let after_gc_stats = engine.get_memory_stats().await;
assert!(after_gc_stats.total_allocated <= after_inference_stats.total_allocated);
}
#[tokio::test]
async fn test_graph_optimization() {
let mut config = create_test_config();
config.optimization_level = 3; // Maximum optimization
let mut engine = InferenceEngine::new(config).await.unwrap();
// Load model
let mock_weights = MockModelWeights::new().unwrap();
engine
.load_model("test_model", &mock_weights.layers, &mock_weights.config)
.await
.unwrap();
// Get optimization stats
let stats = engine.get_optimization_stats("test_model").await.unwrap();
assert!(stats.optimizations_applied > 0);
assert!(stats.nodes_eliminated > 0);
assert!(stats.memory_saved > 0);
assert!(stats.estimated_speedup > 1.0);
// Compare optimized vs unoptimized performance
let request = InferenceRequest::new("test_model".to_string(), vec![1, 2, 3, 4, 5], 10);
// Measure with optimization
let start = Instant::now();
let _optimized_result = engine.infer(request.clone()).await.unwrap();
let optimized_time = start.elapsed();
// Disable optimization temporarily
engine
.set_optimization_level("test_model", 0)
.await
.unwrap();
// Measure without optimization
let start = Instant::now();
let _unoptimized_result = engine.infer(request.clone()).await.unwrap();
let unoptimized_time = start.elapsed();
// Optimized should be faster
assert!(
optimized_time < unoptimized_time,
"Optimized inference ({:?}) should be faster than unoptimized ({:?})",
optimized_time,
unoptimized_time
);
}
#[tokio::test]
async fn test_quantization_integration() {
let mut config = create_test_config();
config.enable_quantization = true;
let mut engine = InferenceEngine::new(config).await.unwrap();
// Load model
let mock_weights = MockModelWeights::new().unwrap();
engine
.load_model("test_model", &mock_weights.layers, &mock_weights.config)
.await
.unwrap();
// Apply quantization
let quant_config = QuantizationConfig::new()
.with_scheme(QuantizationScheme::INT8)
.with_calibration_samples(100);
engine
.quantize_model("test_model", quant_config)
.await
.unwrap();
// Check quantization was applied
let model_info = engine.get_model_info("test_model").await.unwrap();
assert!(model_info.is_quantized);
assert!(model_info.memory_usage < model_info.original_memory_usage);
// Inference should still work
let request = InferenceRequest::new("test_model".to_string(), vec![1, 2, 3], 5);
let result = engine.infer(request).await.unwrap();
assert!(!result.output_tokens.is_empty());
}
#[tokio::test]
async fn test_performance_monitoring() {
let config = create_test_config();
let mut engine = InferenceEngine::new(config).await.unwrap();
// Load model
let mock_weights = MockModelWeights::new().unwrap();
engine
.load_model("test_model", &mock_weights.layers, &mock_weights.config)
.await
.unwrap();
// Run multiple inferences to generate metrics
for i in 0..10 {
let request = InferenceRequest::new("test_model".to_string(), vec![1, 2, i as i32], 5);
let _result = engine.infer(request).await.unwrap();
}
// Get performance metrics
let metrics = engine.get_performance_metrics().await;
assert!(metrics.total_requests >= 10);
assert!(metrics.total_tokens_generated > 0);
assert!(metrics.average_latency > Duration::ZERO);
assert!(metrics.average_throughput > 0.0);
assert!(metrics.success_rate >= 0.9); // Should be high
// Get model-specific metrics
let model_metrics = engine.get_model_metrics("test_model").await.unwrap();
assert!(model_metrics.inference_count >= 10);
assert!(model_metrics.total_inference_time > Duration::ZERO);
assert!(model_metrics.average_tokens_per_request > 0.0);
}
#[tokio::test]
async fn test_error_handling() {
let config = create_test_config();
let mut engine = InferenceEngine::new(config).await.unwrap();
// Test inference without model
let request = InferenceRequest::new("nonexistent_model".to_string(), vec![1, 2, 3], 5);
let result = engine.infer(request).await;
assert!(result.is_err());
assert!(matches!(
result.unwrap_err(),
InferenceError::ModelNotFound { .. }
));
// Load model
let mock_weights = MockModelWeights::new().unwrap();
engine
.load_model("test_model", &mock_weights.layers, &mock_weights.config)
.await
.unwrap();
// Test with invalid input
let invalid_request = InferenceRequest::new("test_model".to_string(), vec![], 5); // Empty input
let result = engine.infer(invalid_request).await;
assert!(result.is_err());
assert!(matches!(
result.unwrap_err(),
InferenceError::InvalidRequest { .. }
));
// Test with too large input
let large_input = vec![1i32; 10000]; // Exceed max sequence length
let large_request = InferenceRequest::new("test_model".to_string(), large_input, 5);
let result = engine.infer(large_request).await;
assert!(result.is_err());
assert!(matches!(
result.unwrap_err(),
InferenceError::InvalidRequest { .. }
));
}
#[tokio::test]
async fn test_concurrent_inference() {
let mut config = create_test_config();
config.max_batch_size = 8;
let mut engine_temp = InferenceEngine::new(config).await.unwrap();
// Load model
let mock_weights = MockModelWeights::new().unwrap();
engine_temp
.load_model("test_model", &mock_weights.layers, &mock_weights.config)
.await
.unwrap();
let engine = Arc::new(engine_temp);
// Run concurrent inferences
let mut handles = vec![];
for i in 0..20 {
let engine_clone = engine.clone();
let handle = tokio::spawn(async move {
let request = InferenceRequest::new("test_model".to_string(), vec![i, i + 1, i + 2], 3);
engine_clone.infer(request).await
});
handles.push(handle);
}
// Wait for all to complete
let results = future::join_all(handles).await;
// Check all succeeded
let mut success_count = 0;
for result in results {
if result.unwrap().is_ok() {
success_count += 1;
}
}
assert_eq!(
success_count, 20,
"All concurrent inferences should succeed"
);
// Check performance metrics show concurrent processing
let metrics = engine.get_performance_metrics().await;
assert!(metrics.concurrent_requests_handled > 0);
assert!(metrics.peak_batch_size > 1);
}
#[tokio::test]
async fn test_model_version_management() {
let config = create_test_config();
let mut engine = InferenceEngine::new(config).await.unwrap();
// Load model version 1
let mock_weights_v1 = MockModelWeights::new().unwrap();
engine
.load_model_version(
"test_model",
"v1.0",
&mock_weights_v1.layers,
&mock_weights_v1.config,
)
.await
.unwrap();
// Load model version 2
let mock_weights_v2 = MockModelWeights::new().unwrap();
engine
.load_model_version(
"test_model",
"v2.0",
&mock_weights_v2.layers,
&mock_weights_v2.config,
)
.await
.unwrap();
// List versions
let versions = engine.list_model_versions("test_model").await.unwrap();
assert!(versions.contains(&"v1.0".to_string()));
assert!(versions.contains(&"v2.0".to_string()));
// Switch to specific version
engine
.switch_model_version("test_model", "v1.0")
.await
.unwrap();
let current_version = engine
.get_current_model_version("test_model")
.await
.unwrap();
assert_eq!(current_version, "v1.0");
// Test inference works with specific version
let request = InferenceRequest::new("test_model".to_string(), vec![1, 2, 3], 5);
let result = engine.infer(request).await.unwrap();
assert!(!result.output_tokens.is_empty());
}
#[tokio::test]
async fn test_health_checks() {
let config = create_test_config();
let mut engine = InferenceEngine::new(config).await.unwrap();
// Initial health should be ok
let health = engine.health_check().await;
assert!(health.is_healthy);
assert_eq!(health.status, "Ready");
// Load model
let mock_weights = MockModelWeights::new().unwrap();
engine
.load_model("test_model", &mock_weights.layers, &mock_weights.config)
.await
.unwrap();
// Health should still be good
let health = engine.health_check().await;
assert!(health.is_healthy);
assert!(health.models_loaded > 0);
// Detailed health check
let detailed_health = engine.detailed_health_check().await;
assert!(detailed_health.is_healthy);
assert!(detailed_health.memory_usage.total_allocated > 0);
assert!(detailed_health.memory_usage.utilization < 1.0);
assert!(detailed_health.model_health.contains_key("test_model"));
}
// Integration tests combining multiple features
#[tokio::test]
async fn test_full_pipeline_integration() {
let mut config = create_test_config();
config.max_batch_size = 4;
config.enable_caching = true;
config.enable_quantization = true;
config.optimization_level = 2;
let mut engine = InferenceEngine::new(config).await.unwrap();
// Load and optimize model
let mock_weights = MockModelWeights::new().unwrap();
engine
.load_model("test_model", &mock_weights.layers, &mock_weights.config)
.await
.unwrap();
// Apply quantization
let quant_config = QuantizationConfig::new().with_scheme(QuantizationScheme::INT8);
engine
.quantize_model("test_model", quant_config)
.await
.unwrap();
// Run inference workload
let mut results = vec![];
for i in 0..10 {
let request = InferenceRequest::new("test_model".to_string(), vec![i, i + 1, i + 2], 8);
let result = engine.infer(request).await.unwrap();
results.push(result);
}
// Validate all succeeded
assert_eq!(results.len(), 10);
for result in &results {
assert!(!result.output_tokens.is_empty());
assert!(result.metrics.is_some());
}
// Check performance metrics
let metrics = engine.get_performance_metrics().await;
assert!(metrics.total_requests >= 10);
assert!(metrics.success_rate >= 0.9);
// Validate memory efficiency
let memory_stats = engine.get_memory_stats().await;
assert!(memory_stats.fragmentation < 0.2); // Low fragmentation
assert!(memory_stats.utilization > 0.1); // Some memory used
}
// Note: TensorError -> InferenceError conversion is handled internally by the crate