219 lines
7.1 KiB
Rust
219 lines
7.1 KiB
Rust
use rtx_auto::{
|
|
error::AutoError,
|
|
proposal::{Proposal, ProposalType},
|
|
rollback::{Checkpoint, CheckpointType, RollbackManager},
|
|
};
|
|
use rtx_runtime::Runtime;
|
|
use std::collections::HashMap;
|
|
use std::sync::Arc;
|
|
|
|
#[tokio::test]
|
|
async fn test_rollback_manager_creation() {
|
|
let runtime = Arc::new(Runtime::new().expect("Failed to create runtime"));
|
|
let mut manager = RollbackManager::new(runtime.clone());
|
|
assert!(manager.is_ok(), "Failed to create RollbackManager");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_create_checkpoint() {
|
|
let runtime = Arc::new(Runtime::new().expect("Failed to create runtime"));
|
|
let mut manager = RollbackManager::new(runtime.clone()).unwrap();
|
|
|
|
let mut state = HashMap::new();
|
|
state.insert("model_weights".to_string(), vec![1.0, 2.0, 3.0]);
|
|
state.insert("optimizer_state".to_string(), vec![0.1, 0.2, 0.3]);
|
|
|
|
let checkpoint = manager
|
|
.create_checkpoint(
|
|
CheckpointType::BeforeOptimization,
|
|
state,
|
|
"Initial model state before applying quantization".to_string(),
|
|
)
|
|
.await;
|
|
|
|
assert!(checkpoint.is_ok());
|
|
let checkpoint = checkpoint.unwrap();
|
|
assert_eq!(
|
|
checkpoint.checkpoint_type(),
|
|
CheckpointType::BeforeOptimization
|
|
);
|
|
assert!(checkpoint.id().len() > 0);
|
|
assert!(checkpoint.timestamp() > 0);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_save_and_load_checkpoint() {
|
|
let runtime = Arc::new(Runtime::new().expect("Failed to create runtime"));
|
|
let mut manager = RollbackManager::new(runtime.clone()).unwrap();
|
|
|
|
let mut original_state = HashMap::new();
|
|
original_state.insert("layer1".to_string(), vec![1.0, 2.0, 3.0, 4.0]);
|
|
original_state.insert("layer2".to_string(), vec![5.0, 6.0, 7.0, 8.0]);
|
|
|
|
let checkpoint = manager
|
|
.create_checkpoint(
|
|
CheckpointType::BeforeOptimization,
|
|
original_state.clone(),
|
|
"Test checkpoint".to_string(),
|
|
)
|
|
.await
|
|
.unwrap();
|
|
|
|
let save_result = manager.save_checkpoint(&checkpoint).await;
|
|
assert!(save_result.is_ok());
|
|
|
|
let loaded_checkpoint = manager.load_checkpoint(checkpoint.id()).await;
|
|
assert!(loaded_checkpoint.is_ok());
|
|
|
|
let loaded = loaded_checkpoint.unwrap();
|
|
assert_eq!(loaded.id(), checkpoint.id());
|
|
assert_eq!(loaded.state(), checkpoint.state());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_rollback_to_checkpoint() {
|
|
let runtime = Arc::new(Runtime::new().expect("Failed to create runtime"));
|
|
let mut manager = RollbackManager::new(runtime.clone()).unwrap();
|
|
|
|
// Create original state
|
|
let mut original_state = HashMap::new();
|
|
original_state.insert("accuracy".to_string(), vec![0.95]);
|
|
original_state.insert("loss".to_string(), vec![0.1]);
|
|
|
|
let checkpoint = manager
|
|
.create_checkpoint(
|
|
CheckpointType::BeforeOptimization,
|
|
original_state,
|
|
"Before optimization".to_string(),
|
|
)
|
|
.await
|
|
.unwrap();
|
|
|
|
manager.save_checkpoint(&checkpoint).await.unwrap();
|
|
|
|
// Simulate applying an optimization that degrades performance
|
|
let mut degraded_state = HashMap::new();
|
|
degraded_state.insert("accuracy".to_string(), vec![0.70]); // Worse accuracy
|
|
degraded_state.insert("loss".to_string(), vec![0.5]); // Higher loss
|
|
|
|
manager.update_current_state(degraded_state).await.unwrap();
|
|
|
|
// Rollback to the checkpoint
|
|
let rollback_result = manager.rollback_to_checkpoint(checkpoint.id()).await;
|
|
assert!(rollback_result.is_ok());
|
|
|
|
let current_state = manager.get_current_state().await.unwrap();
|
|
assert_eq!(current_state["accuracy"], vec![0.95]);
|
|
assert_eq!(current_state["loss"], vec![0.1]);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_automatic_rollback_on_degradation() {
|
|
let runtime = Arc::new(Runtime::new().expect("Failed to create runtime"));
|
|
let mut manager = RollbackManager::new(runtime.clone()).unwrap();
|
|
|
|
// Set up automatic rollback threshold
|
|
manager
|
|
.set_rollback_threshold("accuracy", 0.90)
|
|
.await
|
|
.unwrap();
|
|
|
|
let mut good_state = HashMap::new();
|
|
good_state.insert("accuracy".to_string(), vec![0.95]);
|
|
|
|
let checkpoint = manager
|
|
.create_checkpoint(
|
|
CheckpointType::Automatic,
|
|
good_state,
|
|
"Good performance state".to_string(),
|
|
)
|
|
.await
|
|
.unwrap();
|
|
|
|
manager.save_checkpoint(&checkpoint).await.unwrap();
|
|
|
|
// Apply changes that degrade performance below threshold
|
|
let mut bad_state = HashMap::new();
|
|
bad_state.insert("accuracy".to_string(), vec![0.85]); // Below threshold
|
|
|
|
let should_rollback = manager
|
|
.should_trigger_automatic_rollback(&bad_state)
|
|
.await
|
|
.unwrap();
|
|
assert!(
|
|
should_rollback,
|
|
"Should trigger automatic rollback when performance degrades"
|
|
);
|
|
|
|
if should_rollback {
|
|
let rollback_result = manager.automatic_rollback().await;
|
|
assert!(rollback_result.is_ok());
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_checkpoint_cleanup() {
|
|
let runtime = Arc::new(Runtime::new().expect("Failed to create runtime"));
|
|
let mut manager = RollbackManager::new(runtime.clone()).unwrap();
|
|
|
|
// Create multiple checkpoints
|
|
let mut checkpoints = Vec::new();
|
|
for i in 0..5 {
|
|
let mut state = HashMap::new();
|
|
state.insert(format!("metric_{}", i), vec![i as f32]);
|
|
|
|
let checkpoint = manager
|
|
.create_checkpoint(CheckpointType::Manual, state, format!("Checkpoint {}", i))
|
|
.await
|
|
.unwrap();
|
|
|
|
manager.save_checkpoint(&checkpoint).await.unwrap();
|
|
checkpoints.push(checkpoint);
|
|
}
|
|
|
|
// Cleanup old checkpoints (keep only 3 most recent)
|
|
let cleanup_result = manager.cleanup_old_checkpoints(3).await;
|
|
assert!(cleanup_result.is_ok());
|
|
|
|
let remaining_checkpoints = manager.list_checkpoints().await.unwrap();
|
|
assert!(
|
|
remaining_checkpoints.len() <= 3,
|
|
"Should keep only 3 most recent checkpoints"
|
|
);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_proposal_application_with_rollback() {
|
|
let runtime = Arc::new(Runtime::new().expect("Failed to create runtime"));
|
|
let mut manager = RollbackManager::new(runtime.clone()).unwrap();
|
|
|
|
let proposal = Proposal::new(
|
|
ProposalType::Quantization,
|
|
"Apply aggressive quantization".to_string(),
|
|
2.0,
|
|
);
|
|
|
|
// Create checkpoint before applying proposal
|
|
let mut initial_state = HashMap::new();
|
|
initial_state.insert("model_size".to_string(), vec![1000.0]);
|
|
initial_state.insert("accuracy".to_string(), vec![0.95]);
|
|
|
|
let checkpoint = manager
|
|
.create_checkpoint_for_proposal(&proposal, initial_state)
|
|
.await;
|
|
|
|
assert!(checkpoint.is_ok());
|
|
let checkpoint = checkpoint.unwrap();
|
|
manager.save_checkpoint(&checkpoint).await.unwrap();
|
|
|
|
// Simulate proposal application result
|
|
let mut new_state = HashMap::new();
|
|
new_state.insert("model_size".to_string(), vec![250.0]); // Smaller model
|
|
new_state.insert("accuracy".to_string(), vec![0.85]); // Lower accuracy
|
|
|
|
let validation_result = manager
|
|
.validate_proposal_result(&proposal, &new_state)
|
|
.await;
|
|
assert!(validation_result.is_ok());
|
|
}
|