Initial commit
This commit is contained in:
@@ -0,0 +1,171 @@
|
||||
//! Tests for elastic training enhancements.
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::super::*;
|
||||
|
||||
#[test]
|
||||
fn test_event_manager() {
|
||||
let manager = ElasticEventManager::new(100);
|
||||
manager.set_step(42);
|
||||
manager.emit(ElasticEventType::WorkerJoined, Some(1), 2);
|
||||
|
||||
let history = manager.history();
|
||||
assert_eq!(history.len(), 1);
|
||||
assert_eq!(history[0].event_type, ElasticEventType::WorkerJoined);
|
||||
assert_eq!(history[0].training_step, 42);
|
||||
assert_eq!(history[0].new_world_size, 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_batch_size_constant_global() {
|
||||
let event_mgr = shared_event_manager(100);
|
||||
let config = BatchSizeConfig {
|
||||
strategy: BatchSizeStrategy::ConstantGlobal,
|
||||
base_batch_size: 32,
|
||||
global_batch_target: Some(128), // 128 total
|
||||
min_batch_size: 4,
|
||||
max_batch_size: 64,
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let adapter = BatchSizeAdapter::new(config, event_mgr);
|
||||
|
||||
// With 4 workers: 128 / 4 = 32 per worker
|
||||
assert_eq!(adapter.adapt(4, 0), 32);
|
||||
|
||||
// With 2 workers: 128 / 2 = 64 per worker
|
||||
assert_eq!(adapter.adapt(2, 1), 64);
|
||||
|
||||
// With 8 workers: 128 / 8 = 16 per worker
|
||||
assert_eq!(adapter.adapt(8, 2), 16);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_lr_linear_scaling() {
|
||||
let event_mgr = shared_event_manager(100);
|
||||
let config = LRAdaptationConfig {
|
||||
strategy: LRAdaptationStrategy::LinearScaling,
|
||||
base_lr: 0.001,
|
||||
base_batch_size: 32,
|
||||
min_lr: 1e-6,
|
||||
max_lr: 0.1,
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let adapter = LRAdapter::new(config, event_mgr);
|
||||
|
||||
// 2x batch size = 2x LR
|
||||
adapter.adapt(64, 1);
|
||||
assert!((adapter.current_lr() - 0.002).abs() < 1e-6);
|
||||
|
||||
// 4x batch size = 4x LR
|
||||
adapter.adapt(32, 4);
|
||||
assert!((adapter.current_lr() - 0.004).abs() < 1e-6);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_loss_scaling_recovery() {
|
||||
let event_mgr = shared_event_manager(100);
|
||||
let config = LossScalingRecoveryConfig {
|
||||
initial_scale: 1024.0,
|
||||
scale_down_factor: 0.5,
|
||||
growth_interval: 10,
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let recovery = LossScalingRecovery::new(config, event_mgr);
|
||||
|
||||
// Initial scale
|
||||
assert_eq!(recovery.current_scale(), 1024.0);
|
||||
|
||||
// Report overflow
|
||||
recovery.report_overflow();
|
||||
assert_eq!(recovery.current_scale(), 512.0);
|
||||
|
||||
// Checkpoint
|
||||
let checkpoint = recovery.checkpoint();
|
||||
assert_eq!(checkpoint.checkpoint_scale, 512.0);
|
||||
|
||||
// More overflows
|
||||
recovery.report_overflow();
|
||||
assert_eq!(recovery.current_scale(), 256.0);
|
||||
|
||||
// Restore from checkpoint
|
||||
recovery.restore(&checkpoint);
|
||||
assert_eq!(recovery.current_scale(), 512.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_elastic_agent() {
|
||||
let event_mgr = shared_event_manager(100);
|
||||
let config = ElasticAgentConfig {
|
||||
min_workers: 1,
|
||||
max_workers: 8,
|
||||
max_restarts: 3,
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let agent = ElasticAgent::new(config, event_mgr);
|
||||
|
||||
// Register workers
|
||||
agent.register_worker(1, 0, 0).unwrap();
|
||||
agent.register_worker(2, 1, 1).unwrap();
|
||||
assert_eq!(agent.world_size(), 2);
|
||||
|
||||
// Heartbeat
|
||||
agent.heartbeat(1).unwrap();
|
||||
|
||||
// Activate
|
||||
agent.activate_worker(1).unwrap();
|
||||
|
||||
// Remove worker
|
||||
agent.remove_worker(2).unwrap();
|
||||
assert_eq!(agent.world_size(), 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_agent_worker_failure_restart() {
|
||||
let event_mgr = shared_event_manager(100);
|
||||
let config = ElasticAgentConfig {
|
||||
max_restarts: 2,
|
||||
auto_restart: true,
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let agent = ElasticAgent::new(config, event_mgr);
|
||||
agent.register_worker(1, 0, 0).unwrap();
|
||||
|
||||
// First failure - should restart
|
||||
let should_restart = agent.handle_worker_failure(1).unwrap();
|
||||
assert!(should_restart);
|
||||
assert_eq!(agent.total_restarts(), 1);
|
||||
|
||||
// Second failure - should restart
|
||||
let should_restart = agent.handle_worker_failure(1).unwrap();
|
||||
assert!(should_restart);
|
||||
assert_eq!(agent.total_restarts(), 2);
|
||||
|
||||
// Third failure - exceeded max restarts, worker removed
|
||||
let should_restart = agent.handle_worker_failure(1).unwrap();
|
||||
assert!(!should_restart);
|
||||
assert_eq!(agent.world_size(), 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_event_history_limit() {
|
||||
let manager = ElasticEventManager::new(3);
|
||||
|
||||
// Emit 5 events
|
||||
for i in 0..5 {
|
||||
manager.set_step(i as u64);
|
||||
manager.emit(ElasticEventType::WorkerJoined, Some(i as u64), i + 1);
|
||||
}
|
||||
|
||||
// Only last 3 should be kept
|
||||
let history = manager.history();
|
||||
assert_eq!(history.len(), 3);
|
||||
assert_eq!(history[0].training_step, 2);
|
||||
assert_eq!(history[2].training_step, 4);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user