Initial commit
This commit is contained in:
@@ -0,0 +1,466 @@
|
||||
//! Comprehensive tests for Ranger optimizer implementation
|
||||
//! Following strict TDD - tests written first, then implementation
|
||||
//!
|
||||
//! Ranger combines RAdam (Rectified Adam) with Lookahead for improved convergence
|
||||
|
||||
use super::*;
|
||||
use crate::optimizers::{Optimizer, RangerConfig, RangerOptimizer};
|
||||
use rtx_tensor::{Tensor, Device, Shape, DType};
|
||||
use std::collections::HashMap;
|
||||
use approx::assert_relative_eq;
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_ranger_config_default_values() {
|
||||
let config = RangerConfig::default();
|
||||
|
||||
assert_eq!(config.lr, 0.001);
|
||||
assert_eq!(config.beta1, 0.95);
|
||||
assert_eq!(config.beta2, 0.999);
|
||||
assert_eq!(config.eps, 1e-8);
|
||||
assert_eq!(config.weight_decay, 0.0);
|
||||
assert_eq!(config.k, 5);
|
||||
assert_eq!(config.alpha, 0.5);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_ranger_config_validation() {
|
||||
// Valid config should succeed
|
||||
let config = RangerConfig {
|
||||
lr: 1e-3,
|
||||
beta1: 0.95,
|
||||
beta2: 0.999,
|
||||
eps: 1e-8,
|
||||
weight_decay: 1e-4,
|
||||
k: 5,
|
||||
alpha: 0.5,
|
||||
};
|
||||
|
||||
let optimizer = RangerOptimizer::new(config);
|
||||
assert!(optimizer.is_ok());
|
||||
|
||||
// Invalid learning rate should fail
|
||||
let invalid_config = RangerConfig {
|
||||
lr: -0.1,
|
||||
..RangerConfig::default()
|
||||
};
|
||||
assert!(RangerOptimizer::new(invalid_config).is_err());
|
||||
|
||||
// Invalid beta1 should fail
|
||||
let invalid_config = RangerConfig {
|
||||
beta1: 1.5,
|
||||
..RangerConfig::default()
|
||||
};
|
||||
assert!(RangerOptimizer::new(invalid_config).is_err());
|
||||
|
||||
// Invalid beta2 should fail
|
||||
let invalid_config = RangerConfig {
|
||||
beta2: -0.1,
|
||||
..RangerConfig::default()
|
||||
};
|
||||
assert!(RangerOptimizer::new(invalid_config).is_err());
|
||||
|
||||
// Invalid eps should fail
|
||||
let invalid_config = RangerConfig {
|
||||
eps: -1e-8,
|
||||
..RangerConfig::default()
|
||||
};
|
||||
assert!(RangerOptimizer::new(invalid_config).is_err());
|
||||
|
||||
// Invalid weight decay should fail
|
||||
let invalid_config = RangerConfig {
|
||||
weight_decay: -0.1,
|
||||
..RangerConfig::default()
|
||||
};
|
||||
assert!(RangerOptimizer::new(invalid_config).is_err());
|
||||
|
||||
// Invalid k should fail
|
||||
let invalid_config = RangerConfig {
|
||||
k: 0,
|
||||
..RangerConfig::default()
|
||||
};
|
||||
assert!(RangerOptimizer::new(invalid_config).is_err());
|
||||
|
||||
// Invalid alpha should fail
|
||||
let invalid_config = RangerConfig {
|
||||
alpha: 1.5,
|
||||
..RangerConfig::default()
|
||||
};
|
||||
assert!(RangerOptimizer::new(invalid_config).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_ranger_state_initialization() {
|
||||
let config = RangerConfig::default();
|
||||
let mut optimizer = RangerOptimizer::new(config).unwrap();
|
||||
|
||||
let param = Tensor::ones(vec![2, 3], Device::Cpu).unwrap();
|
||||
let grad = Tensor::ones(vec![2, 3], Device::Cpu).unwrap();
|
||||
|
||||
// Initially should have no state
|
||||
assert!(!optimizer.has_state("param1"));
|
||||
|
||||
// First step should initialize state
|
||||
let _updated_param = optimizer.step_param("param1", ¶m, &grad).unwrap();
|
||||
|
||||
// Now should have state
|
||||
assert!(optimizer.has_state("param1"));
|
||||
assert_eq!(optimizer.get_step_count("param1").unwrap(), 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_ranger_state_reset() {
|
||||
let config = RangerConfig::default();
|
||||
let mut optimizer = RangerOptimizer::new(config).unwrap();
|
||||
|
||||
let param = Tensor::ones(vec![2, 3], Device::Cpu).unwrap();
|
||||
let grad = Tensor::ones(vec![2, 3], Device::Cpu).unwrap();
|
||||
|
||||
// Initialize state
|
||||
let _updated_param = optimizer.step_param("param1", ¶m, &grad).unwrap();
|
||||
assert!(optimizer.has_state("param1"));
|
||||
|
||||
// Reset specific parameter state
|
||||
optimizer.reset_state("param1").unwrap();
|
||||
assert!(!optimizer.has_state("param1"));
|
||||
|
||||
// Initialize multiple parameters
|
||||
let _updated_param1 = optimizer.step_param("param1", ¶m, &grad).unwrap();
|
||||
let _updated_param2 = optimizer.step_param("param2", ¶m, &grad).unwrap();
|
||||
assert!(optimizer.has_state("param1"));
|
||||
assert!(optimizer.has_state("param2"));
|
||||
|
||||
// Reset all state
|
||||
optimizer.reset_all_state();
|
||||
assert!(!optimizer.has_state("param1"));
|
||||
assert!(!optimizer.has_state("param2"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_ranger_learning_rate_management() {
|
||||
let config = RangerConfig::default();
|
||||
let mut optimizer = RangerOptimizer::new(config).unwrap();
|
||||
|
||||
// Check initial learning rate
|
||||
assert_eq!(optimizer.learning_rate(), 0.001);
|
||||
|
||||
// Update learning rate
|
||||
optimizer.set_learning_rate(2e-3).unwrap();
|
||||
assert_eq!(optimizer.learning_rate(), 2e-3);
|
||||
|
||||
// Invalid learning rate should fail
|
||||
assert!(optimizer.set_learning_rate(-1e-3).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_ranger_momentum_and_variance_update() {
|
||||
let config = RangerConfig {
|
||||
lr: 1e-2,
|
||||
beta1: 0.9,
|
||||
beta2: 0.999,
|
||||
..RangerConfig::default()
|
||||
};
|
||||
let mut optimizer = RangerOptimizer::new(config).unwrap();
|
||||
|
||||
let param = Tensor::ones(vec![2, 2], Device::Cpu).unwrap();
|
||||
let grad = Tensor::ones(vec![2, 2], Device::Cpu).unwrap() * 0.1;
|
||||
|
||||
// First step
|
||||
let updated_param1 = optimizer.step_param("param1", ¶m, &grad).unwrap();
|
||||
|
||||
// Second step with same gradient
|
||||
let updated_param2 = optimizer.step_param("param1", &updated_param1, &grad).unwrap();
|
||||
|
||||
// Third step with same gradient
|
||||
let updated_param3 = optimizer.step_param("param1", &updated_param2, &grad).unwrap();
|
||||
|
||||
// Parameters should be decreasing (gradient descent)
|
||||
let param_data = param.to_cpu().unwrap();
|
||||
let updated_data1 = updated_param1.to_cpu().unwrap();
|
||||
let updated_data2 = updated_param2.to_cpu().unwrap();
|
||||
let updated_data3 = updated_param3.to_cpu().unwrap();
|
||||
|
||||
// All parameters should decrease due to positive gradient
|
||||
for i in 0..param_data.len() {
|
||||
assert!(updated_data1[i] < param_data[i]);
|
||||
assert!(updated_data2[i] < updated_data1[i]);
|
||||
assert!(updated_data3[i] < updated_data2[i]);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_ranger_rectification_mechanism() {
|
||||
let config = RangerConfig {
|
||||
lr: 1e-2,
|
||||
beta1: 0.9,
|
||||
beta2: 0.999,
|
||||
..RangerConfig::default()
|
||||
};
|
||||
let mut optimizer = RangerOptimizer::new(config).unwrap();
|
||||
|
||||
let param = Tensor::ones(vec![2, 2], Device::Cpu).unwrap();
|
||||
let grad = Tensor::ones(vec![2, 2], Device::Cpu).unwrap() * 0.1;
|
||||
|
||||
// Early steps should use rectified learning rate
|
||||
let initial_param = param.clone();
|
||||
let mut current_param = param.clone();
|
||||
|
||||
// Perform several steps to test rectification behavior
|
||||
for step in 1..=10 {
|
||||
current_param = optimizer.step_param("param1", ¤t_param, &grad).unwrap();
|
||||
assert_eq!(optimizer.get_step_count("param1").unwrap(), step as i64);
|
||||
}
|
||||
|
||||
// RAdam should provide stable updates even in early training
|
||||
let param_data = initial_param.to_cpu().unwrap();
|
||||
let updated_data = current_param.to_cpu().unwrap();
|
||||
|
||||
for i in 0..param_data.len() {
|
||||
assert!(updated_data[i] < param_data[i]); // Should decrease with positive gradient
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_ranger_lookahead_mechanism() {
|
||||
let config = RangerConfig {
|
||||
lr: 1e-2,
|
||||
k: 3, // Short lookahead for testing
|
||||
alpha: 0.8, // Strong lookahead update
|
||||
..RangerConfig::default()
|
||||
};
|
||||
let mut optimizer = RangerOptimizer::new(config).unwrap();
|
||||
|
||||
let param = Tensor::ones(vec![2, 2], Device::Cpu).unwrap();
|
||||
let grad = Tensor::ones(vec![2, 2], Device::Cpu).unwrap() * 0.1;
|
||||
|
||||
let initial_param = param.clone();
|
||||
let mut current_param = param.clone();
|
||||
|
||||
// Perform k+1 steps to trigger lookahead update
|
||||
for _ in 1..=4 {
|
||||
current_param = optimizer.step_param("param1", ¤t_param, &grad).unwrap();
|
||||
}
|
||||
|
||||
// After k steps, lookahead should have updated slow weights
|
||||
let param_data = initial_param.to_cpu().unwrap();
|
||||
let updated_data = current_param.to_cpu().unwrap();
|
||||
|
||||
for i in 0..param_data.len() {
|
||||
assert!(updated_data[i] < param_data[i]); // Should decrease
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_ranger_weight_decay() {
|
||||
let config = RangerConfig {
|
||||
lr: 1e-2,
|
||||
weight_decay: 1e-4,
|
||||
..RangerConfig::default()
|
||||
};
|
||||
let mut optimizer = RangerOptimizer::new(config).unwrap();
|
||||
|
||||
let param = Tensor::ones(vec![2, 2], Device::Cpu).unwrap() * 2.0;
|
||||
let grad = Tensor::zeros(vec![2, 2], Device::Cpu).unwrap(); // Zero gradient
|
||||
|
||||
// With zero gradient, only weight decay should affect parameters
|
||||
let updated_param = optimizer.step_param("param1", ¶m, &grad).unwrap();
|
||||
|
||||
let param_data = param.to_cpu().unwrap();
|
||||
let updated_data = updated_param.to_cpu().unwrap();
|
||||
|
||||
// Parameters should decrease due to weight decay
|
||||
for i in 0..param_data.len() {
|
||||
assert!(updated_data[i] < param_data[i]);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_ranger_gradient_shapes_mismatch() {
|
||||
let config = RangerConfig::default();
|
||||
let mut optimizer = RangerOptimizer::new(config).unwrap();
|
||||
|
||||
let param = Tensor::ones(vec![2, 3], Device::Cpu).unwrap();
|
||||
let grad = Tensor::ones(vec![3, 2], Device::Cpu).unwrap();
|
||||
|
||||
// Mismatched shapes should fail
|
||||
let result = optimizer.step_param("param1", ¶m, &grad);
|
||||
assert!(result.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_ranger_stored_gradients() {
|
||||
let config = RangerConfig::default();
|
||||
let mut optimizer = RangerOptimizer::new(config).unwrap();
|
||||
|
||||
let param1 = Tensor::ones(vec![2, 2], Device::Cpu).unwrap();
|
||||
let param2 = Tensor::ones(vec![3, 3], Device::Cpu).unwrap();
|
||||
let grad1 = Tensor::ones(vec![2, 2], Device::Cpu).unwrap() * 0.1;
|
||||
let grad2 = Tensor::ones(vec![3, 3], Device::Cpu).unwrap() * 0.2;
|
||||
|
||||
let mut gradients = HashMap::new();
|
||||
gradients.insert("param1".to_string(), grad1.clone());
|
||||
gradients.insert("param2".to_string(), grad2.clone());
|
||||
|
||||
// Store gradients
|
||||
optimizer.set_gradients(gradients).unwrap();
|
||||
|
||||
// Process stored gradients
|
||||
let updates = optimizer.step(1e-3).unwrap();
|
||||
|
||||
assert!(updates.contains_key("param1"));
|
||||
assert!(updates.contains_key("param2"));
|
||||
assert_eq!(updates.len(), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_ranger_optimizer_type() {
|
||||
let config = RangerConfig::default();
|
||||
let optimizer = RangerOptimizer::new(config).unwrap();
|
||||
|
||||
assert_eq!(optimizer.optimizer_type(), "Ranger");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_ranger_zero_grad() {
|
||||
let config = RangerConfig::default();
|
||||
let mut optimizer = RangerOptimizer::new(config).unwrap();
|
||||
|
||||
let grad1 = Tensor::ones(vec![2, 2], Device::Cpu).unwrap() * 0.1;
|
||||
let grad2 = Tensor::ones(vec![3, 3], Device::Cpu).unwrap() * 0.2;
|
||||
|
||||
let mut gradients = HashMap::new();
|
||||
gradients.insert("param1".to_string(), grad1);
|
||||
gradients.insert("param2".to_string(), grad2);
|
||||
|
||||
// Store gradients
|
||||
optimizer.set_gradients(gradients).unwrap();
|
||||
|
||||
// Zero gradients
|
||||
optimizer.zero_grad();
|
||||
|
||||
// Should have no stored gradients after zero_grad
|
||||
let updates = optimizer.step(1e-3).unwrap();
|
||||
assert_eq!(updates.len(), 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_ranger_lookahead_k_steps() {
|
||||
let config = RangerConfig {
|
||||
lr: 1e-2,
|
||||
k: 2, // Very short for testing
|
||||
alpha: 0.5,
|
||||
..RangerConfig::default()
|
||||
};
|
||||
let mut optimizer = RangerOptimizer::new(config).unwrap();
|
||||
|
||||
let param = Tensor::ones(vec![1, 1], Device::Cpu).unwrap();
|
||||
let grad = Tensor::ones(vec![1, 1], Device::Cpu).unwrap() * 0.1;
|
||||
|
||||
let mut param_current = param.clone();
|
||||
|
||||
// Step 1: Fast weight update only
|
||||
let param1 = optimizer.step_param("param1", ¶m_current, &grad).unwrap();
|
||||
|
||||
// Step 2: Should trigger lookahead update (k=2)
|
||||
let param2 = optimizer.step_param("param1", ¶m1, &grad).unwrap();
|
||||
|
||||
// Step 3: Should trigger lookahead again
|
||||
let param3 = optimizer.step_param("param1", ¶m2, &grad).unwrap();
|
||||
|
||||
// All steps should produce valid updates
|
||||
let param0_data = param.to_cpu().unwrap();
|
||||
let param1_data = param1.to_cpu().unwrap();
|
||||
let param2_data = param2.to_cpu().unwrap();
|
||||
let param3_data = param3.to_cpu().unwrap();
|
||||
|
||||
// Progressive decrease due to gradient descent
|
||||
assert!(param1_data[0] < param0_data[0]);
|
||||
assert!(param2_data[0] < param1_data[0]);
|
||||
assert!(param3_data[0] < param2_data[0]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_ranger_different_alpha_values() {
|
||||
let param = Tensor::ones(vec![2, 2], Device::Cpu).unwrap();
|
||||
let grad = Tensor::ones(vec![2, 2], Device::Cpu).unwrap() * 0.2;
|
||||
|
||||
// Test with small alpha (weak lookahead)
|
||||
let config_weak = RangerConfig {
|
||||
lr: 1e-2,
|
||||
k: 2,
|
||||
alpha: 0.1, // Weak lookahead
|
||||
..RangerConfig::default()
|
||||
};
|
||||
let mut optimizer_weak = RangerOptimizer::new(config_weak).unwrap();
|
||||
|
||||
// Test with large alpha (strong lookahead)
|
||||
let config_strong = RangerConfig {
|
||||
lr: 1e-2,
|
||||
k: 2,
|
||||
alpha: 0.9, // Strong lookahead
|
||||
..RangerConfig::default()
|
||||
};
|
||||
let mut optimizer_strong = RangerOptimizer::new(config_strong).unwrap();
|
||||
|
||||
// Perform enough steps to trigger lookahead
|
||||
let mut param_weak = param.clone();
|
||||
let mut param_strong = param.clone();
|
||||
|
||||
for _ in 0..3 {
|
||||
param_weak = optimizer_weak.step_param("param1", ¶m_weak, &grad).unwrap();
|
||||
param_strong = optimizer_strong.step_param("param1", ¶m_strong, &grad).unwrap();
|
||||
}
|
||||
|
||||
let weak_data = param_weak.to_cpu().unwrap();
|
||||
let strong_data = param_strong.to_cpu().unwrap();
|
||||
|
||||
// Different alpha values should produce different results
|
||||
let mut different = false;
|
||||
for i in 0..weak_data.len() {
|
||||
if (weak_data[i] - strong_data[i]).abs() > 1e-6 {
|
||||
different = true;
|
||||
break;
|
||||
}
|
||||
}
|
||||
assert!(different);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_ranger_bias_correction() {
|
||||
let config = RangerConfig {
|
||||
lr: 1e-2,
|
||||
beta1: 0.9,
|
||||
beta2: 0.999,
|
||||
..RangerConfig::default()
|
||||
};
|
||||
let mut optimizer = RangerOptimizer::new(config).unwrap();
|
||||
|
||||
let param = Tensor::ones(vec![2, 2], Device::Cpu).unwrap();
|
||||
let grad = Tensor::ones(vec![2, 2], Device::Cpu).unwrap() * 0.1;
|
||||
|
||||
// Early steps should have different behavior due to bias correction
|
||||
let updated_param1 = optimizer.step_param("param1", ¶m, &grad).unwrap();
|
||||
let diff1 = (¶m - &updated_param1).unwrap();
|
||||
let norm1_squared: f64 = diff1.to_cpu().unwrap().iter().map(|&x| (x as f64).powi(2)).sum();
|
||||
|
||||
// Reset and test later step
|
||||
optimizer.reset_state("param1").unwrap();
|
||||
|
||||
// Advance step count
|
||||
let mut current_param = param.clone();
|
||||
for _ in 0..50 {
|
||||
current_param = optimizer.step_param("param1", ¤t_param, &grad).unwrap();
|
||||
}
|
||||
|
||||
// Another step after many iterations
|
||||
let updated_param_late = optimizer.step_param("param1", ¤t_param, &grad).unwrap();
|
||||
let diff_late = (¤t_param - &updated_param_late).unwrap();
|
||||
let norm_late_squared: f64 = diff_late.to_cpu().unwrap().iter().map(|&x| (x as f64).powi(2)).sum();
|
||||
|
||||
// Update magnitudes should be different due to accumulated state
|
||||
assert_ne!(norm1_squared, norm_late_squared);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user