use optirs_core::memory_efficient::{InPlaceAdam, InPlaceOptimizer, InPlaceSGD};
use optirs_core::optimizers::{Adam, SGD};
use optirs_core::Optimizer;
use scirs2_core::ndarray::Array2;
use std::error::Error;
use std::time::Instant;
#[allow(dead_code)]
fn estimate_memory_usage<T>(arrays: &[&T]) -> usize {
std::mem::size_of_val(arrays)
}
#[allow(dead_code)]
fn compute_loss_and_gradient(params: &Array2<f64>, target: &Array2<f64>) -> (f64, Array2<f64>) {
let diff = params - target;
let loss = diff.mapv(|x| x * x).sum() / (params.nrows() * params.ncols()) as f64;
let grad = diff * 2.0 / (params.nrows() * params.ncols()) as f64;
(loss, grad)
}
#[allow(dead_code)]
fn main() -> Result<(), Box<dyn Error>> {
let size = 1000;
let mut params = Array2::from_elem((size, size), 0.5);
let mut params_inplace = params.clone();
let target = Array2::from_elem((size, size), 1.0);
let mut adam = Adam::new(0.01);
let mut inplace_adam = InPlaceAdam::new(0.01);
let mut sgd = SGD::new(0.1);
let mut inplace_sgd = InPlaceSGD::new(0.1);
let mut regular_adam_memory = vec![];
let mut inplace_adam_memory = vec![];
let mut regular_sgd_memory = vec![];
let mut inplace_sgd_memory = vec![];
println!("Memory-Efficient Optimizer Comparison");
println!("====================================");
println!("Parameter size: {}x{}", size, size);
println!();
println!("Adam Optimizer Comparison:");
println!("-------------------------");
let start = Instant::now();
for step in 0..10 {
let (loss, grad) = compute_loss_and_gradient(¶ms, &target);
params = adam.step(¶ms, &grad)?;
regular_adam_memory.push(estimate_memory_usage(&[¶ms, &grad]));
if step % 3 == 0 {
println!(
" Step {}: Loss = {:.6}, Approx Memory = {} bytes",
step,
loss,
regular_adam_memory.last().expect("unwrap failed")
);
}
}
let regular_adam_time = start.elapsed();
let start = Instant::now();
for step in 0..10 {
let (loss, grad) = compute_loss_and_gradient(¶ms_inplace, &target);
inplace_adam.step_inplace(&mut params_inplace, &grad)?;
inplace_adam_memory.push(estimate_memory_usage(&[¶ms_inplace, &grad]));
if step % 3 == 0 {
println!(
" Step {} (in-place): Loss = {:.6}, Approx Memory = {} bytes",
step,
loss,
inplace_adam_memory.last().expect("unwrap failed")
);
}
}
let inplace_adam_time = start.elapsed();
println!("\nAdam Performance:");
println!(" Regular: {:?}", regular_adam_time);
println!(" In-place: {:?}", inplace_adam_time);
println!(
" Memory efficiency: ~{}% reduction",
((regular_adam_memory[0] - inplace_adam_memory[0]) * 100) / regular_adam_memory[0]
);
params = Array2::from_elem((size, size), 0.5);
params_inplace = params.clone();
println!("\nSGD Optimizer Comparison:");
println!("------------------------");
let start = Instant::now();
for step in 0..10 {
let (loss, grad) = compute_loss_and_gradient(¶ms, &target);
params = sgd.step(¶ms, &grad)?;
regular_sgd_memory.push(estimate_memory_usage(&[¶ms, &grad]));
if step % 3 == 0 {
println!(
" Step {}: Loss = {:.6}, Approx Memory = {} bytes",
step,
loss,
regular_sgd_memory.last().expect("unwrap failed")
);
}
}
let regular_sgd_time = start.elapsed();
let start = Instant::now();
for step in 0..10 {
let (loss, grad) = compute_loss_and_gradient(¶ms_inplace, &target);
inplace_sgd.step_inplace(&mut params_inplace, &grad)?;
inplace_sgd_memory.push(estimate_memory_usage(&[¶ms_inplace, &grad]));
if step % 3 == 0 {
println!(
" Step {} (in-place): Loss = {:.6}, Approx Memory = {} bytes",
step,
loss,
inplace_sgd_memory.last().expect("unwrap failed")
);
}
}
let inplace_sgd_time = start.elapsed();
println!("\nSGD Performance:");
println!(" Regular: {:?}", regular_sgd_time);
println!(" In-place: {:?}", inplace_sgd_time);
println!(
" Memory efficiency: ~{}% reduction",
((regular_sgd_memory[0] - inplace_sgd_memory[0]) * 100) / regular_sgd_memory[0]
);
println!("\nUtility Functions Demo:");
println!("----------------------");
use optirs_core::memory_efficient::{clip_inplace, normalize_inplace, scale_inplace};
let mut test_array = Array2::from_elem((3, 3), 2.0);
println!("Original array:\n{}", test_array);
scale_inplace(&mut test_array, 0.5);
println!("After scaling by 0.5:\n{}", test_array);
clip_inplace(&mut test_array, -0.5, 0.5);
println!("After clipping to [-0.5, 0.5]:\n{}", test_array);
normalize_inplace(&mut test_array);
println!("After normalizing:\n{}", test_array);
Ok(())
}