Files
rustytorch/crates/training/rtx-transformers/examples/adabound_demo.rs
T
2026-03-04 00:08:42 +00:00

228 lines
7.2 KiB
Rust

//! AdaBound Optimizer Demo
//!
//! This example demonstrates the usage of the AdaBound optimizer
//! with different configurations and scenarios.
use rtx_tensor::{Device, Tensor};
use rtx_transformers::optimizers::{AdaBoundConfig, AdaBoundOptimizer, Optimizer};
fn main() -> Result<(), Box<dyn std::error::Error>> {
println!("🚀 AdaBound Optimizer Demo");
println!("================================");
// Demo 1: Basic AdaBound configuration
demo_basic_adabound()?;
// Demo 2: AMSBound variant
demo_amsbound_variant()?;
// Demo 3: Different hyperparameters
demo_hyperparameter_variations()?;
// Demo 4: Bounds convergence demonstration
demo_bounds_convergence()?;
println!("\n✅ All AdaBound demos completed successfully!");
Ok(())
}
fn demo_basic_adabound() -> Result<(), Box<dyn std::error::Error>> {
println!("\n📊 Demo 1: Basic AdaBound Configuration");
println!("--------------------------------------");
let config = AdaBoundConfig {
learning_rate: 1e-3,
final_lr: 0.1,
beta1: 0.9,
beta2: 0.999,
gamma: 1e-3,
eps: 1e-8,
weight_decay: 1e-4,
amsbound: false,
};
let mut optimizer = AdaBoundOptimizer::new(config)?;
println!("✓ Created AdaBound optimizer");
println!(" Learning Rate: {}", optimizer.learning_rate());
println!(" Final LR: {}", optimizer.final_lr());
println!(" Beta1: {}", optimizer.beta1());
println!(" Beta2: {}", optimizer.beta2());
println!(" Gamma: {}", optimizer.gamma());
println!(" Weight Decay: {}", optimizer.weight_decay());
println!(" AMSBound: {}", optimizer.amsbound());
println!(" Optimizer Type: {}", optimizer.optimizer_type());
// Create sample parameters and gradients
let param = Tensor::randn(vec![10, 5], Device::cuda(0).unwrap_or(Device::default()))?;
let grad = Tensor::randn(vec![10, 5], Device::cuda(0).unwrap_or(Device::default()))? * 0.1;
// Perform optimization step
let updated_param = optimizer.step_param("demo_param", &param, &grad)?;
println!("✓ Performed optimization step");
println!(" Step count: {}", optimizer.get_step_count("demo_param")?);
println!(" Has state: {}", optimizer.has_state("demo_param"));
// Verify parameter was updated
let param_norm = param.norm()?.to_cpu()?[0];
let updated_norm = updated_param.norm()?.to_cpu()?[0];
println!(" Original param norm: {:.6}", param_norm);
println!(" Updated param norm: {:.6}", updated_norm);
Ok(())
}
fn demo_amsbound_variant() -> Result<(), Box<dyn std::error::Error>> {
println!("\n📈 Demo 2: AMSBound Variant");
println!("---------------------------");
let config = AdaBoundConfig {
learning_rate: 2e-3,
final_lr: 0.05,
beta1: 0.9,
beta2: 0.999,
gamma: 5e-3,
eps: 1e-8,
weight_decay: 0.0,
amsbound: true, // Enable AMSBound
};
let mut optimizer = AdaBoundOptimizer::new(config)?;
println!("✓ Created AMSBound optimizer (AdaBound with AMSGrad)");
println!(" AMSBound enabled: {}", optimizer.amsbound());
// Create parameters with varying gradient scales
let param = Tensor::ones(vec![5, 5], Device::cuda(0).unwrap_or(Device::default()))?;
// Simulate training with varying gradients
let gradients = vec![
Tensor::ones(vec![5, 5], Device::cuda(0).unwrap_or(Device::default()))? * 0.1, // Small gradient
Tensor::ones(vec![5, 5], Device::cuda(0).unwrap_or(Device::default()))? * 1.0, // Large gradient
Tensor::ones(vec![5, 5], Device::cuda(0).unwrap_or(Device::default()))? * 0.05, // Small gradient again
];
let mut current_param = param.clone();
for (step, grad) in gradients.iter().enumerate() {
current_param = optimizer.step_param("ams_param", &current_param, grad)?;
let step_count = optimizer.get_step_count("ams_param")?;
let param_norm = current_param.norm()?.to_cpu()?[0];
println!(" Step {}: param_norm={:.6}", step_count, param_norm);
}
println!("✓ AMSBound maintains max of past squared gradients for stability");
Ok(())
}
fn demo_hyperparameter_variations() -> Result<(), Box<dyn std::error::Error>> {
println!("\n⚙️ Demo 3: Hyperparameter Variations");
println!("-------------------------------------");
let scenarios = vec![
(
"Conservative",
AdaBoundConfig {
learning_rate: 1e-4,
final_lr: 0.01,
gamma: 1e-4,
..AdaBoundConfig::default()
},
),
(
"Aggressive",
AdaBoundConfig {
learning_rate: 1e-2,
final_lr: 0.5,
gamma: 1e-2,
..AdaBoundConfig::default()
},
),
(
"High Momentum",
AdaBoundConfig {
beta1: 0.95,
beta2: 0.9999,
..AdaBoundConfig::default()
},
),
(
"With Weight Decay",
AdaBoundConfig {
weight_decay: 1e-3,
..AdaBoundConfig::default()
},
),
];
for (name, config) in scenarios {
println!("\n {} Configuration:", name);
let optimizer = AdaBoundOptimizer::new(config)?;
println!(
" LR: {:.6}, Final LR: {:.3}, Gamma: {:.6}",
optimizer.learning_rate(),
optimizer.final_lr(),
optimizer.gamma()
);
println!(
" Beta1: {:.3}, Beta2: {:.4}, Weight Decay: {:.6}",
optimizer.beta1(),
optimizer.beta2(),
optimizer.weight_decay()
);
}
println!("\n✓ Different configurations allow adapting to various training scenarios");
Ok(())
}
fn demo_bounds_convergence() -> Result<(), Box<dyn std::error::Error>> {
println!("\n📉 Demo 4: Dynamic Bounds Convergence");
println!("-------------------------------------");
let config = AdaBoundConfig {
learning_rate: 1e-2,
final_lr: 0.1,
gamma: 1e-2, // Higher gamma for visible convergence
..AdaBoundConfig::default()
};
let optimizer = AdaBoundOptimizer::new(config)?;
println!("Demonstrating how bounds tighten over training steps:");
println!("Step Lower Bound Upper Bound Bound Width");
println!("---- ----------- ----------- -----------");
for step in [1, 5, 10, 20, 50, 100, 200, 500, 1000] {
let (lower, upper) = optimizer.calculate_bounds(step);
let width = upper - lower;
println!(
"{:4} {:11.6} {:11.6} {:11.6}",
step, lower, upper, width
);
}
println!("\n✓ Bounds converge from adaptive (wide) to SGD-like (narrow) over time");
println!(" This provides the benefits of both adaptive methods and SGD");
Ok(())
}
#[cfg(test)]
mod demo_tests {
use super::*;
#[test]
fn test_demo_functions() {
// Test that our demo functions run without panicking
assert!(demo_basic_adabound().is_ok());
assert!(demo_amsbound_variant().is_ok());
assert!(demo_hyperparameter_variations().is_ok());
assert!(demo_bounds_convergence().is_ok());
}
}