Files
rustytorch/crates/training/rtx-auto/tests/rollback_tests.rs
T
2026-03-04 00:08:42 +00:00

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());
}