use std::fmt;
use scirs2_core::ndarray::{Array, Dimension, ScalarOperand};
use scirs2_core::numeric::Float;
use super::{HardwareOptimizationConfig, HardwarePlatform, MemoryStrategy, QuantizationSupport};
use crate::error::{OptimError, Result};
use crate::optimizers::{Adam, Lion, Optimizer, LAMB, SGD};
use crate::schedulers::{ConstantScheduler, LearningRateScheduler};
use crate::utils::scalar_or;
pub const DEFAULT_BASE_LEARNING_RATE: f64 = 1e-3;
const LARGE_BATCH_THRESHOLD: usize = 512;
const LOW_POWER_BUDGET_WATTS: f64 = 5.0;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum HardwareOptimizerKind {
Sgd,
Lion,
Adam,
Lamb,
}
impl HardwareOptimizerKind {
pub fn state_buffers_per_parameter(self) -> usize {
match self {
Self::Sgd => 1,
Self::Lion => 1,
Self::Adam => 2,
Self::Lamb => 2,
}
}
pub fn name(self) -> &'static str {
match self {
Self::Sgd => "sgd",
Self::Lion => "lion",
Self::Adam => "adam",
Self::Lamb => "lamb",
}
}
pub fn recommend_for<A: Float>(
platform: &HardwarePlatform,
config: &HardwareOptimizationConfig<A>,
) -> Self {
let offloading = matches!(config.memory_strategy, MemoryStrategy::CPUOffloading { .. });
match platform {
HardwarePlatform::Edge {
power_budget,
quantization_support,
..
} => {
if offloading || *power_budget < LOW_POWER_BUDGET_WATTS {
Self::Sgd
} else if matches!(
quantization_support,
QuantizationSupport::Int4 | QuantizationSupport::Int8
) {
Self::Lion
} else if config.batch_size >= LARGE_BATCH_THRESHOLD {
Self::Lamb
} else {
Self::Adam
}
}
_ => {
if offloading {
Self::Sgd
} else if config.batch_size >= LARGE_BATCH_THRESHOLD {
Self::Lamb
} else {
Self::Adam
}
}
}
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct HardwareStepReport<A: Float> {
pub applied: bool,
pub learning_rate: A,
pub step_count: usize,
pub accumulated_micro_steps: usize,
}
pub struct OptimizationState<A: Float + 'static, D: Dimension + 'static> {
parameters: Array<A, D>,
optimizer: Box<dyn Optimizer<A, D> + Send + Sync>,
optimizer_kind: HardwareOptimizerKind,
lr_schedule: Box<dyn LearningRateScheduler<A> + Send + Sync>,
base_learning_rate: A,
step_count: usize,
accumulated_micro_steps: usize,
accumulation_steps: usize,
gradient_accumulator: Option<Array<A, D>>,
}
impl<A, D> fmt::Debug for OptimizationState<A, D>
where
A: Float + fmt::Debug + 'static,
D: Dimension + 'static,
{
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("OptimizationState")
.field("parameter_count", &self.parameters.len())
.field("optimizer_kind", &self.optimizer_kind)
.field("base_learning_rate", &self.base_learning_rate)
.field("step_count", &self.step_count)
.field("accumulation_steps", &self.accumulation_steps)
.field("accumulated_micro_steps", &self.accumulated_micro_steps)
.finish()
}
}
impl<A, D> OptimizationState<A, D>
where
A: Float + ScalarOperand + std::fmt::Debug + Send + Sync + 'static,
D: Dimension + 'static,
{
pub fn new(
parameters: Array<A, D>,
kind: HardwareOptimizerKind,
base_learning_rate: A,
accumulation_steps: usize,
) -> Self {
Self {
parameters,
optimizer: build_optimizer(kind, base_learning_rate),
optimizer_kind: kind,
lr_schedule: Box::new(ConstantScheduler::new(base_learning_rate)),
base_learning_rate,
step_count: 0,
accumulated_micro_steps: 0,
accumulation_steps: accumulation_steps.max(1),
gradient_accumulator: None,
}
}
pub fn step(&mut self, gradients: &Array<A, D>) -> Result<HardwareStepReport<A>> {
if gradients.raw_dim() != self.parameters.raw_dim() {
return Err(OptimError::DimensionMismatch(format!(
"hardware-aware step: parameters have shape {:?} but the gradient has shape {:?}",
self.parameters.raw_dim().slice(),
gradients.raw_dim().slice()
)));
}
let effective_gradient = if self.accumulation_steps == 1 {
gradients.clone()
} else {
let shape = self.parameters.raw_dim();
let accumulator = self
.gradient_accumulator
.get_or_insert_with(|| Array::zeros(shape));
for (slot, &g) in accumulator.iter_mut().zip(gradients.iter()) {
*slot = *slot + g;
}
self.accumulated_micro_steps += 1;
if self.accumulated_micro_steps < self.accumulation_steps {
return Ok(HardwareStepReport {
applied: false,
learning_rate: self.lr_schedule.get_learning_rate(),
step_count: self.step_count,
accumulated_micro_steps: self.accumulated_micro_steps,
});
}
let window = scalar_or(self.accumulated_micro_steps, A::one());
let averaged = accumulator.mapv(|g| g / window);
accumulator.fill(A::zero());
self.accumulated_micro_steps = 0;
averaged
};
let learning_rate = self.lr_schedule.get_learning_rate();
self.optimizer.set_learning_rate(learning_rate);
self.parameters = self.optimizer.step(&self.parameters, &effective_gradient)?;
self.step_count += 1;
self.lr_schedule.step();
Ok(HardwareStepReport {
applied: true,
learning_rate,
step_count: self.step_count,
accumulated_micro_steps: 0,
})
}
pub fn parameters(&self) -> &Array<A, D> {
&self.parameters
}
pub fn step_count(&self) -> usize {
self.step_count
}
pub fn optimizer_kind(&self) -> HardwareOptimizerKind {
self.optimizer_kind
}
pub fn learning_rate(&self) -> A {
self.lr_schedule.get_learning_rate()
}
pub fn base_learning_rate(&self) -> A {
self.base_learning_rate
}
pub fn accumulation_steps(&self) -> usize {
self.accumulation_steps
}
pub fn accumulated_micro_steps(&self) -> usize {
self.accumulated_micro_steps
}
pub fn set_accumulation_steps(&mut self, accumulation_steps: usize) {
self.accumulation_steps = accumulation_steps.max(1);
if let Some(accumulator) = self.gradient_accumulator.as_mut() {
accumulator.fill(A::zero());
}
self.accumulated_micro_steps = 0;
}
pub fn set_lr_scheduler(&mut self, schedule: Box<dyn LearningRateScheduler<A> + Send + Sync>) {
self.lr_schedule = schedule;
}
pub fn rebuild_optimizer(&mut self, kind: HardwareOptimizerKind) {
self.optimizer = build_optimizer(kind, self.lr_schedule.get_learning_rate());
self.optimizer_kind = kind;
}
pub fn set_optimizer(
&mut self,
kind: HardwareOptimizerKind,
optimizer: Box<dyn Optimizer<A, D> + Send + Sync>,
) {
self.optimizer = optimizer;
self.optimizer_kind = kind;
}
}
fn build_optimizer<A, D>(
kind: HardwareOptimizerKind,
learning_rate: A,
) -> Box<dyn Optimizer<A, D> + Send + Sync>
where
A: Float + ScalarOperand + std::fmt::Debug + Send + Sync + 'static,
D: Dimension + 'static,
{
match kind {
HardwareOptimizerKind::Sgd => Box::new(SGD::new_with_config(
learning_rate,
scalar_or(0.9, A::zero()),
A::zero(),
)),
HardwareOptimizerKind::Lion => Box::new(Lion::new(learning_rate)),
HardwareOptimizerKind::Adam => Box::new(Adam::new(learning_rate)),
HardwareOptimizerKind::Lamb => Box::new(LAMB::new(learning_rate)),
}
}
pub(super) fn accumulation_steps_for(strategy: &MemoryStrategy) -> usize {
match strategy {
MemoryStrategy::GradientAccumulation { accumulation_steps } => (*accumulation_steps).max(1),
MemoryStrategy::Mixed { strategies, .. } => strategies
.iter()
.map(accumulation_steps_for)
.max()
.unwrap_or(1),
_ => 1,
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::schedulers::ExponentialDecay;
use scirs2_core::ndarray::{Array1, Ix1};
fn quadratic_gradient(parameters: &Array1<f64>) -> Array1<f64> {
parameters.mapv(|x| 2.0 * x)
}
fn quadratic_loss(parameters: &Array1<f64>) -> f64 {
parameters.iter().map(|&x| x * x).sum()
}
#[test]
fn steps_reduce_a_quadratic_loss() {
for kind in [
HardwareOptimizerKind::Sgd,
HardwareOptimizerKind::Lion,
HardwareOptimizerKind::Adam,
HardwareOptimizerKind::Lamb,
] {
let start = Array1::from_vec(vec![1.0, -2.0, 3.0]);
let initial_loss = quadratic_loss(&start);
let mut state: OptimizationState<f64, Ix1> =
OptimizationState::new(start, kind, 0.05, 1);
for _ in 0..200 {
let gradient = quadratic_gradient(state.parameters());
let report = state.step(&gradient).expect("step must succeed");
assert!(report.applied, "{} did not apply an update", kind.name());
}
let final_loss = quadratic_loss(state.parameters());
assert_eq!(state.step_count(), 200, "{}", kind.name());
assert!(
final_loss < initial_loss * 0.5,
"{}: loss did not decrease ({initial_loss} -> {final_loss})",
kind.name()
);
}
}
#[test]
fn mismatched_gradient_shape_is_reported() {
let mut state: OptimizationState<f64, Ix1> = OptimizationState::new(
Array1::from_vec(vec![1.0, 2.0]),
HardwareOptimizerKind::Adam,
0.01,
1,
);
let error = state
.step(&Array1::from_vec(vec![1.0, 2.0, 3.0]))
.expect_err("a shape mismatch must be reported");
assert!(
matches!(error, OptimError::DimensionMismatch(_)),
"{error:?}"
);
}
#[test]
fn gradient_accumulation_updates_once_per_window() {
let mut state: OptimizationState<f64, Ix1> = OptimizationState::new(
Array1::from_vec(vec![0.0, 0.0]),
HardwareOptimizerKind::Sgd,
0.1,
3,
);
let gradient = Array1::from_vec(vec![1.0, 1.0]);
for micro in 1..=2 {
let report = state.step(&gradient).expect("accumulating step");
assert!(!report.applied, "micro-batch {micro} must not update");
assert_eq!(report.accumulated_micro_steps, micro);
assert_eq!(state.parameters()[0], 0.0);
}
let report = state.step(&gradient).expect("closing step");
assert!(report.applied, "the full window must apply an update");
assert_eq!(state.step_count(), 1);
assert!(state.parameters()[0] < 0.0);
}
#[test]
fn accumulation_applies_the_window_mean() {
let gradient = Array1::from_vec(vec![1.0, -0.5]);
let mut direct: OptimizationState<f64, Ix1> = OptimizationState::new(
Array1::from_vec(vec![0.0, 0.0]),
HardwareOptimizerKind::Sgd,
0.1,
1,
);
direct.step(&gradient).expect("direct step");
let mut accumulated: OptimizationState<f64, Ix1> = OptimizationState::new(
Array1::from_vec(vec![0.0, 0.0]),
HardwareOptimizerKind::Sgd,
0.1,
3,
);
for _ in 0..3 {
accumulated.step(&gradient).expect("accumulated step");
}
for (index, (&direct_value, &accumulated_value)) in direct
.parameters()
.iter()
.zip(accumulated.parameters().iter())
.enumerate()
{
assert!(
(direct_value - accumulated_value).abs() < 1e-12,
"coordinate {index}: {direct_value} != {accumulated_value}"
);
}
}
#[test]
fn the_schedule_drives_the_optimizer_learning_rate() {
let mut state: OptimizationState<f64, Ix1> = OptimizationState::new(
Array1::from_vec(vec![1.0]),
HardwareOptimizerKind::Sgd,
0.1,
1,
);
state.set_lr_scheduler(Box::new(ExponentialDecay::new(0.1, 0.5, 1)));
assert!((state.learning_rate() - 0.1).abs() < 1e-12);
assert!((state.base_learning_rate() - 0.1).abs() < 1e-12);
let first = state
.step(&Array1::from_vec(vec![1.0]))
.expect("first step");
assert!((first.learning_rate - 0.1).abs() < 1e-12);
let second = state
.step(&Array1::from_vec(vec![1.0]))
.expect("second step");
assert!(
second.learning_rate < first.learning_rate,
"the schedule did not decay: {} -> {}",
first.learning_rate,
second.learning_rate
);
assert!(
(state.base_learning_rate() - 0.1).abs() < 1e-12,
"the construction-time rate must stay put so the decay is measurable"
);
}
#[test]
fn each_family_reports_its_optimizer_state_footprint() {
assert_eq!(HardwareOptimizerKind::Sgd.state_buffers_per_parameter(), 1);
assert_eq!(HardwareOptimizerKind::Lion.state_buffers_per_parameter(), 1);
assert_eq!(HardwareOptimizerKind::Adam.state_buffers_per_parameter(), 2);
assert_eq!(HardwareOptimizerKind::Lamb.state_buffers_per_parameter(), 2);
assert!(
HardwareOptimizerKind::Lion.state_buffers_per_parameter()
< HardwareOptimizerKind::Adam.state_buffers_per_parameter(),
"Lion is recommended for edge devices precisely because it is cheaper"
);
}
#[test]
fn optimizer_recommendation_follows_the_platform() {
let edge = HardwarePlatform::Edge {
power_budget: 2.0,
memory_limit: 256 * 1024 * 1024,
quantization_support: QuantizationSupport::Int8,
};
let mut config: HardwareOptimizationConfig<f64> = HardwareOptimizationConfig {
batch_size: 16,
memory_strategy: MemoryStrategy::Standard,
parallelization: super::super::ParallelizationStrategy::SingleThread,
precision: super::super::PrecisionStrategy::FP32,
optimizer_params: std::collections::HashMap::new(),
communication: None,
};
assert_eq!(
HardwareOptimizerKind::recommend_for(&edge, &config),
HardwareOptimizerKind::Sgd
);
let roomy_edge = HardwarePlatform::Edge {
power_budget: 30.0,
memory_limit: 4 * 1024 * 1024 * 1024,
quantization_support: QuantizationSupport::Int8,
};
assert_eq!(
HardwareOptimizerKind::recommend_for(&roomy_edge, &config),
HardwareOptimizerKind::Lion
);
let gpu = HardwarePlatform::GPU {
memory: 16 * 1024 * 1024 * 1024,
compute_units: 80,
memory_bandwidth: 900.0,
architecture: super::super::GPUArchitecture::Ampere,
};
config.batch_size = 128;
assert_eq!(
HardwareOptimizerKind::recommend_for(&gpu, &config),
HardwareOptimizerKind::Adam
);
config.batch_size = 4096;
assert_eq!(
HardwareOptimizerKind::recommend_for(&gpu, &config),
HardwareOptimizerKind::Lamb
);
config.memory_strategy = MemoryStrategy::CPUOffloading { offload_ratio: 0.8 };
assert_eq!(
HardwareOptimizerKind::recommend_for(&gpu, &config),
HardwareOptimizerKind::Sgd
);
}
#[test]
fn accumulation_window_survives_a_mixed_memory_strategy() {
assert_eq!(accumulation_steps_for(&MemoryStrategy::Standard), 1);
assert_eq!(
accumulation_steps_for(&MemoryStrategy::GradientAccumulation {
accumulation_steps: 0
}),
1,
"a zero window would mean the parameters never move"
);
assert_eq!(
accumulation_steps_for(&MemoryStrategy::Mixed {
strategies: vec![
MemoryStrategy::Standard,
MemoryStrategy::GradientAccumulation {
accumulation_steps: 4
},
],
strategy_weights: vec![0.5, 0.5],
}),
4
);
}
}