Initial commit

This commit is contained in:
redclawsystems
2026-03-04 00:08:42 +00:00
commit 4d88dc0584
4449 changed files with 1556714 additions and 0 deletions
@@ -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);
}
}